From 5aa2c41e8d1df84ead2b7f3d008dc21ea1faff75 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 12:07:46 +0200 Subject: [PATCH 01/16] Credit Vincent Lovero for his ARM SME kernel work --- CONTRIBUTORS.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/CONTRIBUTORS.md b/CONTRIBUTORS.md index 39ae88cf7a..fc612cfbd4 100644 --- a/CONTRIBUTORS.md +++ b/CONTRIBUTORS.md @@ -283,4 +283,6 @@ hheei * Aadityansha Verma * [2026-07-14] Add independent transpose support for C in GEADD (sgeadd/dgeadd/cgeadd/zgeadd). +* Vincent Lovero + * [2026-08-11] ARM v9.2 SME GEMM kernels for Apple M From 2f915c5e230a7930e122b0303b1f59977b47f15d Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 12:10:03 +0200 Subject: [PATCH 02/16] Add f64f64 extension to VortexM4 options --- Makefile.arm64 | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Makefile.arm64 b/Makefile.arm64 index 68aabb3dc4..bfa0199bbe 100644 --- a/Makefile.arm64 +++ b/Makefile.arm64 @@ -310,7 +310,7 @@ endif ifeq ($(CORE), VORTEXM4) ifneq ($(C_COMPILER), GCC) -CCOMMON_OPT += -march=armv8.4-a+sme +CCOMMON_OPT += -march=armv8.4-a+sme+sme-f64f64 #ifneq ($(APPLECLANG),1) #override LDFLAGS += -lclang_rt_builtins-aarch64 #endif From f2dc74796f4b47b75b00d772cc44799e0b87eb2d Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 12:10:47 +0200 Subject: [PATCH 03/16] Integrate SME GEMM kernels --- interface/gemm.c | 71 +++++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 67 insertions(+), 4 deletions(-) diff --git a/interface/gemm.c b/interface/gemm.c index 2aba7bed6c..71bf12ac38 100644 --- a/interface/gemm.c +++ b/interface/gemm.c @@ -45,6 +45,12 @@ #include "functable.h" #endif +#ifdef ARCH_ARM64 +void sme_SGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const float *alpha, const float *a, const BLASLONG *lda, const float *b, const BLASLONG *ldb, const float *beta, float *c, const BLASLONG *ldc); +void sme_DGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const double *alpha, const double *a, const BLASLONG *lda, const double *b, const BLASLONG *ldb, const double *beta, double *c, const BLASLONG *ldc); +void sme_CGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const float _Complex alpha, const float *a, const BLASLONG *lda, const float *b, const BLASLONG *ldb, const float _Complex beta, float *c, const BLASLONG *ldc); +void sme_ZGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const double _Complex alpha, const double *a, const BLASLONG *lda, const double *b, const BLASLONG *ldb, const double _Complex beta, double *c, const BLASLONG *ldc); +#endif #ifndef COMPLEX #define SMP_THRESHOLD_MIN 65536.0 #ifdef XDOUBLE @@ -268,7 +274,6 @@ void NAME(char *TRANSA, char *TRANSB, int transa, transb, nrowa, nrowb; blasint info; - int order = -1; char transA, transB; IFLOAT *buffer; @@ -346,7 +351,6 @@ void NAME(char *TRANSA, char *TRANSB, if (transB == 'R') transb = 2; if (transB == 'C') transb = 3; #endif - nrowa = args.m; if (transa & 1) nrowa = args.k; nrowb = args.k; @@ -562,8 +566,8 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 || strcmp(gotoblas_corename(), "vortexm4") == 0 #endif ) -// if (support_sme1()) #endif + if (order == CblasRowMajor && k==lda && n==ldb && n==ldc && beta == 0 && alpha == 1.0 && TransA == CblasNoTrans && TransB == CblasNoTrans && SGEMM_DIRECT_PERFORMANT(m,n,k)) { SGEMM_DIRECT(m, n, k, a, lda, b, ldb, c, ldc); return; @@ -574,10 +578,69 @@ else return; } +#endif //defined arm64 +#endif //defined complex + + +#endif //defined CBLAS + +#if !defined(BFLOAT16) && !defined(HFLOAT16) +#if defined(ARCH_ARM64) && (defined(USE_SGEMM_KERNEL_DIRECT)||defined(DYNAMIC_ARCH)) +#if defined(DYNAMIC_ARCH) +if (strcmp(gotoblas_corename(), "armv9sme") == 0 +#if defined(__clang__) + || strcmp(gotoblas_corename(), "vortexm4") == 0 #endif +) +#endif //defined dynarch +{ +char* TA,*TB; + if (transa & 1) + TA = "T"; + else + TA= "N"; + if (transb & 1) + TB = "T"; + else + TB= "N"; +#ifndef COMPLEX + if (transa == 3) + TA= "T"; + if (transb == 3) + TB= "T"; +FLOAT* al=(FLOAT*)args.alpha; +FLOAT* be=(FLOAT*)args.beta; +#ifndef DOUBLE + sme_SGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, al, args.a, &args.lda, args.b, &args.ldb, be, args.c, &args.ldc); +#else + sme_DGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, al, args.a, &args.lda, args.b, &args.ldb, be, args.c, &args.ldc); +#endif +#else + if (transa == 2) + TA= "R"; + if (transb == 2) + TB= "R"; + if (transa == 3) + TA= "C"; + if (transb == 3) + TB= "C"; +FLOAT* al=(FLOAT*)args.alpha; +FLOAT* be=(FLOAT*)args.beta; +#ifndef DOUBLE +float _Complex c_al={al[0],al[1]}; +float _Complex c_be={be[0],be[1]}; + sme_CGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, c_al, args.a, &args.lda, args.b, &args.ldb, c_be, args.c, &args.ldc); +#else +double _Complex c_al={al[0],al[1]}; +double _Complex c_be={be[0],be[1]}; + sme_ZGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, c_al, args.a, &args.lda, args.b, &args.ldb, c_be, args.c, &args.ldc); #endif - #endif + return; +} +#endif //defined arm64 + +#endif //defined b/hfloat16 #if defined(__linux__) && defined(__x86_64__) && defined(BFLOAT16) #if defined(DYNAMIC_ARCH) From 78f03216de33c5959ec11a0f306d87d63eec42d9 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 12:12:29 +0200 Subject: [PATCH 04/16] Add SME GEMM kernels ported from vlovero's ARMv9.2-GEMM project --- kernel/arm64/sme_cgemm_kernel.c | 546 ++++++++++++++ kernel/arm64/sme_dgemm_kernel.c | 1206 +++++++++++++++++++++++++++++++ kernel/arm64/sme_sgemm_kernel.c | 651 +++++++++++++++++ kernel/arm64/sme_zgemm_kernel.c | 547 ++++++++++++++ 4 files changed, 2950 insertions(+) create mode 100644 kernel/arm64/sme_cgemm_kernel.c create mode 100644 kernel/arm64/sme_dgemm_kernel.c create mode 100644 kernel/arm64/sme_sgemm_kernel.c create mode 100644 kernel/arm64/sme_zgemm_kernel.c diff --git a/kernel/arm64/sme_cgemm_kernel.c b/kernel/arm64/sme_cgemm_kernel.c new file mode 100644 index 0000000000..e2622fcda4 --- /dev/null +++ b/kernel/arm64/sme_cgemm_kernel.c @@ -0,0 +1,546 @@ +//#include +#include +//#include +#include +#include +#include +#include "common.h" +#ifndef stdmin +#define stdmin(a,b) (a>b? b:a) +#endif +typedef float _Complex cfloat; +cfloat CMUL(cfloat a, cfloat b,bool conja, bool conjb) { +float ra=creal(a); +float rb=creal(b); +float ia=conja ? -cimag(a) : cimag(a); +float ib=conjb ? -cimag(b) : cimag(b); +float r1=ra*rb; +float r2=ia*ib; +float r=r1-r2; +float i=(ra+ia)*(rb+ib)-r1-r2; +cfloat res={r,i}; +return res; +} +#define KERNEL_ALPHA 0 +#define USE_VECTORIZED_PACKING 1 + +//using cfloat = std::complex; + +cfloat czero={0.,0.}; +cfloat cone={1.,0.}; +#define MC 256 +#define KC 512 +#define NC 1024 + +static inline void cgemm_sme_compute_16x16_tile(blasint current_K, const float *A_ptr, const float *B_ptr, cfloat *C_ptr, size_t ldc, int beta_mode, const cfloat *beta_ptr) +{ + size_t ldc_bytes = ldc * sizeof(cfloat); + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + + asm volatile("smstart\n\t" + "ptrue p0.s\n\t" + "zero {za}\n\t" + "cmp %w[beta_mode], #0\n\t" + "b.eq 19f\n\t" + + // FIX 1: Correct SVE mnemonics for loading 32-bit floats (Real/Imag) + "cmp %w[beta_mode], #1\n\t" + "b.ne 10f\n\t" + "ld1rw z30.s, p0/z, [%[beta_ptr]]\n\t" // Load beta.real() + "add x15, %[beta_ptr], #4\n\t" // 4-byte offset for float + "ld1rw z31.s, p0/z, [x15]\n\t" // Load beta.imag() + "10:\n\t" + + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "100:\n\t" + "ld1w z0.s, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1w z1.s, p0/z, [x14]\n\t" + "uzp1 z2.s, z0.s, z1.s\n\t" // z2 = C_re + "uzp2 z3.s, z0.s, z1.s\n\t" // z3 = C_im + + "cmp %w[beta_mode], #1\n\t" + "b.ne 101f\n\t" + + // FIX 2: Fully implemented Complex Beta Multiplication using movprfx + "movprfx z4, z2\n\t" + "fmul z4.s, p0/m, z4.s, z30.s\n\t" // z4 = Cre * Bre + "movprfx z5, z3\n\t" + "fmul z5.s, p0/m, z5.s, z31.s\n\t" // z5 = Cim * Bim + "movprfx z6, z4\n\t" + "fsub z6.s, p0/m, z6.s, z5.s\n\t" // z6 = Cre' (Cre*Bre - Cim*Bim) + + "movprfx z8, z2\n\t" + "fmul z8.s, p0/m, z8.s, z31.s\n\t" // z8 = Cre * Bim + "movprfx z9, z3\n\t" + "fmul z9.s, p0/m, z9.s, z30.s\n\t" // z9 = Cim * Bre + "movprfx z7, z8\n\t" + "fadd z7.s, p0/m, z7.s, z9.s\n\t" // z7 = Cim' (Cre*Bim + Cim*Bre) + + "mova za0v.s[w12, 0], p0/m, z6.s\n\t" + "mova za2v.s[w12, 0], p0/m, z7.s\n\t" + "b 102f\n\t" + + "101:\n\t" // Fallback: beta == 1.0 + "mova za0v.s[w12, 0], p0/m, z2.s\n\t" + "mova za2v.s[w12, 0], p0/m, z3.s\n\t" + + "102:\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #16\n\t" + "b.ne 100b\n\t" + + "19:\n\t" + "mov w10, %w[k]\n\t" + "cbz w10, 3f\n\t" + "11:\n\t" + "ld1w z0.s, p0/z, [%[a], #0, mul vl]\n\t" + "ld1w z1.s, p0/z, [%[a], #1, mul vl]\n\t" + "ld1w z2.s, p0/z, [%[b], #0, mul vl]\n\t" + "ld1w z3.s, p0/z, [%[b], #1, mul vl]\n\t" + "fmopa za0.s, p0/m, p0/m, z0.s, z2.s\n\t" + "fmopa za1.s, p0/m, p0/m, z1.s, z3.s\n\t" + "fmopa za2.s, p0/m, p0/m, z0.s, z3.s\n\t" + "fmopa za3.s, p0/m, p0/m, z1.s, z2.s\n\t" + "add %[a], %[a], #128\n\t" + "add %[b], %[b], #128\n\t" + "subs w10, w10, #1\n\t" + "b.ne 11b\n\t" + + "3:\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "200:\n\t" + "mova z0.s, p0/m, za0v.s[w12, 0]\n\t" + "mova z1.s, p0/m, za1v.s[w12, 0]\n\t" + "mova z2.s, p0/m, za2v.s[w12, 0]\n\t" + "mova z3.s, p0/m, za3v.s[w12, 0]\n\t" + + // FIX 3: Non-destructive SVE arithmetic on store + "movprfx z4, z0\n\t" + "fsub z4.s, p0/m, z4.s, z1.s\n\t" + "movprfx z5, z2\n\t" + "fadd z5.s, p0/m, z5.s, z3.s\n\t" + + "zip1 z6.s, z4.s, z5.s\n\t" + "zip2 z7.s, z4.s, z5.s\n\t" + "st1w z6.s, p0, [x13]\n\t" + "add x14, x13, #64\n\t" + "st1w z7.s, p0, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #16\n\t" + "b.ne 200b\n\t" + "smstop\n\t" + : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) + // Updated Clobber list to cover x15 and z8-z9, z30-z31 + : "p0", "x10", "w12", "x13", "x14", "x15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z30", "z31", "za" ); + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + +} + +void cgemm_sme_NN(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +{ + + if (alpha == czero || K == 0) { + if (beta == czero) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = czero; + } + } + } else { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] *= beta; + } + } + } + return; + } + + + + alignas(256) /*thread_local*/ static float A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta==czero) ? 0 : (beta != cone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 16) { + float *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + for (int bc = 0; bc < 16; ++bc) { + cfloat val = (jj + bc < current_N) ? b[(j + jj + bc) * ldb + k + kk] : czero; + B_out[kk * 32 + bc] = creal(val); + B_out[kk * 32 + 16 + bc] = conjb ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + size_t panel_stride_A = 32 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 16) { + float *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + for (int br = 0; br < 16; ++br) { + cfloat val = (ii + br < current_M) ? CMUL(A[(k + kk) * lda + i + ii + br] , alpha, conja, 0) : czero; + A_out[kk * 32 + br] = creal(val); + A_out[kk * 32 + 16 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + float *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + float *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + cfloat *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 16 && current_N_block == 16) { + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) cfloat C_buffer[256]; + if (beta_mode == 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = czero; + } + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +void cgemm_sme_TN(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +{ + if (alpha == czero || K == 0) { + if (beta == czero) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = czero; + } + } + } else { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = CMUL(C[i+j*ldc],beta,0,0); + } + } + } + return; + } + alignas(256) /*thread_local*/ static float A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta==czero) ? 0 : (beta != cone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 16) { + float *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + for (int bc = 0; bc < 16; ++bc) { + cfloat val = (jj + bc < current_N) ? b[(j + jj + bc) * ldb + k + kk] : czero; + B_out[kk * 32 + bc] = creal(val); + B_out[kk * 32 + 16 + bc] = conjb ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + size_t panel_stride_A = 32 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 16) { + float *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + for (int br = 0; br < 16; ++br) { + cfloat val = (ii + br < current_M) ? CMUL(A[(i + ii + br) * lda + k + kk],alpha,conja,0) : czero; + A_out[kk * 32 + br] = creal(val); + A_out[kk * 32 + 16 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + float *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + float *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + cfloat *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 16 && current_N_block == 16) { + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) cfloat C_buffer[256]; + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = czero; + } + if (beta_mode != 0) { + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + + } + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +void cgemm_sme_NT(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +{ + if (alpha == czero || K == 0) { + if (beta == czero) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = czero; + } + } + } else { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] *= beta; + } + } + } + return; + } + alignas(256) /*thread_local*/ static float A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta==czero) ? 0 : (beta != cone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 16) { + float *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; kk++) { + for (int bc = 0; bc < 16; ++bc) { + cfloat val = (jj + bc < current_N) ? b[(k + kk) * ldb + j + jj + bc] : czero ; + B_out[kk * 32 + bc] = creal(val); + B_out[kk * 32 + 16 + bc] = (conjb) ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + size_t panel_stride_A = 32 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 16) { + float *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; kk++) { + for (int br = 0; br < 16; ++br) { + cfloat val = (ii + br < current_M) ? CMUL(A[(k + kk) * lda + i + ii + br],alpha,conja,0) : czero; + A_out[kk * 32 + br] = creal(val); + A_out[kk * 32 + 16 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + float *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + float *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + cfloat *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 16 && current_N_block == 16) { + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) cfloat C_buffer[256]; + if (beta_mode != 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = czero; + } + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +void cgemm_sme_TT(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +{ + if (alpha == czero || K == 0) { + if (beta == czero) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = czero; + } + } + } else { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] *= beta; + } + } + } + return; + } + alignas(256) /*thread_local*/ static float A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == czero) ? 0 : (beta != cone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 16) { + float *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; kk++) { + for (int bc = 0; bc < 16; ++bc) { + cfloat val = (jj + bc < current_N) ? b[(k + kk) * ldb + j + jj + bc] : czero; + B_out[kk * 32 + bc] = creal(val); + B_out[kk * 32 + 16 + bc] = (conjb) ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + size_t panel_stride_A = 32 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 16) { + float *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; kk++) { + for (int br = 0; br < 16; ++br) { + cfloat val = (ii + br < current_M) ? CMUL( A[(i + ii + br) * lda + k + kk],alpha,conja,0) : czero; + A_out[kk * 32 + br] = creal(val); + A_out[kk * 32 + 16 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + float *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + float *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + cfloat *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 16 && current_N_block == 16) { + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) cfloat C_buffer[256]; + if (beta_mode != 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = czero; + } + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + + cgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +void sme_CGEMM_KERNEL(const char *transa, const char *transb, const blasint *m, const blasint *n, const blasint *k, const cfloat alpha, const cfloat *a, const blasint *lda, const cfloat *b, const blasint *ldb, const cfloat beta, cfloat *c, const blasint *ldc) +{ + bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); + bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); + if (!trans_a && !trans_b) { + bool conja=(*transa == 'R' || *transa == 'r'); + bool conjb=(*transb == 'R' || *transb == 'r'); + cgemm_sme_NN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } + else if (trans_a && !trans_b) { + bool conja=(*transa == 'C' || *transa == 'c'); + bool conjb=(*transb == 'R' || *transb == 'r'); + cgemm_sme_TN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } + else if (!trans_a && trans_b) { + bool conja=(*transa == 'R' || *transa == 'r'); + bool conjb=(*transb == 'C' || *transb == 'c'); + cgemm_sme_NT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } + else { + bool conja=(*transa == 'C' || *transa == 'c'); + bool conjb=(*transb == 'C' || *transb == 'c'); + cgemm_sme_TT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } +} diff --git a/kernel/arm64/sme_dgemm_kernel.c b/kernel/arm64/sme_dgemm_kernel.c new file mode 100644 index 0000000000..7333c668c6 --- /dev/null +++ b/kernel/arm64/sme_dgemm_kernel.c @@ -0,0 +1,1206 @@ +//#include +#include +//#include +#include +#include +#include +#include +#include "common.h" +#ifndef stdmin +#define stdmin(a,b) (a>b? b:a) +#endif + +#define KERNEL_ALPHA 0 +#define USE_VECTORIZED_PACKING 1 + + +// ======================================================================== +// OPTIMIZED CACHE BLOCKING PARAMETERS +// ======================================================================== +#define MC 256 +#define KC 512 +#define NC 1024 + +// ======================================================================== +// 16x16 PURE-SME MACRO-KERNEL (Fused Beta Scaling) +// ======================================================================== +static inline void dgemm_sme_compute_16x16_tile(int current_K, const double *A_ptr, const double *B_ptr, double *C_ptr, ptrdiff_t ldc, int beta_mode, const double *beta_ptr) +{ + ptrdiff_t ldc_bytes = ldc * sizeof(double); + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + + asm volatile( + "smstart\n\t" + // Enable all 64-bit (double precision) lanes in predicate register p0 + "ptrue p0.d\n\t" + + // ========================================================= + // PHASE 1: INITIALIZE 'ZA' ACCUMULATOR + // ========================================================= + "cmp %w[beta_mode], #0\n\t" + "b.eq 10f\n\t" // Jump to zero_za + "cmp %w[beta_mode], #1\n\t" + "b.eq 11f\n\t" // Jump to scale_beta + "b 12f\n\t" // Jump to direct_load (default) + + // --- PATH 0: beta == 0.0 (Instant Zero) --- + "10:\n\t" + "zero {za}\n\t" + "b 19f\n\t" // Done with Phase 1 + + // --- PATH 1: beta != 1.0 (Vectorized Load & Scale) --- + "11:\n\t" + "ld1rd z31.d, p0/z, [%[beta_ptr]]\n\t" // Load and replicate beta into z31 + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + + "110:\n\t" + "ld1d z0.d, p0/z, [x13]\n\t" + "fmul z0.d, p0/m, z0.d, z31.d\n\t" + "mova za0v.d[w12, 0], p0/m, z0.d\n\t" + + "add x14, x13, #64\n\t" + "ld1d z1.d, p0/z, [x14]\n\t" + "fmul z1.d, p0/m, z1.d, z31.d\n\t" + "mova za2v.d[w12, 0], p0/m, z1.d\n\t" + + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #8\n\t" + "b.ne 110b\n\t" + + "mov w15, #0\n\t" + "111:\n\t" + "ld1d z0.d, p0/z, [x13]\n\t" + "fmul z0.d, p0/m, z0.d, z31.d\n\t" + "mova za1v.d[w15, 0], p0/m, z0.d\n\t" + + "add x14, x13, #64\n\t" + "ld1d z1.d, p0/z, [x14]\n\t" + "fmul z1.d, p0/m, z1.d, z31.d\n\t" + "mova za3v.d[w15, 0], p0/m, z1.d\n\t" + + "add x13, x13, %[ldc_bytes]\n\t" + "add w15, w15, #1\n\t" + "cmp w15, #8\n\t" + "b.ne 111b\n\t" + "b 19f\n\t" // Done with Phase 1 + + // --- PATH 2: k > 0 or beta == 1.0 (Standard Direct Load) --- + "12:\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + + "100:\n\t" + "ld1d {za0v.d[w12, 0]}, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1d {za2v.d[w12, 0]}, p0/z, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #8\n\t" + "b.ne 100b\n\t" + + "mov w15, #0\n\t" + "101:\n\t" + "ld1d {za1v.d[w15, 0]}, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1d {za3v.d[w15, 0]}, p0/z, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w15, w15, #1\n\t" + "cmp w15, #8\n\t" + "b.ne 101b\n\t" + + "19:\n\t" // Entry point for Phase 2 + + // ========================================================= + // PHASE 2: COMPUTE (ZA += A * B) + // ========================================================= + "mov w10, %w[k]\n\t" + "cbz w10, 3f\n\t" + + "lsr w11, w10, #2\n\t" + "cbz w11, 12f\n\t" + + "11:\n\t" + "ld1d z0.d, p0/z, [%[a], #0, mul vl]\n\t" + "ld1d z1.d, p0/z, [%[a], #1, mul vl]\n\t" + "ld1d z2.d, p0/z, [%[b], #0, mul vl]\n\t" + "ld1d z3.d, p0/z, [%[b], #1, mul vl]\n\t" + + "ld1d z4.d, p0/z, [%[a], #2, mul vl]\n\t" + "ld1d z5.d, p0/z, [%[a], #3, mul vl]\n\t" + "ld1d z6.d, p0/z, [%[b], #2, mul vl]\n\t" + "ld1d z7.d, p0/z, [%[b], #3, mul vl]\n\t" + + "ld1d z8.d, p0/z, [%[a], #4, mul vl]\n\t" + "ld1d z9.d, p0/z, [%[a], #5, mul vl]\n\t" + "ld1d z10.d, p0/z, [%[b], #4, mul vl]\n\t" + "ld1d z11.d, p0/z, [%[b], #5, mul vl]\n\t" + + "ld1d z12.d, p0/z, [%[a], #6, mul vl]\n\t" + "ld1d z13.d, p0/z, [%[a], #7, mul vl]\n\t" + "ld1d z14.d, p0/z, [%[b], #6, mul vl]\n\t" + "ld1d z15.d, p0/z, [%[b], #7, mul vl]\n\t" + + "fmopa za0.d, p0/m, p0/m, z0.d, z2.d\n\t" + "fmopa za1.d, p0/m, p0/m, z0.d, z3.d\n\t" + "fmopa za2.d, p0/m, p0/m, z1.d, z2.d\n\t" + "fmopa za3.d, p0/m, p0/m, z1.d, z3.d\n\t" + + "fmopa za0.d, p0/m, p0/m, z4.d, z6.d\n\t" + "fmopa za1.d, p0/m, p0/m, z4.d, z7.d\n\t" + "fmopa za2.d, p0/m, p0/m, z5.d, z6.d\n\t" + "fmopa za3.d, p0/m, p0/m, z5.d, z7.d\n\t" + + "fmopa za0.d, p0/m, p0/m, z8.d, z10.d\n\t" + "fmopa za1.d, p0/m, p0/m, z8.d, z11.d\n\t" + "fmopa za2.d, p0/m, p0/m, z9.d, z10.d\n\t" + "fmopa za3.d, p0/m, p0/m, z9.d, z11.d\n\t" + + "fmopa za0.d, p0/m, p0/m, z12.d, z14.d\n\t" + "fmopa za1.d, p0/m, p0/m, z12.d, z15.d\n\t" + "fmopa za2.d, p0/m, p0/m, z13.d, z14.d\n\t" + "fmopa za3.d, p0/m, p0/m, z13.d, z15.d\n\t" + + "add %[a], %[a], #512\n\t" + "add %[b], %[b], #512\n\t" + "subs w11, w11, #1\n\t" + "b.ne 11b\n\t" + + // --- Remainder Loop --- + "12:\n\t" + "and w10, w10, #3\n\t" + "cbz w10, 3f\n\t" + + "13:\n\t" + "ld1d z0.d, p0/z, [%[a], #0, mul vl]\n\t" + "ld1d z1.d, p0/z, [%[a], #1, mul vl]\n\t" + "ld1d z2.d, p0/z, [%[b], #0, mul vl]\n\t" + "ld1d z3.d, p0/z, [%[b], #1, mul vl]\n\t" + "fmopa za0.d, p0/m, p0/m, z0.d, z2.d\n\t" + "fmopa za1.d, p0/m, p0/m, z0.d, z3.d\n\t" + "fmopa za2.d, p0/m, p0/m, z1.d, z2.d\n\t" + "fmopa za3.d, p0/m, p0/m, z1.d, z3.d\n\t" + "add %[a], %[a], #128\n\t" + "add %[b], %[b], #128\n\t" + "subs w10, w10, #1\n\t" + "b.ne 13b\n\t" + + // ========================================================= + // PHASE 3: STORE 'ZA' ACCUMULATOR BACK TO MATRIX 'C' + // ========================================================= + "3:\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + + "200:\n\t" + "st1d {za0v.d[w12, 0]}, p0, [x13]\n\t" + "add x14, x13, #64\n\t" + "st1d {za2v.d[w12, 0]}, p0, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #8\n\t" + "b.ne 200b\n\t" + + "mov w15, #0\n\t" + "201:\n\t" + "st1d {za1v.d[w15, 0]}, p0, [x13]\n\t" + "add x14, x13, #64\n\t" + "st1d {za3v.d[w15, 0]}, p0, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w15, w15, #1\n\t" + "cmp w15, #8\n\t" + "b.ne 201b\n\t" + "smstop\n\t" + : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) + : "p0","x10", "x11", "w12", "x13", "x14", "w15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", "z31", "za"); + + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + +} + +// ======================================================================== +// VERSION 0: C = alpha * A * B + beta * C (NN) +// ======================================================================== +void dgemm_sme_NN(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +{ +#if PROFILING + CALI_CXX_MARK_FUNCTION; +#endif + // Hardware bypass for skipped main loops + if (alpha == 0.0 || K == 0) { + if (beta != 1.0) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0) ? 0.0 : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { +#if PROFILING + CALI_MARK_BEGIN("loop_j"); +#endif + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + + for (int k = 0; k < K; k += KC) { +#if PROFILING + CALI_MARK_BEGIN("loop_k"); +#endif + int current_K = stdmin(KC, K - k); + ptrdiff_t panel_stride_B = 16 * (ptrdiff_t)current_K; + + // Route the beta scaling natively inside the SME Assembly + int beta_mode = 2; // Default: K>0, directly load C + if (k == 0) { + if (beta == 0.0) { + beta_mode = 0; + } + else if (beta != 1.0) { + beta_mode = 1; + } + } + +#if PROFILING + CALI_MARK_BEGIN("pack_B"); +#endif +#if USE_VECTORIZED_PACKING + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_in_ptr = &B[(j + jj) * (ptrdiff_t)ldb + k + kk]; + double *__restrict B_out_ptr = &B_out[kk * 16]; + + ptrdiff_t ldb_sz = (ptrdiff_t)ldb; + + B_out_ptr[0] = B_in_ptr[0]; + B_out_ptr[1] = B_in_ptr[1 * ldb_sz]; + B_out_ptr[2] = B_in_ptr[2 * ldb_sz]; + B_out_ptr[3] = B_in_ptr[3 * ldb_sz]; + B_out_ptr[4] = B_in_ptr[4 * ldb_sz]; + B_out_ptr[5] = B_in_ptr[5 * ldb_sz]; + B_out_ptr[6] = B_in_ptr[6 * ldb_sz]; + B_out_ptr[7] = B_in_ptr[7 * ldb_sz]; + B_out_ptr[8] = B_in_ptr[8 * ldb_sz]; + B_out_ptr[9] = B_in_ptr[9 * ldb_sz]; + B_out_ptr[10] = B_in_ptr[10 * ldb_sz]; + B_out_ptr[11] = B_in_ptr[11 * ldb_sz]; + B_out_ptr[12] = B_in_ptr[12 * ldb_sz]; + B_out_ptr[13] = B_in_ptr[13 * ldb_sz]; + B_out_ptr[14] = B_in_ptr[14 * ldb_sz]; + B_out_ptr[15] = B_in_ptr[15 * ldb_sz]; + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int bc = 0; bc < 16; ++bc) { + if (jj + bc < current_N) { + const double *__restrict B_col = &B[(j + jj + bc) * (ptrdiff_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = B_col[kk]; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = 0.0; + } + } + } + } +#else + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int bc = 0; bc < 16; ++bc) { + const double *__restrict B_col = &B[(j + jj + bc) * (ptrdiff_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = B_col[kk]; + } + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int bc = 0; bc < 16; ++bc) { + if (jj + bc < current_N) { + const double *__restrict B_col = &B[(j + jj + bc) * (ptrdiff_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = B_col[kk]; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = 0.0; + } + } + } + } +#endif +#if PROFILING + CALI_MARK_END("pack_B"); +#endif + + for (int i = 0; i < M; i += MC) { +#if PROFILING + CALI_MARK_BEGIN("loop_i"); +#endif + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + ptrdiff_t panel_stride_A = 16 * (ptrdiff_t)current_K; + +#if PROFILING + CALI_MARK_BEGIN("pack_A"); +#endif +#if USE_VECTORIZED_PACKING + float64x2_t valpha = vdupq_n_f64(alpha); + + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + double *__restrict A_out_ptr = &A_out[kk * 16]; + + // Alpha is mathematically confined here (M*K operations exact) + float64x2_t v0 = vmulq_f64(vld1q_f64(&A_col[0]), valpha); + float64x2_t v1 = vmulq_f64(vld1q_f64(&A_col[2]), valpha); + float64x2_t v2 = vmulq_f64(vld1q_f64(&A_col[4]), valpha); + float64x2_t v3 = vmulq_f64(vld1q_f64(&A_col[6]), valpha); + float64x2_t v4 = vmulq_f64(vld1q_f64(&A_col[8]), valpha); + float64x2_t v5 = vmulq_f64(vld1q_f64(&A_col[10]), valpha); + float64x2_t v6 = vmulq_f64(vld1q_f64(&A_col[12]), valpha); + float64x2_t v7 = vmulq_f64(vld1q_f64(&A_col[14]), valpha); + + vst1q_f64(&A_out_ptr[0], v0); + vst1q_f64(&A_out_ptr[2], v1); + vst1q_f64(&A_out_ptr[4], v2); + vst1q_f64(&A_out_ptr[6], v3); + vst1q_f64(&A_out_ptr[8], v4); + vst1q_f64(&A_out_ptr[10], v5); + vst1q_f64(&A_out_ptr[12], v6); + vst1q_f64(&A_out_ptr[14], v7); + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + int br = 0; + for (; br < current_M - ii; ++br) { + A_out[kk * 16 + br] = A_col[br] * alpha; + } + for (; br < 16; ++br) { + A_out[kk * 16 + br] = 0.0; + } + } + } +#else + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + for (int br = 0; br < 16; ++br) { + A_out[kk * 16 + br] = A_col[br] * alpha; + } + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + int br = 0; + for (; br < current_M - ii; ++br) { + A_out[kk * 16 + br] = A_col[br] * alpha; + } + for (; br < 16; ++br) { + A_out[kk * 16 + br] = 0.0; + } + } + } +#endif +#if PROFILING + CALI_MARK_END("pack_A"); +#endif + + for (int jj = 0; jj < N_pad; jj += 16) { +#if PROFILING + CALI_MARK_BEGIN("loop_jj"); +#endif + int current_N_block = stdmin(16, current_N - jj); + double *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + + for (int ii = 0; ii < M_pad; ii += 16) { +#if PROFILING + CALI_MARK_BEGIN("loop_ii"); +#endif + int current_M_block = stdmin(16, current_M - ii); + double *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + double *C_ptr = &C[(i + ii) + (j + jj) * ldc]; +#if PROFILING + CALI_MARK_BEGIN("kernel"); +#endif + if (current_M_block == 16 && current_N_block == 16) { + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) double C_buffer[256]; + // We strictly only require initialization if the kernel natively relies on reading C. + // If beta == 0 (beta_mode == 0), the kernel entirely skips reading the buffer and writes over it. + if (beta_mode != 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = 0.0; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + } + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } +#if PROFILING + CALI_MARK_END("kernel"); +#endif +#if PROFILING + CALI_MARK_END("loop_ii"); +#endif + } +#if PROFILING + CALI_MARK_END("loop_jj"); +#endif + } +#if PROFILING + CALI_MARK_END("loop_i"); +#endif + } +#if PROFILING + CALI_MARK_END("loop_k"); +#endif + } +#if PROFILING + CALI_MARK_END("loop_j"); +#endif + } +} + +// ======================================================================== +// VERSION 1: C = alpha * A^T * B + beta * C (TN) +// ======================================================================== +void dgemm_sme_TN(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +{ + if (alpha == 0.0 || K == 0) { + if (beta != 1.0) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0) ? 0.0 : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + ptrdiff_t panel_stride_B = 16 * (ptrdiff_t)current_K; + + int beta_mode = 2; + if (k == 0) { + if (beta == 0.0) { + beta_mode = 0; + } + else if (beta != 1.0) { + beta_mode = 1; + } + } + +#if USE_VECTORIZED_PACKING + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_in_ptr = &B[(j + jj) * (ptrdiff_t)ldb + k + kk]; + double *__restrict B_out_ptr = &B_out[kk * 16]; + + ptrdiff_t ldb_sz = (ptrdiff_t)ldb; + + B_out_ptr[0] = B_in_ptr[0]; + B_out_ptr[1] = B_in_ptr[1 * ldb_sz]; + B_out_ptr[2] = B_in_ptr[2 * ldb_sz]; + B_out_ptr[3] = B_in_ptr[3 * ldb_sz]; + B_out_ptr[4] = B_in_ptr[4 * ldb_sz]; + B_out_ptr[5] = B_in_ptr[5 * ldb_sz]; + B_out_ptr[6] = B_in_ptr[6 * ldb_sz]; + B_out_ptr[7] = B_in_ptr[7 * ldb_sz]; + B_out_ptr[8] = B_in_ptr[8 * ldb_sz]; + B_out_ptr[9] = B_in_ptr[9 * ldb_sz]; + B_out_ptr[10] = B_in_ptr[10 * ldb_sz]; + B_out_ptr[11] = B_in_ptr[11 * ldb_sz]; + B_out_ptr[12] = B_in_ptr[12 * ldb_sz]; + B_out_ptr[13] = B_in_ptr[13 * ldb_sz]; + B_out_ptr[14] = B_in_ptr[14 * ldb_sz]; + B_out_ptr[15] = B_in_ptr[15 * ldb_sz]; + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int bc = 0; bc < 16; ++bc) { + if (jj + bc < current_N) { + const double *__restrict B_col = &B[(j + jj + bc) * (ptrdiff_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = B_col[kk]; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = 0.0; + } + } + } + } +#else + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int bc = 0; bc < 16; ++bc) { + const double *__restrict B_col = &B[(j + jj + bc) * (ptrdiff_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = B_col[kk]; + } + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int bc = 0; bc < 16; ++bc) { + if (jj + bc < current_N) { + const double *__restrict B_col = &B[(j + jj + bc) * (ptrdiff_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = B_col[kk]; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 16 + bc] = 0.0; + } + } + } + } +#endif + + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + ptrdiff_t panel_stride_A = 16 * (ptrdiff_t)current_K; + +#if USE_VECTORIZED_PACKING + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_in_ptr = &A[(i + ii) * (ptrdiff_t)lda + k + kk]; + double *__restrict A_out_ptr = &A_out[kk * 16]; + ptrdiff_t lda_sz = (ptrdiff_t)lda; + + A_out_ptr[0] = A_in_ptr[0] * alpha; + A_out_ptr[1] = A_in_ptr[1 * lda_sz] * alpha; + A_out_ptr[2] = A_in_ptr[2 * lda_sz] * alpha; + A_out_ptr[3] = A_in_ptr[3 * lda_sz] * alpha; + A_out_ptr[4] = A_in_ptr[4 * lda_sz] * alpha; + A_out_ptr[5] = A_in_ptr[5 * lda_sz] * alpha; + A_out_ptr[6] = A_in_ptr[6 * lda_sz] * alpha; + A_out_ptr[7] = A_in_ptr[7 * lda_sz] * alpha; + A_out_ptr[8] = A_in_ptr[8 * lda_sz] * alpha; + A_out_ptr[9] = A_in_ptr[9 * lda_sz] * alpha; + A_out_ptr[10] = A_in_ptr[10 * lda_sz] * alpha; + A_out_ptr[11] = A_in_ptr[11 * lda_sz] * alpha; + A_out_ptr[12] = A_in_ptr[12 * lda_sz] * alpha; + A_out_ptr[13] = A_in_ptr[13 * lda_sz] * alpha; + A_out_ptr[14] = A_in_ptr[14 * lda_sz] * alpha; + A_out_ptr[15] = A_in_ptr[15 * lda_sz] * alpha; + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int br = 0; br < 16; ++br) { + if (ii + br < current_M) { + const double *__restrict A_row = &A[(i + ii + br) * (ptrdiff_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = A_row[kk] * alpha; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = 0.0; + } + } + } + } +#else + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int br = 0; br < 16; ++br) { + const double *__restrict A_row = &A[(i + ii + br) * (ptrdiff_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = A_row[kk] * alpha; + } + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int br = 0; br < 16; ++br) { + if (ii + br < current_M) { + const double *__restrict A_row = &A[(i + ii + br) * (ptrdiff_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = A_row[kk] * alpha; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = 0.0; + } + } + } + } +#endif + + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + double *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + double *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + double *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 16 && current_N_block == 16) { + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) double C_buffer[256]; + if (beta_mode != 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = 0.0; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + } + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +// ======================================================================== +// VERSION 2: C = alpha * A * B^T + beta * C (NT) +// ======================================================================== +void dgemm_sme_NT(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +{ + if (alpha == 0.0 || K == 0) { + if (beta != 1.0) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0) ? 0.0 : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + ptrdiff_t panel_stride_B = 16 * (ptrdiff_t)current_K; + + int beta_mode = 2; + if (k == 0) { + if (beta == 0.0) { + beta_mode = 0; + } + else if (beta != 1.0) { + beta_mode = 1; + } + } + +#if USE_VECTORIZED_PACKING + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_col = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + double *__restrict B_out_ptr = &B_out[kk * 16]; + + float64x2_t v0 = vld1q_f64(&B_col[0]); + float64x2_t v1 = vld1q_f64(&B_col[2]); + float64x2_t v2 = vld1q_f64(&B_col[4]); + float64x2_t v3 = vld1q_f64(&B_col[6]); + float64x2_t v4 = vld1q_f64(&B_col[8]); + float64x2_t v5 = vld1q_f64(&B_col[10]); + float64x2_t v6 = vld1q_f64(&B_col[12]); + float64x2_t v7 = vld1q_f64(&B_col[14]); + + vst1q_f64(&B_out_ptr[0], v0); + vst1q_f64(&B_out_ptr[2], v1); + vst1q_f64(&B_out_ptr[4], v2); + vst1q_f64(&B_out_ptr[6], v3); + vst1q_f64(&B_out_ptr[8], v4); + vst1q_f64(&B_out_ptr[10], v5); + vst1q_f64(&B_out_ptr[12], v6); + vst1q_f64(&B_out_ptr[14], v7); + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_row = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + int bc = 0; + for (; bc < current_N - jj; ++bc) { + B_out[kk * 16 + bc] = B_row[bc]; + } + for (; bc < 16; ++bc) { + B_out[kk * 16 + bc] = 0.0; + } + } + } +#else + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_row = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + for (int bc = 0; bc < 16; ++bc) { + B_out[kk * 16 + bc] = B_row[bc]; + } + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_row = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + int bc = 0; + for (; bc < current_N - jj; ++bc) { + B_out[kk * 16 + bc] = B_row[bc]; + } + for (; bc < 16; ++bc) { + B_out[kk * 16 + bc] = 0.0; + } + } + } +#endif + + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + ptrdiff_t panel_stride_A = 16 * (ptrdiff_t)current_K; + +#if USE_VECTORIZED_PACKING + float64x2_t valpha = vdupq_n_f64(alpha); + + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + double *__restrict A_out_ptr = &A_out[kk * 16]; + + float64x2_t v0 = vmulq_f64(vld1q_f64(&A_col[0]), valpha); + float64x2_t v1 = vmulq_f64(vld1q_f64(&A_col[2]), valpha); + float64x2_t v2 = vmulq_f64(vld1q_f64(&A_col[4]), valpha); + float64x2_t v3 = vmulq_f64(vld1q_f64(&A_col[6]), valpha); + float64x2_t v4 = vmulq_f64(vld1q_f64(&A_col[8]), valpha); + float64x2_t v5 = vmulq_f64(vld1q_f64(&A_col[10]), valpha); + float64x2_t v6 = vmulq_f64(vld1q_f64(&A_col[12]), valpha); + float64x2_t v7 = vmulq_f64(vld1q_f64(&A_col[14]), valpha); + + vst1q_f64(&A_out_ptr[0], v0); + vst1q_f64(&A_out_ptr[2], v1); + vst1q_f64(&A_out_ptr[4], v2); + vst1q_f64(&A_out_ptr[6], v3); + vst1q_f64(&A_out_ptr[8], v4); + vst1q_f64(&A_out_ptr[10], v5); + vst1q_f64(&A_out_ptr[12], v6); + vst1q_f64(&A_out_ptr[14], v7); + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + int br = 0; + for (; br < current_M - ii; ++br) { + A_out[kk * 16 + br] = A_col[br] * alpha; + } + for (; br < 16; ++br) { + A_out[kk * 16 + br] = 0.0; + } + } + } +#else + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + for (int br = 0; br < 16; ++br) { + A_out[kk * 16 + br] = A_col[br] * alpha; + } + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_col = &A[(k + kk) * (ptrdiff_t)lda + i + ii]; + int br = 0; + for (; br < current_M - ii; ++br) { + A_out[kk * 16 + br] = A_col[br] * alpha; + } + for (; br < 16; ++br) { + A_out[kk * 16 + br] = 0.0; + } + } + } +#endif + + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + double *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + double *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + double *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 16 && current_N_block == 16) { + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) double C_buffer[256]; + if (beta_mode != 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = 0.0; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + } + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +// ======================================================================== +// VERSION 3: C = alpha * A^T * B^T + beta * C (TT) +// ======================================================================== +void dgemm_sme_TT(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +{ + if (alpha == 0.0 || K == 0) { + if (beta != 1.0) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0) ? 0.0 : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 15) & ~15; + + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + ptrdiff_t panel_stride_B = 16 * (ptrdiff_t)current_K; + + int beta_mode = 2; + if (k == 0) { + if (beta == 0.0) { + beta_mode = 0; + } + else if (beta != 1.0) { + beta_mode = 1; + } + } + +#if USE_VECTORIZED_PACKING + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_col = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + double *__restrict B_out_ptr = &B_out[kk * 16]; + + vst1q_f64(&B_out_ptr[0], vld1q_f64(&B_col[0])); + vst1q_f64(&B_out_ptr[2], vld1q_f64(&B_col[2])); + vst1q_f64(&B_out_ptr[4], vld1q_f64(&B_col[4])); + vst1q_f64(&B_out_ptr[6], vld1q_f64(&B_col[6])); + vst1q_f64(&B_out_ptr[8], vld1q_f64(&B_col[8])); + vst1q_f64(&B_out_ptr[10], vld1q_f64(&B_col[10])); + vst1q_f64(&B_out_ptr[12], vld1q_f64(&B_col[12])); + vst1q_f64(&B_out_ptr[14], vld1q_f64(&B_col[14])); + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_row = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + int bc = 0; + for (; bc < current_N - jj; ++bc) { + B_out[kk * 16 + bc] = B_row[bc]; + } + for (; bc < 16; ++bc) { + B_out[kk * 16 + bc] = 0.0; + } + } + } +#else + int N_main = current_N & ~15; + for (int jj = 0; jj < N_main; jj += 16) { + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_row = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + for (int bc = 0; bc < 16; ++bc) { + B_out[kk * 16 + bc] = B_row[bc]; + } + } + } + if (N_main < current_N) { + int jj = N_main; + double *__restrict B_out = &B_pack[(jj / 16) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict B_row = &B[(k + kk) * (ptrdiff_t)ldb + j + jj]; + int bc = 0; + for (; bc < current_N - jj; ++bc) { + B_out[kk * 16 + bc] = B_row[bc]; + } + for (; bc < 16; ++bc) { + B_out[kk * 16 + bc] = 0.0; + } + } + } +#endif + + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 15) & ~15; + ptrdiff_t panel_stride_A = 16 * (ptrdiff_t)current_K; + +#if USE_VECTORIZED_PACKING + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const double *__restrict A_in_ptr = &A[(i + ii) * (ptrdiff_t)lda + k + kk]; + double *__restrict A_out_ptr = &A_out[kk * 16]; + ptrdiff_t lda_sz = (ptrdiff_t)lda; + + A_out_ptr[0] = A_in_ptr[0] * alpha; + A_out_ptr[1] = A_in_ptr[1 * lda_sz] * alpha; + A_out_ptr[2] = A_in_ptr[2 * lda_sz] * alpha; + A_out_ptr[3] = A_in_ptr[3 * lda_sz] * alpha; + A_out_ptr[4] = A_in_ptr[4 * lda_sz] * alpha; + A_out_ptr[5] = A_in_ptr[5 * lda_sz] * alpha; + A_out_ptr[6] = A_in_ptr[6 * lda_sz] * alpha; + A_out_ptr[7] = A_in_ptr[7 * lda_sz] * alpha; + A_out_ptr[8] = A_in_ptr[8 * lda_sz] * alpha; + A_out_ptr[9] = A_in_ptr[9 * lda_sz] * alpha; + A_out_ptr[10] = A_in_ptr[10 * lda_sz] * alpha; + A_out_ptr[11] = A_in_ptr[11 * lda_sz] * alpha; + A_out_ptr[12] = A_in_ptr[12 * lda_sz] * alpha; + A_out_ptr[13] = A_in_ptr[13 * lda_sz] * alpha; + A_out_ptr[14] = A_in_ptr[14 * lda_sz] * alpha; + A_out_ptr[15] = A_in_ptr[15 * lda_sz] * alpha; + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int br = 0; br < 16; ++br) { + if (ii + br < current_M) { + const double *__restrict A_row = &A[(i + ii + br) * (ptrdiff_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = A_row[kk] * alpha; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = 0.0; + } + } + } + } +#else + int M_main = current_M & ~15; + for (int ii = 0; ii < M_main; ii += 16) { + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int br = 0; br < 16; ++br) { + const double *__restrict A_row = &A[(i + ii + br) * (ptrdiff_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = A_row[kk] * alpha; + } + } + } + if (M_main < current_M) { + int ii = M_main; + double *__restrict A_out = &A_pack[(ii / 16) * panel_stride_A]; + for (int br = 0; br < 16; ++br) { + if (ii + br < current_M) { + const double *__restrict A_row = &A[(i + ii + br) * (ptrdiff_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = A_row[kk] * alpha; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 16 + br] = 0.0; + } + } + } + } +#endif + + for (int jj = 0; jj < N_pad; jj += 16) { + int current_N_block = stdmin(16, current_N - jj); + double *B_ptr = &B_pack[(jj / 16) * panel_stride_B]; + + for (int ii = 0; ii < M_pad; ii += 16) { + int current_M_block = stdmin(16, current_M - ii); + double *A_ptr = &A_pack[(ii / 16) * panel_stride_A]; + double *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 16 && current_N_block == 16) { + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) double C_buffer[256]; + if (beta_mode != 0) { + for (int idx = 0; idx < 256; ++idx) { + C_buffer[idx] = 0.0; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 16] = C_ptr[br + bc * ldc]; + } + } + } + dgemm_sme_compute_16x16_tile(current_K, A_ptr, B_ptr, C_buffer, 16, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 16]; + } + } + } + } + } + } + } + } +} + +// ======================================================================== +// WRAPPER: BLAS ABI Compatible dgemm +// ======================================================================== +void sme_DGEMM_KERNEL(const char *transa, const char *transb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const double *alpha, const double *a, const BLASLONG *lda, const double *b, const BLASLONG *ldb, const double *beta, double *c, const BLASLONG *ldc) +{ + bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); + bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); + if (!trans_a && !trans_b) { + dgemm_sme_NN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } + else if (trans_a && !trans_b) { + dgemm_sme_TN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } + else if (!trans_a && trans_b) { + dgemm_sme_NT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } + else { + dgemm_sme_TT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } +} diff --git a/kernel/arm64/sme_sgemm_kernel.c b/kernel/arm64/sme_sgemm_kernel.c new file mode 100644 index 0000000000..5d5925ad8c --- /dev/null +++ b/kernel/arm64/sme_sgemm_kernel.c @@ -0,0 +1,651 @@ +//#include +#include +//#include +#include +#include +#include +#include "common.h" +#ifndef stdmin +#define stdmin(a,b) (a>b? b:a) +#endif + void SMEStart() + { + asm volatile("smstart\n\t" ::: "memory"); + } + void SMEStop() + { + asm volatile("smstop\n\t" ::: "memory"); + } +// SMEGuard(const SMEGuard &) = delete; +// SMEGuard &operator=(const SMEGuard &) = delete; +/* +const int MC = 512; +const int KC = 1024; +const int NC = 2048; +*/ +#define MC 512 +#define KC 1024 +#define NC 2048 + +static inline void sgemm_sme_compute_32x32_tile(blasint current_K, const float *A_ptr, const float *B_ptr, float *C_ptr, size_t ldc, blasint beta_mode, const float *beta_ptr) +{ + size_t ldc_bytes = ldc * sizeof(float); + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + + //SMEGuard __stream_guard; + asm volatile("smstart\n\t" + "ptrue p0.s\n\t" + "cmp %w[beta_mode], #0\n\t" + "b.eq 10f\n\t" + "cmp %w[beta_mode], #1\n\t" + "b.eq 11f\n\t" + "b 12f\n\t" + + "10:\n\t" + "zero {za}\n\t" + "b 19f\n\t" + + "11:\n\t" + "ld1rw z31.s, p0/z, [%[beta_ptr]]\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "110:\n\t" + "ld1w z0.s, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1w z1.s, p0/z, [x14]\n\t" + "fmul z0.s, p0/m, z0.s, z31.s\n\t" + "fmul z1.s, p0/m, z1.s, z31.s\n\t" + "mova za0v.s[w12, 0], p0/m, z0.s\n\t" + "mova za2v.s[w12, 0], p0/m, z1.s\n\t" // FIXED: za2 (Bottom-Left) + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #16\n\t" + "b.ne 110b\n\t" + + "mov w15, #0\n\t" + "111:\n\t" + "ld1w z0.s, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1w z1.s, p0/z, [x14]\n\t" + "fmul z0.s, p0/m, z0.s, z31.s\n\t" + "fmul z1.s, p0/m, z1.s, z31.s\n\t" + "mova za1v.s[w15, 0], p0/m, z0.s\n\t" // FIXED: za1 (Top-Right) + "mova za3v.s[w15, 0], p0/m, z1.s\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w15, w15, #1\n\t" + "cmp w15, #16\n\t" + "b.ne 111b\n\t" + "b 19f\n\t" + + "12:\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "100:\n\t" + "ld1w {za0v.s[w12, 0]}, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1w {za2v.s[w12, 0]}, p0/z, [x14]\n\t" // FIXED: za2 + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #16\n\t" + "b.ne 100b\n\t" + + "mov w15, #0\n\t" + "101:\n\t" + "ld1w {za1v.s[w15, 0]}, p0/z, [x13]\n\t" // FIXED: za1 + "add x14, x13, #64\n\t" + "ld1w {za3v.s[w15, 0]}, p0/z, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w15, w15, #1\n\t" + "cmp w15, #16\n\t" + "b.ne 101b\n\t" + + "19:\n\t" + "mov w10, %w[k]\n\t" + "cbz w10, 3f\n\t" + "11:\n\t" + "ld1w z0.s, p0/z, [%[a], #0, mul vl]\n\t" + "ld1w z1.s, p0/z, [%[a], #1, mul vl]\n\t" + "ld1w z2.s, p0/z, [%[b], #0, mul vl]\n\t" + "ld1w z3.s, p0/z, [%[b], #1, mul vl]\n\t" + "fmopa za0.s, p0/m, p0/m, z0.s, z2.s\n\t" + "fmopa za1.s, p0/m, p0/m, z0.s, z3.s\n\t" + "fmopa za2.s, p0/m, p0/m, z1.s, z2.s\n\t" + "fmopa za3.s, p0/m, p0/m, z1.s, z3.s\n\t" + "add %[a], %[a], #128\n\t" + "add %[b], %[b], #128\n\t" + "subs w10, w10, #1\n\t" + "b.ne 11b\n\t" + + "3:\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "200:\n\t" + "st1w {za0v.s[w12, 0]}, p0, [x13]\n\t" + "add x14, x13, #64\n\t" + "st1w {za2v.s[w12, 0]}, p0, [x14]\n\t" // FIXED: za2 + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #16\n\t" + "b.ne 200b\n\t" + + "mov w15, #0\n\t" + "201:\n\t" + "st1w {za1v.s[w15, 0]}, p0, [x13]\n\t" // FIXED: za1 + "add x14, x13, #64\n\t" + "st1w {za3v.s[w15, 0]}, p0, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w15, w15, #1\n\t" + "cmp w15, #16\n\t" + "b.ne 201b\n\t" + "smstop\n\t" + : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) + : "p0", "x10", "w12", "x13", "x14", "w15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z31", "za"); + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + +} + +void sgemm_sme_NN(blasint M, blasint N, blasint K, float alpha, const float *A, blasint lda, const float *B, blasint ldb, float beta, float *C, blasint ldc) +{ + if (alpha == 0.0f || K == 0) { + if (beta != 1.0f) { + for (blasint j = 0; j < N; ++j) { + for (blasint i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0f) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static float A_pack[MC * KC]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (blasint j = 0; j < N; j += NC) { + blasint current_N = stdmin(NC, N - j); + blasint N_pad = (current_N + 31) & ~31; + for (blasint k = 0; k < K; k += KC) { + blasint current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == 0.0f) ? 0 : (beta != 1.0f ? 1 : 2)) : 2; + + blasint N_main = current_N & ~31; + for (blasint jj = 0; jj < N_main; jj += 32) { + float *__restrict B_out = &B_pack[(jj / 32) * panel_stride_B]; + for (blasint kk = 0; kk < current_K; ++kk) { + const float *__restrict B_in = &B[(j + jj) * (size_t)ldb + k + kk]; + for (int bc = 0; bc < 32; ++bc) { + B_out[kk * 32 + bc] = B_in[bc * ldb]; + } + } + } + if (N_main < current_N) { + float *__restrict B_out = &B_pack[(N_main / 32) * panel_stride_B]; + for (int bc = 0; bc < 32; ++bc) { + if (N_main + bc < current_N) { + const float *__restrict B_col = &B[(j + N_main + bc) * (size_t)ldb + k]; + for (blasint kk = 0; kk < current_K; ++kk) { + B_out[kk * 32 + bc] = B_col[kk]; + } + } + else { + for (blasint kk = 0; kk < current_K; ++kk) { + B_out[kk * 32 + bc] = 0.0f; + } + } + } + } + + for (blasint i = 0; i < M; i += MC) { + blasint current_M = stdmin(MC, M - i); + blasint M_pad = (current_M + 31) & ~31; + size_t panel_stride_A = 32 * (size_t)current_K; + float32x4_t valpha = vdupq_n_f32(alpha); + + blasint M_main = current_M & ~31; + for (blasint ii = 0; ii < M_main; ii += 32) { + float *__restrict A_out = &A_pack[(ii / 32) * panel_stride_A]; + for (blasint kk = 0; kk < current_K; ++kk) { + const float *__restrict A_col = &A[(k + kk) * (size_t)lda + i + ii]; + float *__restrict A_out_ptr = &A_out[kk * 32]; + for (int v = 0; v < 8; ++v) { + vst1q_f32(&A_out_ptr[v * 4], vmulq_f32(vld1q_f32(&A_col[v * 4]), valpha)); + } + } + } + if (M_main < current_M) { + float *__restrict A_out = &A_pack[(M_main / 32) * panel_stride_A]; + for (blasint kk = 0; kk < current_K; ++kk) { + const float *__restrict A_col = &A[(k + kk) * (size_t)lda + i + M_main]; + for (blasint br = 0; br < current_M - M_main; ++br) { + A_out[kk * 32 + br] = A_col[br] * alpha; + } + for (int br = current_M - M_main; br < 32; ++br) { + A_out[kk * 32 + br] = 0.0f; + } + } + } + + for (blasint jj = 0; jj < N_pad; jj += 32) { + blasint current_N_block = stdmin(32, current_N - jj); + float *B_ptr = &B_pack[(jj / 32) * panel_stride_B]; + for (blasint ii = 0; ii < M_pad; ii += 32) { + int current_M_block = stdmin(32, current_M - ii); + float *A_ptr = &A_pack[(ii / 32) * panel_stride_A]; + float *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 32 && current_N_block == 32) { + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) float C_buffer[1024]; + if (beta_mode != 0) { + for (int idx = 0; idx < 1024; ++idx) { + C_buffer[idx] = 0.0f; + } + for (blasint bc = 0; bc < current_N_block; ++bc) { + for (blasint br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 32] = C_ptr[br + bc * ldc]; + } + } + } + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_buffer, 32, beta_mode, &beta); + for (blasint bc = 0; bc < current_N_block; ++bc) { + for (blasint br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 32]; + } + } + } + } + } + } + } + } +} + +void sgemm_sme_TN(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) +{ + if (alpha == 0.0f || K == 0) { + if (beta != 1.0f) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0f) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static float A_pack[MC * KC]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 31) & ~31; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == 0.0f) ? 0 : (beta != 1.0f ? 1 : 2)) : 2; + + int N_main = current_N & ~31; + for (int jj = 0; jj < N_main; jj += 32) { + float *__restrict B_out = &B_pack[(jj / 32) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict B_in = &B[(j + jj) * (size_t)ldb + k + kk]; + for (int bc = 0; bc < 32; ++bc) { + B_out[kk * 32 + bc] = B_in[bc * ldb]; + } + } + } + if (N_main < current_N) { + float *__restrict B_out = &B_pack[(N_main / 32) * panel_stride_B]; + for (int bc = 0; bc < 32; ++bc) { + if (N_main + bc < current_N) { + const float *__restrict B_col = &B[(j + N_main + bc) * (size_t)ldb + k]; + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 32 + bc] = B_col[kk]; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + B_out[kk * 32 + bc] = 0.0f; + } + } + } + } + + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 31) & ~31; + size_t panel_stride_A = 32 * (size_t)current_K; + + int M_main = current_M & ~31; + for (int ii = 0; ii < M_main; ii += 32) { + float *__restrict A_out = &A_pack[(ii / 32) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict A_in = &A[(i + ii) * (size_t)lda + k + kk]; + for (int br = 0; br < 32; ++br) { + A_out[kk * 32 + br] = A_in[br * lda] * alpha; + } + } + } + if (M_main < current_M) { + float *__restrict A_out = &A_pack[(M_main / 32) * panel_stride_A]; + for (int br = 0; br < 32; ++br) { + if (M_main + br < current_M) { + const float *__restrict A_row = &A[(i + M_main + br) * (size_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 32 + br] = A_row[kk] * alpha; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 32 + br] = 0.0f; + } + } + } + } + + for (int jj = 0; jj < N_pad; jj += 32) { + int current_N_block = stdmin(32, current_N - jj); + float *B_ptr = &B_pack[(jj / 32) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 32) { + int current_M_block = stdmin(32, current_M - ii); + float *A_ptr = &A_pack[(ii / 32) * panel_stride_A]; + float *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 32 && current_N_block == 32) { + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) float C_buffer[1024]; + if (beta_mode != 0) { + for (int idx = 0; idx < 1024; ++idx) { + C_buffer[idx] = 0.0f; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 32] = C_ptr[br + bc * ldc]; + } + } + } + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_buffer, 32, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 32]; + } + } + } + } + } + } + } + } +} + +void sgemm_sme_NT(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) +{ + if (alpha == 0.0f || K == 0) { + if (beta != 1.0f) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0f) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static float A_pack[MC * KC]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 31) & ~31; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == 0.0f) ? 0 : (beta != 1.0f ? 1 : 2)) : 2; + + int N_main = current_N & ~31; + for (int jj = 0; jj < N_main; jj += 32) { + float *__restrict B_out = &B_pack[(jj / 32) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict B_col = &B[(k + kk) * (size_t)ldb + j + jj]; + float *__restrict B_out_ptr = &B_out[kk * 32]; + for (int v = 0; v < 8; ++v) { + vst1q_f32(&B_out_ptr[v * 4], vld1q_f32(&B_col[v * 4])); + } + } + } + if (N_main < current_N) { + float *__restrict B_out = &B_pack[(N_main / 32) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict B_row = &B[(k + kk) * (size_t)ldb + j + N_main]; + for (int bc = 0; bc < current_N - N_main; ++bc) { + B_out[kk * 32 + bc] = B_row[bc]; + } + for (int bc = current_N - N_main; bc < 32; ++bc) { + B_out[kk * 32 + bc] = 0.0f; + } + } + } + + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 31) & ~31; + size_t panel_stride_A = 32 * (size_t)current_K; + float32x4_t valpha = vdupq_n_f32(alpha); + + int M_main = current_M & ~31; + for (int ii = 0; ii < M_main; ii += 32) { + float *__restrict A_out = &A_pack[(ii / 32) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict A_col = &A[(k + kk) * (size_t)lda + i + ii]; + float *__restrict A_out_ptr = &A_out[kk * 32]; + for (int v = 0; v < 8; ++v) { + vst1q_f32(&A_out_ptr[v * 4], vmulq_f32(vld1q_f32(&A_col[v * 4]), valpha)); + } + } + } + if (M_main < current_M) { + float *__restrict A_out = &A_pack[(M_main / 32) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict A_col = &A[(k + kk) * (size_t)lda + i + M_main]; + for (int br = 0; br < current_M - M_main; ++br) { + A_out[kk * 32 + br] = A_col[br] * alpha; + } + for (int br = current_M - M_main; br < 32; ++br) { + A_out[kk * 32 + br] = 0.0f; + } + } + } + + for (int jj = 0; jj < N_pad; jj += 32) { + int current_N_block = stdmin(32, current_N - jj); + float *B_ptr = &B_pack[(jj / 32) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 32) { + int current_M_block = stdmin(32, current_M - ii); + float *A_ptr = &A_pack[(ii / 32) * panel_stride_A]; + float *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 32 && current_N_block == 32) { + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) float C_buffer[1024]; + if (beta_mode != 0) { + for (int idx = 0; idx < 1024; ++idx) { + C_buffer[idx] = 0.0f; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 32] = C_ptr[br + bc * ldc]; + } + } + } + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_buffer, 32, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 32]; + } + } + } + } + } + } + } + } +} + +void sgemm_sme_TT(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) +{ + if (alpha == 0.0f || K == 0) { + if (beta != 1.0f) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == 0.0f) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static float A_pack[MC * KC]; + alignas(256) /*thread_local*/ static float B_pack[KC * NC]; + +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 31) & ~31; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 32 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == 0.0f) ? 0 : (beta != 1.0f ? 1 : 2)) : 2; + + int N_main = current_N & ~31; + for (int jj = 0; jj < N_main; jj += 32) { + float *__restrict B_out = &B_pack[(jj / 32) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict B_col = &B[(k + kk) * (size_t)ldb + j + jj]; + float *__restrict B_out_ptr = &B_out[kk * 32]; + for (int v = 0; v < 8; ++v) { + vst1q_f32(&B_out_ptr[v * 4], vld1q_f32(&B_col[v * 4])); + } + } + } + if (N_main < current_N) { + float *__restrict B_out = &B_pack[(N_main / 32) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict B_row = &B[(k + kk) * (size_t)ldb + j + N_main]; + for (int bc = 0; bc < current_N - N_main; ++bc) { + B_out[kk * 32 + bc] = B_row[bc]; + } + for (int bc = current_N - N_main; bc < 32; ++bc) { + B_out[kk * 32 + bc] = 0.0f; + } + } + } + + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 31) & ~31; + size_t panel_stride_A = 32 * (size_t)current_K; + + int M_main = current_M & ~31; + for (int ii = 0; ii < M_main; ii += 32) { + float *__restrict A_out = &A_pack[(ii / 32) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + const float *__restrict A_in = &A[(i + ii) * (size_t)lda + k + kk]; + for (int br = 0; br < 32; ++br) { + A_out[kk * 32 + br] = A_in[br * lda] * alpha; + } + } + } + if (M_main < current_M) { + float *__restrict A_out = &A_pack[(M_main / 32) * panel_stride_A]; + for (int br = 0; br < 32; ++br) { + if (M_main + br < current_M) { + const float *__restrict A_row = &A[(i + M_main + br) * (size_t)lda + k]; + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 32 + br] = A_row[kk] * alpha; + } + } + else { + for (int kk = 0; kk < current_K; ++kk) { + A_out[kk * 32 + br] = 0.0f; + } + } + } + } + + for (int jj = 0; jj < N_pad; jj += 32) { + int current_N_block = stdmin(32, current_N - jj); + float *B_ptr = &B_pack[(jj / 32) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 32) { + int current_M_block = stdmin(32, current_M - ii); + float *A_ptr = &A_pack[(ii / 32) * panel_stride_A]; + float *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + + if (current_M_block == 32 && current_N_block == 32) { + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) float C_buffer[1024]; + if (beta_mode != 0) { + for (int idx = 0; idx < 1024; ++idx) { + C_buffer[idx] = 0.0f; + } + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 32] = C_ptr[br + bc * ldc]; + } + } + } + sgemm_sme_compute_32x32_tile(current_K, A_ptr, B_ptr, C_buffer, 32, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 32]; + } + } + } + } + } + } + } + } +} + +/*extern "C"*/ void sme_SGEMM_KERNEL(const char *transa, const char *transb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const float *alpha, const float *a, const BLASLONG *lda, const float *b, const BLASLONG *ldb, const float *beta, float *c, const BLASLONG *ldc) +{ + bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); + bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); + if (!trans_a && !trans_b) { + sgemm_sme_NN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } + else if (trans_a && !trans_b) { + sgemm_sme_TN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } + else if (!trans_a && trans_b) { + sgemm_sme_NT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } + else { + sgemm_sme_TT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + } +} diff --git a/kernel/arm64/sme_zgemm_kernel.c b/kernel/arm64/sme_zgemm_kernel.c new file mode 100644 index 0000000000..71744e64c3 --- /dev/null +++ b/kernel/arm64/sme_zgemm_kernel.c @@ -0,0 +1,547 @@ +#include +#include +#include +#include +#include "common.h" +#ifndef stdmin +#define stdmin(a,b) (a>b? b:a) +#endif + +typedef double _Complex zdouble; + +zdouble CDMUL(zdouble a, zdouble b,bool conja, bool conjb) { +double ra=creal(a); +double rb=creal(b); +double ia=conja ? -cimag(a) : cimag(a); +double ib=conjb ? -cimag(b) : cimag(b); +double r1=ra*rb; +double r2=ia*ib; +double r=r1-r2; +double i=(ra+ia)*(rb+ib)-r1-r2; +zdouble res={r,i}; +return res; +} + + +zdouble cdzero={0.,0.}; +zdouble cdone={1.,0.}; + + + + +#define MC 128 +#define KC 256 +#define NC 512 + +inline void zgemm_sme_compute_8x8_tile(int current_K, const double *A_ptr, const double *B_ptr, zdouble *C_ptr, size_t ldc, int beta_mode, const zdouble *beta_ptr) +{ + size_t ldc_bytes = ldc * sizeof(zdouble); + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + + + asm volatile("smstart\n\t" + "ptrue p0.d\n\t" + "zero {za}\n\t" + "cmp %w[beta_mode], #0\n\t" + "b.eq 19f\n\t" + + // FIX: Load Beta Real and Imaginary components and replicate them + "cmp %w[beta_mode], #1\n\t" + "b.ne 10f\n\t" + "ld1rd z30.d, p0/z, [%[beta_ptr]]\n\t" // Load beta.real() + "add x15, %[beta_ptr], #8\n\t" + "ld1rd z31.d, p0/z, [x15]\n\t" // Load beta.imag() + "10:\n\t" + + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "100:\n\t" + "ld1d z0.d, p0/z, [x13]\n\t" + "add x14, x13, #64\n\t" + "ld1d z1.d, p0/z, [x14]\n\t" + "uzp1 z2.d, z0.d, z1.d\n\t" // z2 = C_re + "uzp2 z3.d, z0.d, z1.d\n\t" // z3 = C_im + + "cmp %w[beta_mode], #1\n\t" + "b.ne 101f\n\t" + + // FIX: Fully implemented Complex Beta Multiplication using movprfx + "movprfx z4, z2\n\t" + "fmul z4.d, p0/m, z4.d, z30.d\n\t" // z4 = Cre * Bre + "movprfx z5, z3\n\t" + "fmul z5.d, p0/m, z5.d, z31.d\n\t" // z5 = Cim * Bim + "movprfx z6, z4\n\t" + "fsub z6.d, p0/m, z6.d, z5.d\n\t" // z6 = Cre' (Cre*Bre - Cim*Bim) + + "movprfx z8, z2\n\t" + "fmul z8.d, p0/m, z8.d, z31.d\n\t" // z8 = Cre * Bim + "movprfx z9, z3\n\t" + "fmul z9.d, p0/m, z9.d, z30.d\n\t" // z9 = Cim * Bre + "movprfx z7, z8\n\t" + "fadd z7.d, p0/m, z7.d, z9.d\n\t" // z7 = Cim' (Cre*Bim + Cim*Bre) + + "mova za0v.d[w12, 0], p0/m, z6.d\n\t" + "mova za2v.d[w12, 0], p0/m, z7.d\n\t" + "b 102f\n\t" + + "101:\n\t" // Fallback: beta == 1.0 + "mova za0v.d[w12, 0], p0/m, z2.d\n\t" + "mova za2v.d[w12, 0], p0/m, z3.d\n\t" + + "102:\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #8\n\t" + "b.ne 100b\n\t" + + "19:\n\t" + "mov w10, %w[k]\n\t" + "cbz w10, 3f\n\t" + "11:\n\t" + "ld1d z0.d, p0/z, [%[a], #0, mul vl]\n\t" + "ld1d z1.d, p0/z, [%[a], #1, mul vl]\n\t" + "ld1d z2.d, p0/z, [%[b], #0, mul vl]\n\t" + "ld1d z3.d, p0/z, [%[b], #1, mul vl]\n\t" + "fmopa za0.d, p0/m, p0/m, z0.d, z2.d\n\t" + "fmopa za1.d, p0/m, p0/m, z1.d, z3.d\n\t" + "fmopa za2.d, p0/m, p0/m, z0.d, z3.d\n\t" + "fmopa za3.d, p0/m, p0/m, z1.d, z2.d\n\t" + "add %[a], %[a], #128\n\t" + "add %[b], %[b], #128\n\t" + "subs w10, w10, #1\n\t" + "b.ne 11b\n\t" + + "3:\n\t" + "mov w12, #0\n\t" + "mov x13, %[c]\n\t" + "200:\n\t" + "mova z0.d, p0/m, za0v.d[w12, 0]\n\t" + "mova z1.d, p0/m, za1v.d[w12, 0]\n\t" + "mova z2.d, p0/m, za2v.d[w12, 0]\n\t" + "mova z3.d, p0/m, za3v.d[w12, 0]\n\t" + + // FIX: Non-destructive SVE arithmetic + "movprfx z4, z0\n\t" + "fsub z4.d, p0/m, z4.d, z1.d\n\t" + "movprfx z5, z2\n\t" + "fadd z5.d, p0/m, z5.d, z3.d\n\t" + + "zip1 z6.d, z4.d, z5.d\n\t" + "zip2 z7.d, z4.d, z5.d\n\t" + "st1d z6.d, p0, [x13]\n\t" + "add x14, x13, #64\n\t" + "st1d z7.d, p0, [x14]\n\t" + "add x13, x13, %[ldc_bytes]\n\t" + "add w12, w12, #1\n\t" + "cmp w12, #8\n\t" + "b.ne 200b\n\t" + "smstop\n\t" + : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) + // Updated Clobber list to cover all the new registers utilized + : "p0", "x10", "w12", "x13", "x14", "x15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z30", "z31", "za"); + + + asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + + +} + +void zgemm_sme_NN(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) +{ + + if (alpha == cdzero || K == 0) { + if (beta != cdone) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == cdzero) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 7) & ~7; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 16 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == cdzero) ? 0 : (beta != cdone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 8) { + double *__restrict B_out = &B_pack[(jj / 8) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + for (int bc = 0; bc < 8; ++bc) { + zdouble val = (jj + bc < current_N) ? b[(j + jj + bc) * ldb + k + kk] : cdzero; + B_out[kk * 16 + bc] = creal(val); + B_out[kk * 16 + 8 + bc] = conjb ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 7) & ~7; + size_t panel_stride_A = 16 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 8) { + double *__restrict A_out = &A_pack[(ii / 8) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + for (int br = 0; br < 8; ++br) { + zdouble val = (ii + br < current_M) ? CDMUL(A[(k + kk) * lda + i + ii + br] ,alpha,conja,0) : cdzero; + A_out[kk * 16 + br] = creal(val); + A_out[kk * 16 + 8 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 8) { + int current_N_block = stdmin(8, current_N - jj); + double *B_ptr = &B_pack[(jj / 8) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 8) { + int current_M_block = stdmin(8, current_M - ii); + double *A_ptr = &A_pack[(ii / 8) * panel_stride_A]; + zdouble *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 8 && current_N_block == 8) { + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) zdouble C_buffer[64]; + + for (int idx = 0; idx < 64; ++idx) { + C_buffer[idx] = cdzero; + } + if (beta_mode != 0) { + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 8] = C_ptr[br + bc * ldc]; + } + } + } + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_buffer, 8, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 8]; + } + } + } + } + } + } + } + } +} + +void zgemm_sme_TN(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) +{ + if (alpha == cdzero || K == 0) { + if (beta != cdone) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == cdzero) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 7) & ~7; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 16 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == cdzero) ? 0 : (beta != cdone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 8) { + double *__restrict B_out = &B_pack[(jj / 8) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + for (int bc = 0; bc < 8; ++bc) { + zdouble val = (jj + bc < current_N) ? b[(j + jj + bc) * ldb + k + kk] : cdzero; + B_out[kk * 16 + bc] = creal(val); + B_out[kk * 16 + 8 + bc] = conjb ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 7) & ~7; + size_t panel_stride_A = 16 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 8) { + double *__restrict A_out = &A_pack[(ii / 8) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + for (int br = 0; br < 8; ++br) { + zdouble val = (ii + br < current_M) ? CDMUL(A[(i + ii + br) * lda + k + kk],alpha,conja,0) : cdzero; + A_out[kk * 16 + br] = creal(val); + A_out[kk * 16 + 8 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 8) { + int current_N_block = stdmin(8, current_N - jj); + double *B_ptr = &B_pack[(jj / 8) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 8) { + int current_M_block = stdmin(8, current_M - ii); + double *A_ptr = &A_pack[(ii / 8) * panel_stride_A]; + zdouble *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 8 && current_N_block == 8) { + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) zdouble C_buffer[64]; + + for (int idx = 0; idx < 64; ++idx) { + C_buffer[idx] = cdzero; + } + if (beta_mode != 0) { + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 8] = C_ptr[br + bc * ldc]; + } + } + } + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_buffer, 8, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 8]; + } + } + } + } + } + } + } + } +} + +void zgemm_sme_NT(blasint M, blasint N, blasint K, const zdouble alpha, const zdouble *A, blasint lda, const zdouble *b, blasint ldb, const zdouble beta, zdouble *C, blasint ldc, bool conja, bool conjb) +{ +#if 0 + if (alpha == cdzero || K == 0) { + if (beta != cdone) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == cdzero) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } +#endif + if (alpha == cdzero || K == 0) { + if (beta == cdzero) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = cdzero; + } + } + } else { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] *= beta; + } + } + } + return; + } + + + alignas(256) /*thread_local*/ static double A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC * 2 *2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 7) & ~7; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 16 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == cdzero) ? 0 : (beta != cdone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 8) { + double *__restrict B_out = &B_pack[(jj / 8) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + for (int bc = 0; bc < 8; ++bc) { + zdouble val = (jj + bc < current_N) ? b[(k + kk) * ldb + j + jj + bc] : cdzero; + B_out[kk * 16 + bc] = creal(val); + B_out[kk * 16 + 8 + bc] = conjb ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 7) & ~7; + size_t panel_stride_A = 16 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 8) { + double *__restrict A_out = &A_pack[(ii / 8) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + for (int br = 0; br < 8; ++br) { + zdouble val = (ii + br < current_M) ? CDMUL(A[(k + kk) * lda + i + ii + br] ,alpha,conja,0) : cdzero; + A_out[kk * 16 + br] = creal(val); + A_out[kk * 16 + 8 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 8) { + int current_N_block = stdmin(8, current_N - jj); + double *B_ptr = &B_pack[(jj / 8) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 8) { + int current_M_block = stdmin(8, current_M - ii); + double *A_ptr = &A_pack[(ii / 8) * panel_stride_A]; + zdouble *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 8 && current_N_block == 8) { + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) zdouble C_buffer[64]; + + for (int idx = 0; idx < 64; ++idx) { + C_buffer[idx] = cdzero; + } + if (beta_mode != 0) { + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 8] = C_ptr[br + bc * ldc]; + } + } + } + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_buffer, 8, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 8]; + } + } + } + } + } + } + } + } +} + +void zgemm_sme_TT(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) +{ + if (alpha == cdzero || K == 0) { + if (beta != cdone) { + for (int j = 0; j < N; ++j) { + for (int i = 0; i < M; ++i) { + C[i + j * ldc] = (beta == cdzero) ? 0.0f : C[i + j * ldc] * beta; + } + } + } + return; + } + + alignas(256) /*thread_local*/ static double A_pack[MC * KC * 2]; + alignas(256) /*thread_local*/ static double B_pack[KC * NC * 2]; +#pragma omp parallel for schedule(dynamic, 1) + for (int j = 0; j < N; j += NC) { + int current_N = stdmin(NC, N - j); + int N_pad = (current_N + 7) & ~7; + for (int k = 0; k < K; k += KC) { + int current_K = stdmin(KC, K - k); + size_t panel_stride_B = 16 * (size_t)current_K; + int beta_mode = (k == 0) ? ((beta == cdzero) ? 0 : (beta != cdone ? 1 : 2)) : 2; + + for (int jj = 0; jj < current_N; jj += 8) { + double *__restrict B_out = &B_pack[(jj / 8) * panel_stride_B]; + for (int kk = 0; kk < current_K; ++kk) { + for (int bc = 0; bc < 8; ++bc) { + zdouble val = (jj + bc < current_N) ? b[(k + kk) * ldb + j + jj + bc] : cdzero; + B_out[kk * 16 + bc] = creal(val); + B_out[kk * 16 + 8 + bc] = conjb ? -cimag(val) : cimag(val); + } + } + } + for (int i = 0; i < M; i += MC) { + int current_M = stdmin(MC, M - i); + int M_pad = (current_M + 7) & ~7; + size_t panel_stride_A = 16 * (size_t)current_K; + for (int ii = 0; ii < current_M; ii += 8) { + double *__restrict A_out = &A_pack[(ii / 8) * panel_stride_A]; + for (int kk = 0; kk < current_K; ++kk) { + for (int br = 0; br < 8; ++br) { + zdouble val = (ii + br < current_M) ? CDMUL(A[(i + ii + br) * lda + k + kk],alpha,conja,0) : cdzero; + A_out[kk * 16 + br] = creal(val); + A_out[kk * 16 + 8 + br] = cimag(val); + } + } + } + for (int jj = 0; jj < N_pad; jj += 8) { + int current_N_block = stdmin(8, current_N - jj); + double *B_ptr = &B_pack[(jj / 8) * panel_stride_B]; + for (int ii = 0; ii < M_pad; ii += 8) { + int current_M_block = stdmin(8, current_M - ii); + double *A_ptr = &A_pack[(ii / 8) * panel_stride_A]; + zdouble *C_ptr = &C[(i + ii) + (j + jj) * ldc]; + if (current_M_block == 8 && current_N_block == 8) { + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_ptr, ldc, beta_mode, &beta); + } + else { + alignas(256) zdouble C_buffer[64]; + + for (int idx = 0; idx < 64; ++idx) { + C_buffer[idx] = cdzero; + } + if (beta_mode != 0) { + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_buffer[br + bc * 8] = C_ptr[br + bc * ldc]; + } + } + } + zgemm_sme_compute_8x8_tile(current_K, A_ptr, B_ptr, C_buffer, 8, beta_mode, &beta); + for (int bc = 0; bc < current_N_block; ++bc) { + for (int br = 0; br < current_M_block; ++br) { + C_ptr[br + bc * ldc] = C_buffer[br + bc * 8]; + } + } + } + } + } + } + } + } +} + +void sme_ZGEMM_KERNEL(const char *transa, const char *transb, const blasint *m, const blasint *n, const blasint *k, const zdouble alpha, const zdouble *a, const blasint *lda, const zdouble *b, const blasint *ldb, const zdouble beta, zdouble *c, const blasint *ldc) +{ + bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); + bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); + if (!trans_a && !trans_b) { + bool conja = (*transa == 'R' || *transa == 'r'); + bool conjb = (*transb == 'R' || *transb == 'r'); + zgemm_sme_NN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } + else if (trans_a && !trans_b) { + bool conja = (*transa == 'C' || *transa == 'c'); + bool conjb = (*transb == 'R' || *transb == 'r'); + zgemm_sme_TN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } + else if (!trans_a && trans_b) { + bool conja = (*transa == 'R' || *transa == 'r'); + bool conjb = (*transb == 'C' || *transb == 'c'); + zgemm_sme_NT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } + else { + bool conja = (*transa == 'C' || *transa == 'c'); + bool conjb = (*transb == 'C' || *transb == 'c'); + zgemm_sme_TT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + } +} + From cf1b7c1ad90e9ef74dc8d26706ee659804acb95d Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 12:14:58 +0200 Subject: [PATCH 05/16] Clean up non-OpenMP build and add CBLAS GEMM benchmark --- benchmark/Makefile | 36 ++++++- benchmark/cblasgemm.c | 233 ++++++++++++++++++++++++++++++++++++++++++ benchmark/gemm.c | 4 +- 3 files changed, 271 insertions(+), 2 deletions(-) create mode 100644 benchmark/cblasgemm.c diff --git a/benchmark/Makefile b/benchmark/Makefile index ec7930b1cc..f773b09d16 100644 --- a/benchmark/Makefile +++ b/benchmark/Makefile @@ -95,10 +95,15 @@ else GOTO_HFLOAT_TARGETS= endif +ifeq ($(USE_OPENMP), 1) +SMALLSCALING=smallscaling +endif + ifeq ($(OSNAME), WINNT) goto :: slinpack.goto dlinpack.goto clinpack.goto zlinpack.goto \ scholesky.goto dcholesky.goto ccholesky.goto zcholesky.goto \ + cblas_sgemm.goto cblas_dgemm.goto cblas_cgemm.goto cblas_zgemm.goto \ sgemm.goto dgemm.goto cgemm.goto zgemm.goto \ strmm.goto dtrmm.goto ctrmm.goto ztrmm.goto \ strsm.goto dtrsm.goto ctrsm.goto ztrsm.goto \ @@ -268,6 +273,7 @@ mkl :: slinpack.mkl dlinpack.mkl clinpack.mkl zlinpack.mkl \ else goto :: sgemm.goto dgemm.goto cgemm.goto zgemm.goto \ + cblas_sgemm.goto cblas_dgemm.goto cblas_cgemm.goto cblas_zgemm.goto \ strmm.goto dtrmm.goto ctrmm.goto ztrmm.goto \ strsm.goto dtrsm.goto ctrsm.goto ztrsm.goto \ sspr.goto dspr.goto \ @@ -301,7 +307,7 @@ goto :: sgemm.goto dgemm.goto cgemm.goto zgemm.goto \ stpsv.goto dtpsv.goto ctpsv.goto ztpsv.goto \ strsv.goto dtrsv.goto ctrsv.goto ztrsv.goto \ ssymm.goto dsymm.goto csymm.goto zsymm.goto \ - smallscaling \ + $(SMALLSCALING) \ isamax.goto idamax.goto icamax.goto izamax.goto \ ismax.goto idmax.goto \ isamin.goto idamin.goto icamin.goto izamin.goto \ @@ -681,6 +687,18 @@ endif sgemm.goto : sgemm.$(SUFFIX) ../$(LIBNAME) $(CC) $(CFLAGS) -o $(@F) $^ $(CEXTRALIB) $(EXTRALIB) $(FEXTRALIB) -lm +cblas_sgemm.goto : cblas_sgemm.$(SUFFIX) ../$(LIBNAME) + $(CC) $(CFLAGS) -o $(@F) $^ $(CEXTRALIB) $(EXTRALIB) $(FEXTRALIB) -lm + +cblas_dgemm.goto : cblas_dgemm.$(SUFFIX) ../$(LIBNAME) + $(CC) $(CFLAGS) -o $(@F) $^ $(CEXTRALIB) $(EXTRALIB) $(FEXTRALIB) -lm + +cblas_cgemm.goto : cblas_cgemm.$(SUFFIX) ../$(LIBNAME) + $(CC) $(CFLAGS) -o $(@F) $^ $(CEXTRALIB) $(EXTRALIB) $(FEXTRALIB) -lm + +cblas_zgemm.goto : cblas_zgemm.$(SUFFIX) ../$(LIBNAME) + $(CC) $(CFLAGS) -o $(@F) $^ $(CEXTRALIB) $(EXTRALIB) $(FEXTRALIB) -lm + sgemm.acml : sgemm.$(SUFFIX) -$(CC) $(CFLAGS) -o $(@F) $^ $(LIBACML) $(CEXTRALIB) $(EXTRALIB) $(FEXTRALIB) @@ -3027,6 +3045,18 @@ cgemm.$(SUFFIX) : gemm.c zgemm.$(SUFFIX) : gemm.c $(CC) $(CFLAGS) -c -DCOMPLEX -DDOUBLE -o $(@F) $^ +cblas_sgemm.$(SUFFIX) : cblasgemm.c + $(CC) $(CFLAGS) -c -UCOMPLEX -UDOUBLE -o $(@F) $^ + +cblas_dgemm.$(SUFFIX) : cblasgemm.c + $(CC) $(CFLAGS) -c -UCOMPLEX -DDOUBLE -o $(@F) $^ + +cblas_cgemm.$(SUFFIX) : cblasgemm.c + $(CC) $(CFLAGS) -c -DCOMPLEX -UDOUBLE -o $(@F) $^ + +cblas_zgemm.$(SUFFIX) : cblasgemm.c + $(CC) $(CFLAGS) -c -DCOMPLEX -DDOUBLE -o $(@F) $^ + ssymm.$(SUFFIX) : symm.c $(CC) $(CFLAGS) -c -UCOMPLEX -UDOUBLE -o $(@F) $^ @@ -3533,7 +3563,11 @@ zomatcopy.$(SUFFIX) : omatcopy.c smallscaling: smallscaling.c ../$(LIBNAME) +ifeq ($(C_COMPILER), GCC) $(CC) $(CFLAGS) -o $(@F) $^ $(EXTRALIB) -fopenmp -lm -lpthread +else + $(CC) $(CFLAGS) -o $(@F) $^ $(EXTRALIB) -quak -openmp -lm -lpthread +endif clean :: @rm -f *.goto *.mkl *.acml *.atlas *.veclib *.essl smallscaling diff --git a/benchmark/cblasgemm.c b/benchmark/cblasgemm.c new file mode 100644 index 0000000000..0eb7d69cf8 --- /dev/null +++ b/benchmark/cblasgemm.c @@ -0,0 +1,233 @@ +/*************************************************************************** +Copyright (c) 2014, The OpenBLAS Project +All rights reserved. +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are +met: +1. Redistributions of source code must retain the above copyright +notice, this list of conditions and the following disclaimer. +2. Redistributions in binary form must reproduce the above copyright +notice, this list of conditions and the following disclaimer in +the documentation and/or other materials provided with the +distribution. +3. Neither the name of the OpenBLAS project nor the names of +its contributors may be used to endorse or promote products +derived from this software without specific prior written permission. +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +ARE DISCLAIMED. IN NO EVENT SHALL THE OPENBLAS PROJECT OR CONTRIBUTORS BE +LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE +USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +*****************************************************************************/ + +#include "bench.h" +#include "cblas.h" +#undef GEMM + +#ifndef COMPLEX + +#ifdef DOUBLE +#define GEMM cblas_dgemm +#elif defined(BFLOAT16) && defined(BGEMM) +#define GEMM cblas_bgemm +#elif defined(BFLOAT16) +#define GEMM cblas_sbgemm +#undef IFLOAT +#define IFLOAT bfloat16 +#elif defined(HFLOAT16) +#define GEMM cblas_shgemm +#undef IFLOAT +#define IFLOAT hfloat16 +#else +#define GEMM cblas_sgemm +#undef IFLOAT +#define IFLOAT float +#endif + +#else + +#ifdef DOUBLE +#define GEMM cblas_zgemm +#else +#define GEMM cblas_cgemm +#endif + +#endif + +int main(int argc, char *argv[]){ + + IFLOAT *a, *b; + //IFLOAT *aa, *bb; + FLOAT *c; + //FLOAT *cc; +#ifdef BGEMM + blasint one=1; + blasint two=2; + float alpha_in[] = {1.0, 0.0}; + float beta_in[] = {0.0, 0.0}; + FLOAT alpha[2], beta[2]; + sbstobf16_(&two, alpha_in, &one, alpha, &one); + sbstobf16_(&two, beta_in, &one, beta, &one); +#else +#ifdef COMPLEX + FLOAT alpha[] = {1.0, 0.0}; + FLOAT beta [] = {0.0, 0.0}; +#else + FLOAT alpha = 1.0; + FLOAT beta = 0.0; +#endif +#endif + CBLAS_TRANSPOSE transa = CblasNoTrans; + CBLAS_TRANSPOSE transb = CblasNoTrans; + char transac, transbc; + blasint m, n, k, i, j, lda, ldb, ldc; + int loops = 1; + int has_param_m = 0; + int has_param_n = 0; + int has_param_k = 0; + int has_param_lda = 0; + int has_param_ldb = 0; + char *p; +//blasint sme=0; + int from = 1; + int to = 200; + int step = 1; + + double time1, timeg; + + argc--;argv++; + + if (argc > 0) { from = atol(*argv); argc--; argv++; } + if (argc > 0) { to = MAX(atol(*argv), from); argc--; argv++; } + if (argc > 0) { step = atol(*argv); argc--; argv++; } + + if ((p = getenv("OPENBLAS_TRANS"))) { + transa=(*p=='N') ? CblasNoTrans : CblasTrans; + transb=(*p=='N') ? CblasNoTrans : CblasTrans; + } + if ((p = getenv("OPENBLAS_TRANSA"))) { + transa=(*p=='N') ? CblasNoTrans : CblasTrans; + } + if ((p = getenv("OPENBLAS_TRANSB"))) { + transb=(*p=='N') ? CblasNoTrans : CblasTrans; + } + + transac=(transa==CblasNoTrans) ? 'N' : 'T'; + transbc=(transb==CblasNoTrans) ? 'N' : 'T'; + fprintf(stderr, "From : %3d To : %3d Step=%d : Transa=%c : Transb=%c\n", from, to, step, transac, transbc); + + p = getenv("OPENBLAS_LOOPS"); + if ( p != NULL ) { + loops = atoi(p); + } + + if ((p = getenv("OPENBLAS_PARAM_M"))) { + m = atoi(p); + has_param_m=1; + } else { + m = to; + } + if ((p = getenv("OPENBLAS_PARAM_N"))) { + n = atoi(p); + has_param_n=1; + } else { + n = to; + } + if ((p = getenv("OPENBLAS_PARAM_K"))) { + k = atoi(p); + has_param_k=1; + } else { + k = to; + } + if ((p = getenv("OPENBLAS_PARAM_LDA"))) { + lda = atoi(p); + has_param_lda=1; + } + if ((p = getenv("OPENBLAS_PARAM_LDB"))) { + ldb = atoi(p); + has_param_ldb=1; + } + + if (( a = (IFLOAT *)malloc(sizeof(IFLOAT) * m * k * COMPSIZE)) == NULL) { + fprintf(stderr,"Out of Memory!!\n");exit(1); + } + if (( b = (IFLOAT *)malloc(sizeof(IFLOAT) * k * n * COMPSIZE)) == NULL) { + fprintf(stderr,"Out of Memory!!\n");exit(1); + } + if (( c = (FLOAT *)malloc(sizeof(FLOAT) * m * n * COMPSIZE)) == NULL) { + fprintf(stderr,"Out of Memory!!\n");exit(1); + } + //if (( aa = (IFLOAT *)malloc(sizeof(IFLOAT) * m * k * COMPSIZE)) == NULL) { + // fprintf(stderr,"Out of Memory!!\n");exit(1); + //} + //if (( bb = (IFLOAT *)malloc(sizeof(IFLOAT) * k * n * COMPSIZE)) == NULL) { + // fprintf(stderr,"Out of Memory!!\n");exit(1); + //} + //if (( cc = (FLOAT *)malloc(sizeof(FLOAT) * m * n * COMPSIZE)) == NULL) { + // fprintf(stderr,"Out of Memory!!\n");exit(1); + //} + +#ifdef __linux + srandom(getpid()); +#endif + + for (i = 0; i < m * k * COMPSIZE; i++) { + a[i] = ((IFLOAT) rand() / (IFLOAT) RAND_MAX) - 0.5; + // aa[i]=a[i]; + } + for (i = 0; i < k * n * COMPSIZE; i++) { + b[i] = ((IFLOAT) rand() / (IFLOAT) RAND_MAX) - 0.5; + // bb[i]=b[i]; + } + for (i = 0; i < m * n * COMPSIZE; i++) { + c[i] = ((FLOAT) rand() / (FLOAT) RAND_MAX) - 0.5; + // cc[i]=c[i]; + } + + fprintf(stderr, " SIZE Flops Time\n"); + + for (i = from; i <= to; i += step) { + + timeg=0; + + if (!has_param_m) { m = i; } + if (!has_param_n) { n = i; } + if (!has_param_k) { k = i; } + + if (!has_param_lda) { + if (transa == CblasNoTrans) { lda = k; } + else { lda = m; } + } + if (!has_param_ldb) { + if (transb == CblasNoTrans) { ldb = n; } + else { ldb = k; } + } + ldc = n; + + fprintf(stderr, " M=%4d, N=%4d, K=%4d : ", (int)m, (int)n, (int)k); + begin(); + + for (j=0; j1.5e-5){fprintf(stderr,"mismatch %d %f != %f: %g\n",ii,c[ii],cc[ii],fabsf(c[ii]-cc[ii]));} + end(); + time1 = getsec(); + + timeg = time1/loops; + fprintf(stderr, + " %10.2f MFlops %10.6f sec\n", + COMPSIZE * COMPSIZE * 2. * (double)k * (double)m * (double)n / timeg * 1.e-6, time1); + + } + + return 0; +} + +// void main(int argc, char *argv[]) __attribute__((weak, alias("MAIN__"))); diff --git a/benchmark/gemm.c b/benchmark/gemm.c index 704e332251..67b3a2e639 100644 --- a/benchmark/gemm.c +++ b/benchmark/gemm.c @@ -1,3 +1,4 @@ +//#pragma clang optimize off /*************************************************************************** Copyright (c) 2014, The OpenBLAS Project All rights reserved. @@ -45,6 +46,7 @@ USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. #define IFLOAT hfloat16 #else #define GEMM BLASFUNC(sgemm) +#undef IFLOAT #define IFLOAT float #endif @@ -186,7 +188,7 @@ int main(int argc, char *argv[]){ timeg = time1/loops; fprintf(stderr, - " %10.2f MFlops %10.6f sec\n", + " %10.2lf MFlops %10.6f sec\n", COMPSIZE * COMPSIZE * 2. * (double)k * (double)m * (double)n / timeg * 1.e-6, time1); } From 39b56e9a05c52a399cbe8c08197fea6e23bf225a Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 12:16:04 +0200 Subject: [PATCH 06/16] Add ARM64 SME GEMM kernels --- kernel/Makefile.L3 | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/kernel/Makefile.L3 b/kernel/Makefile.L3 index da961925d0..f151f5bb5f 100644 --- a/kernel/Makefile.L3 +++ b/kernel/Makefile.L3 @@ -264,7 +264,11 @@ SKERNELOBJS += \ ifdef USE_SME SKERNELOBJS += \ sgemm_direct_sme1_2VLx2VL$(TSUFFIX).$(SUFFIX) \ - sgemm_direct_sme1_preprocess$(TSUFFIX).$(SUFFIX) + sgemm_direct_sme1_preprocess$(TSUFFIX).$(SUFFIX) \ + sme_sgemm_kernel$(TSUFFIX).$(SUFFIX) \ + sme_dgemm_kernel$(TSUFFIX).$(SUFFIX) \ + sme_cgemm_kernel$(TSUFFIX).$(SUFFIX) \ + sme_zgemm_kernel$(TSUFFIX).$(SUFFIX) endif endif endif @@ -1064,6 +1068,14 @@ $(KDIR)sgemm_direct_sme1_2VLx2VL$(TSUFFIX).$(SUFFIX) : $(CC) $(CFLAGS) -c $(KERNELDIR)/sgemm_direct_sme1_2VLx2VL.S -UDOUBLE -UCOMPLEX -o $@ $(KDIR)sgemm_direct_sme1_preprocess$(TSUFFIX).$(SUFFIX) : $(CC) $(CFLAGS) -c $(KERNELDIR)/sgemm_direct_sme1_preprocess.S -UDOUBLE -UCOMPLEX -o $@ +$(KDIR)sme_sgemm_kernel$(TSUFFIX).$(SUFFIX) : + $(CC) $(CFLAGS) -c $(KERNELDIR)/sme_sgemm_kernel.c -UDOUBLE -UCOMPLEX -o $@ +$(KDIR)sme_dgemm_kernel$(TSUFFIX).$(SUFFIX) : + $(CC) $(CFLAGS) -c $(KERNELDIR)/sme_dgemm_kernel.c -DDOUBLE -UCOMPLEX -o $@ +$(KDIR)sme_cgemm_kernel$(TSUFFIX).$(SUFFIX) : + $(CC) $(CFLAGS) -c $(KERNELDIR)/sme_cgemm_kernel.c -UDOUBLE -DCOMPLEX -o $@ +$(KDIR)sme_zgemm_kernel$(TSUFFIX).$(SUFFIX) : + $(CC) $(CFLAGS) -c $(KERNELDIR)/sme_zgemm_kernel.c -DDOUBLE -DCOMPLEX -o $@ endif endif endif From e24e0da7793f9095d2bb570204c02f9ea498c9f0 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Tue, 11 Aug 2026 17:58:17 +0200 Subject: [PATCH 07/16] Add sme-f64f64 to ARMV9SME archflags too --- Makefile.arm64 | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/Makefile.arm64 b/Makefile.arm64 index bfa0199bbe..900956a44f 100644 --- a/Makefile.arm64 +++ b/Makefile.arm64 @@ -59,7 +59,7 @@ endif endif ifeq ($(CORE), ARMV9SME) -CCOMMON_OPT += -march=armv9-a+sve2+sme +CCOMMON_OPT += -march=armv9-a+sve2+sme+sme-f64f64 FCOMMON_OPT += -march=armv9-a+sve2 ifdef OS_WINDOWS ifeq ($(C_COMPILER), CLANG) From 1ed99815fa9ad10f594e254b2835002f3566ff1b Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Wed, 12 Aug 2026 10:53:07 +0200 Subject: [PATCH 08/16] Add SME GEMM kernels --- kernel/CMakeLists.txt | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/kernel/CMakeLists.txt b/kernel/CMakeLists.txt index f09c391a38..1a1443a6a3 100644 --- a/kernel/CMakeLists.txt +++ b/kernel/CMakeLists.txt @@ -290,6 +290,10 @@ function (build_core TARGET_CORE KDIR TSUFFIX KERNEL_DEFINITIONS) if (HAVE_SME) GenerateNamedObjects("${KERNELDIR}/${SGEMMDIRECTSMEKERNEL}" "" "gemm_direct_sme1_2VLx2VL" false "" "" false SINGLE) GenerateNamedObjects("${KERNELDIR}/${SGEMMDIRECTPREKERNEL}" "" "gemm_direct_sme1_preprocess" false "" "" false SINGLE) + GenerateNamedObjects("${KERNELDIR}/sme_sgemm_kernel.c" "" "sme_*gemm_kernel" false "" "" false "SINGLE") + GenerateNamedObjects("${KERNELDIR}/sme_dgemm_kernel.c" "" "sme_*gemm_kernel" false "" "" false "DOUBLE") + GenerateNamedObjects("${KERNELDIR}/sme_cgemm_kernel.c" "" "sme_*gemm_kernel" false "" "" false "COMPLEX") + GenerateNamedObjects("${KERNELDIR}/sme_zgemm_kernel.c" "" "sme_*gemm_kernel" false "" "" false "ZCOMPLEX") endif () endif () endif() From 4ae369ac0a79aa4272c65f34a9cd59ff3c2aad23 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Wed, 12 Aug 2026 10:54:27 +0200 Subject: [PATCH 09/16] Make the compute kernel static --- kernel/arm64/sme_zgemm_kernel.c | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/kernel/arm64/sme_zgemm_kernel.c b/kernel/arm64/sme_zgemm_kernel.c index 71744e64c3..805db84f8c 100644 --- a/kernel/arm64/sme_zgemm_kernel.c +++ b/kernel/arm64/sme_zgemm_kernel.c @@ -33,7 +33,7 @@ zdouble cdone={1.,0.}; #define KC 256 #define NC 512 -inline void zgemm_sme_compute_8x8_tile(int current_K, const double *A_ptr, const double *B_ptr, zdouble *C_ptr, size_t ldc, int beta_mode, const zdouble *beta_ptr) +static inline void zgemm_sme_compute_8x8_tile(int current_K, const double *A_ptr, const double *B_ptr, zdouble *C_ptr, size_t ldc, int beta_mode, const zdouble *beta_ptr) { size_t ldc_bytes = ldc * sizeof(zdouble); From 9724481b599bb75fd260a9e7e991b445ef02a0bd Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 18:38:55 +0200 Subject: [PATCH 10/16] Add declarations for ARM64 SME GEMM kernels --- common_c.h | 2 ++ common_d.h | 2 ++ common_level3.h | 25 +++++++++++++++++++++++++ common_param.h | 11 +++++++++++ common_s.h | 2 ++ common_z.h | 2 ++ 6 files changed, 44 insertions(+) diff --git a/common_c.h b/common_c.h index 6cff610bb5..88f8f5dcdf 100644 --- a/common_c.h +++ b/common_c.h @@ -119,6 +119,7 @@ #endif #define CGEMM_BETA cgemm_beta +#define SME_CGEMM_KERNEL sme_cgemm_kernel #define CGEMM_KERNEL_N cgemm_kernel_n #define CGEMM_KERNEL_L cgemm_kernel_l @@ -326,6 +327,7 @@ #define CTRSM_ILTNCOPY gotoblas -> ctrsm_iltncopy #define CGEMM_BETA gotoblas -> cgemm_beta +#define SME_CGEMM_KERNEL gotoblas -> sme_cgemm_kernel #define CGEMM_KERNEL_N gotoblas -> cgemm_kernel_n #define CGEMM_KERNEL_L gotoblas -> cgemm_kernel_l #define CGEMM_KERNEL_R gotoblas -> cgemm_kernel_r diff --git a/common_d.h b/common_d.h index 1e8c33d7a3..a7a1e1cd14 100644 --- a/common_d.h +++ b/common_d.h @@ -114,6 +114,7 @@ #define DGEMM_BETA dgemm_beta #define DGEMM_KERNEL dgemm_kernel +#define SME_DGEMM_KERNEL sme_dgemm_kernel #define DTRMM_KERNEL_LN dtrmm_kernel_LN #define DTRMM_KERNEL_LT dtrmm_kernel_LT @@ -246,6 +247,7 @@ #define DGEMM_BETA gotoblas -> dgemm_beta #define DGEMM_KERNEL gotoblas -> dgemm_kernel +#define SME_DGEMM_KERNEL gotoblas -> sme_dgemm_kernel #define DTRMM_KERNEL_LN gotoblas -> dtrmm_kernel_LN #define DTRMM_KERNEL_LT gotoblas -> dtrmm_kernel_LT diff --git a/common_level3.h b/common_level3.h index 8399b5b013..3e4edf1950 100644 --- a/common_level3.h +++ b/common_level3.h @@ -135,6 +135,31 @@ void ssyr2k_direct_alpha_betaLT(BLASLONG N, BLASLONG K, float beta, float * R, BLASLONG strideR); +void sme_sgemm_kernel(char*, char*, BLASLONG M, BLASLONG N, BLASLONG K, + float * alpha, + float * A, BLASLONG ldA, + float * B, BLASLONG ldB, + float * beta, + float * C, BLASLONG ldC); +void sme_dgemm_kernel(const char*, const char*, const BLASLONG M, const BLASLONG N, const BLASLONG K, + const double * alpha, + const double * A, const BLASLONG ldA, + const double * B, const BLASLONG ldB, + const double * beta, + double * C, const BLASLONG ldC); +void sme_cgemm_kernel(const char*, const char*, const BLASLONG M, const BLASLONG N, const BLASLONG K, + const float alpha_r, const float alpha_i, + const float * A, const BLASLONG ldA, + const float * B, const BLASLONG ldB, + const float beta_r, const float beta_i, + float * C, const BLASLONG ldC); +void sme_zgemm_kernel(const char*, const char*, const BLASLONG M, const BLASLONG N, const BLASLONG K, + const double alpha_r, const double alpha_i, + const double * A, const BLASLONG ldA, + const double * B, const BLASLONG ldB, + const double beta_r, const double beta_i, + double * C, const BLASLONG ldC); + int sgemm_direct_performant(BLASLONG M, BLASLONG N, BLASLONG K); int shgemm_beta(BLASLONG, BLASLONG, BLASLONG, float, diff --git a/common_param.h b/common_param.h index fd20945c4e..7d21652a5c 100644 --- a/common_param.h +++ b/common_param.h @@ -276,6 +276,7 @@ int (*shgemv_t) (BLASLONG, BLASLONG, float, hfloat16 *, BLASLONG, hfloat16 *, BL void (*ssyr2k_direct_alpha_betaUT) (BLASLONG, BLASLONG, float, float *, BLASLONG, float *, BLASLONG, float, float *, BLASLONG); void (*ssyr2k_direct_alpha_betaLN) (BLASLONG, BLASLONG, float, float *, BLASLONG, float *, BLASLONG, float, float *, BLASLONG); void (*ssyr2k_direct_alpha_betaLT) (BLASLONG, BLASLONG, float, float *, BLASLONG, float *, BLASLONG, float, float *, BLASLONG); + void (*sme_sgemm_kernel) (char*, char*, BLASLONG, BLASLONG, BLASLONG, float*, float *, BLASLONG , float *, BLASLONG ,float*, float *, BLASLONG); #endif @@ -401,6 +402,9 @@ int (*shgemv_t) (BLASLONG, BLASLONG, float, hfloat16 *, BLASLONG, hfloat16 *, BL int (*dsymv_U) (BLASLONG, BLASLONG, double, double *, BLASLONG, double *, BLASLONG, double *, BLASLONG, double *); #endif #if (BUILD_DOUBLE==1) || (BUILD_COMPLEX16==1) +#ifdef ARCH_ARM64 + void (*sme_dgemm_kernel) (const char*, const char*, const BLASLONG, const BLASLONG, const BLASLONG, const double*, const double *, const BLASLONG , const double *, const BLASLONG ,const double*, double *, const BLASLONG); +#endif int (*dgemm_kernel )(BLASLONG, BLASLONG, BLASLONG, double, double *, double *, double *, BLASLONG); int (*dgemm_beta )(BLASLONG, BLASLONG, BLASLONG, double, double *, BLASLONG, double *, BLASLONG, double *, BLASLONG); @@ -616,6 +620,9 @@ int (*shgemv_t) (BLASLONG, BLASLONG, float, hfloat16 *, BLASLONG, hfloat16 *, BL int (*chemv_M) (BLASLONG, BLASLONG, float, float, float *, BLASLONG, float *, BLASLONG, float *, BLASLONG, float *); int (*chemv_V) (BLASLONG, BLASLONG, float, float, float *, BLASLONG, float *, BLASLONG, float *, BLASLONG, float *); +#ifdef ARCH_ARM64 + void (*sme_cgemm_kernel) (const char*, const char*, const BLASLONG, const BLASLONG, const BLASLONG, const float, const float, const float *, const BLASLONG , const float *, const BLASLONG, const float, const float, float *, const BLASLONG); +#endif int (*cgemm_kernel_n )(BLASLONG, BLASLONG, BLASLONG, float, float, float *, float *, float *, BLASLONG); int (*cgemm_kernel_l )(BLASLONG, BLASLONG, BLASLONG, float, float, float *, float *, float *, BLASLONG); int (*cgemm_kernel_r )(BLASLONG, BLASLONG, BLASLONG, float, float, float *, float *, float *, BLASLONG); @@ -826,6 +833,10 @@ int (*shgemv_t) (BLASLONG, BLASLONG, float, hfloat16 *, BLASLONG, hfloat16 *, BL int (*zhemv_M) (BLASLONG, BLASLONG, double, double, double *, BLASLONG, double *, BLASLONG, double *, BLASLONG, double *); int (*zhemv_V) (BLASLONG, BLASLONG, double, double, double *, BLASLONG, double *, BLASLONG, double *, BLASLONG, double *); +#ifdef ARCH_ARM64 + void (*sme_zgemm_kernel) (const char*, const char*, const BLASLONG, const BLASLONG, const BLASLONG, const double, const double, const double *, const BLASLONG , const double *, const BLASLONG, const double, const double, double *, const BLASLONG); +#endif + int (*zgemm_kernel_n )(BLASLONG, BLASLONG, BLASLONG, double, double, double *, double *, double *, BLASLONG); int (*zgemm_kernel_l )(BLASLONG, BLASLONG, BLASLONG, double, double, double *, double *, double *, BLASLONG); int (*zgemm_kernel_r )(BLASLONG, BLASLONG, BLASLONG, double, double, double *, double *, double *, BLASLONG); diff --git a/common_s.h b/common_s.h index e43b74b917..6df9ba0f1e 100644 --- a/common_s.h +++ b/common_s.h @@ -76,6 +76,7 @@ #define SGEMM_ITCOPY sgemm_itcopy #endif +#define SME_SGEMM_KERNEL sme_sgemm_kernel #define STRMM_OUNUCOPY strmm_ounucopy #define STRMM_OUNNCOPY strmm_ounncopy #define STRMM_OUTUCOPY strmm_outucopy @@ -248,6 +249,7 @@ #define SSYR2K_DIRECT_ALPHA_BETA_UT gotoblas -> ssyr2k_direct_alpha_betaUT #define SSYR2K_DIRECT_ALPHA_BETA_LN gotoblas -> ssyr2k_direct_alpha_betaLN #define SSYR2K_DIRECT_ALPHA_BETA_LT gotoblas -> ssyr2k_direct_alpha_betaLT +#define SME_SGEMM_KERNEL gotoblas -> sme_sgemm_kernel #endif #define SGEMM_ONCOPY gotoblas -> sgemm_oncopy diff --git a/common_z.h b/common_z.h index c12d71b390..5a9fb88b77 100644 --- a/common_z.h +++ b/common_z.h @@ -119,6 +119,7 @@ #endif #define ZGEMM_BETA zgemm_beta +#define SME_ZGEMM_KERNEL sme_zgemm_kernel #define ZGEMM_KERNEL_N zgemm_kernel_n #define ZGEMM_KERNEL_L zgemm_kernel_l @@ -326,6 +327,7 @@ #define ZTRSM_ILTNCOPY gotoblas -> ztrsm_iltncopy #define ZGEMM_BETA gotoblas -> zgemm_beta +#define SME_ZGEMM_KERNEL gotoblas -> sme_zgemm_kernel #define ZGEMM_KERNEL_N gotoblas -> zgemm_kernel_n #define ZGEMM_KERNEL_L gotoblas -> zgemm_kernel_l #define ZGEMM_KERNEL_R gotoblas -> zgemm_kernel_r From c073f087b4c3aee0bd93e680bd12050e66aebff7 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 18:39:56 +0200 Subject: [PATCH 11/16] Add ARM64 SME GEMM kernels --- kernel/setparam-ref.c | 27 ++++++++++++++++++++++++++- 1 file changed, 26 insertions(+), 1 deletion(-) diff --git a/kernel/setparam-ref.c b/kernel/setparam-ref.c index 107f1f4f0f..7db921b713 100644 --- a/kernel/setparam-ref.c +++ b/kernel/setparam-ref.c @@ -238,6 +238,11 @@ gotoblas_t TABLE_NAME = { ssyr2k_direct_alpha_betaUTTS, ssyr2k_direct_alpha_betaLNTS, ssyr2k_direct_alpha_betaLTTS, +#ifdef HAVE_SME + sme_sgemm_kernelTS, +#else + NULL, +#endif #endif sgemm_kernelTS, sgemm_betaTS, @@ -332,6 +337,13 @@ gotoblas_t TABLE_NAME = { #endif #if (BUILD_DOUBLE==1) || (BUILD_COMPLEX16==1) +#ifdef ARCH_ARM64 +#ifdef HAVE_SME + sme_dgemm_kernelTS, +#else + NULL, +#endif +#endif dgemm_kernelTS, dgemm_betaTS, #if DGEMM_DEFAULT_UNROLL_M != DGEMM_DEFAULT_UNROLL_N dgemm_incopyTS, dgemm_itcopyTS, @@ -476,6 +488,13 @@ gotoblas_t TABLE_NAME = { chemv_LTS, chemv_UTS, chemv_MTS, chemv_VTS, #endif #if (BUILD_COMPLEX) +#ifdef ARCH_ARM64 +#ifdef HAVE_SME + sme_cgemm_kernelTS, +#else + NULL, +#endif +#endif cgemm_kernel_nTS, cgemm_kernel_lTS, cgemm_kernel_rTS, cgemm_kernel_bTS, cgemm_betaTS, #if CGEMM_DEFAULT_UNROLL_M != CGEMM_DEFAULT_UNROLL_N @@ -631,7 +650,13 @@ gotoblas_t TABLE_NAME = { zgeru_kTS, zgerc_kTS, zgerv_kTS, zgerd_kTS, zsymv_LTS, zsymv_UTS, zhemv_LTS, zhemv_UTS, zhemv_MTS, zhemv_VTS, - +#ifdef ARCH_ARM64 +#ifdef HAVE_SME + sme_zgemm_kernelTS, +#else + NULL, +#endif +#endif zgemm_kernel_nTS, zgemm_kernel_lTS, zgemm_kernel_rTS, zgemm_kernel_bTS, zgemm_betaTS, From 3ead57fd2bdc81b2351d225267d7d663ce1b22b7 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 18:41:33 +0200 Subject: [PATCH 12/16] Improve clobber lists and interfaces --- kernel/arm64/sme_cgemm_kernel.c | 58 ++++++++++++++--------------- kernel/arm64/sme_dgemm_kernel.c | 48 +++++++++++------------- kernel/arm64/sme_sgemm_kernel.c | 58 ++++++++++------------------- kernel/arm64/sme_zgemm_kernel.c | 65 +++++++++++++++------------------ 4 files changed, 98 insertions(+), 131 deletions(-) diff --git a/kernel/arm64/sme_cgemm_kernel.c b/kernel/arm64/sme_cgemm_kernel.c index e2622fcda4..e613b1b343 100644 --- a/kernel/arm64/sme_cgemm_kernel.c +++ b/kernel/arm64/sme_cgemm_kernel.c @@ -9,7 +9,7 @@ #define stdmin(a,b) (a>b? b:a) #endif typedef float _Complex cfloat; -cfloat CMUL(cfloat a, cfloat b,bool conja, bool conjb) { +static cfloat CMUL(cfloat a, cfloat b,bool conja, bool conjb) { float ra=creal(a); float rb=creal(b); float ia=conja ? -cimag(a) : cimag(a); @@ -24,10 +24,10 @@ return res; #define KERNEL_ALPHA 0 #define USE_VECTORIZED_PACKING 1 -//using cfloat = std::complex; -cfloat czero={0.,0.}; -cfloat cone={1.,0.}; +static cfloat czero={0.,0.}; +static cfloat cone={1.,0.}; + #define MC 256 #define KC 512 #define NC 1024 @@ -36,14 +36,6 @@ static inline void cgemm_sme_compute_16x16_tile(blasint current_K, const float * { size_t ldc_bytes = ldc * sizeof(cfloat); - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); - - asm volatile("smstart\n\t" "ptrue p0.s\n\t" "zero {za}\n\t" @@ -141,22 +133,26 @@ static inline void cgemm_sme_compute_16x16_tile(blasint current_K, const float * "cmp w12, #16\n\t" "b.ne 200b\n\t" "smstop\n\t" - : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + "msr fpsr, xzr\n\t" + : [a] "+&r"(A_ptr), [b] "+&r"(B_ptr) : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) // Updated Clobber list to cover x15 and z8-z9, z30-z31 - : "p0", "x10", "w12", "x13", "x14", "x15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z30", "z31", "za" ); - - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + : "x10", "x12", "x13", "x14", "x15", "cc", "memory", + "v0", "v1", "v2", "v3", "v4", "v5", "v6", "v7", + "v8", "v9", "v10", "v11", "v12", "v13", "v14", "v15", + "v16", "v17", "v18", "v19", "v20", "v21", "v22", "v23", + "v24", "v25", "v26", "v27", "v28", "v29", "v30", "v31", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31", "za", + "p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15"); } -void cgemm_sme_NN(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +static void cgemm_sme_NN(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) { if (alpha == czero || K == 0) { @@ -249,7 +245,7 @@ void cgemm_sme_NN(int M, int N, int K, const cfloat alpha, const cfloat *A, int } } -void cgemm_sme_TN(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +static void cgemm_sme_TN(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) { if (alpha == czero || K == 0) { if (beta == czero) { @@ -339,7 +335,7 @@ void cgemm_sme_TN(int M, int N, int K, const cfloat alpha, const cfloat *A, int } } -void cgemm_sme_NT(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +static void cgemm_sme_NT(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) { if (alpha == czero || K == 0) { if (beta == czero) { @@ -429,7 +425,7 @@ void cgemm_sme_NT(int M, int N, int K, const cfloat alpha, const cfloat *A, int } } -void cgemm_sme_TT(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) +static void cgemm_sme_TT(int M, int N, int K, const cfloat alpha, const cfloat *A, int lda, const cfloat *b, int ldb, const cfloat beta, cfloat *C, int ldc, bool conja, bool conjb) { if (alpha == czero || K == 0) { if (beta == czero) { @@ -519,28 +515,30 @@ void cgemm_sme_TT(int M, int N, int K, const cfloat alpha, const cfloat *A, int } } -void sme_CGEMM_KERNEL(const char *transa, const char *transb, const blasint *m, const blasint *n, const blasint *k, const cfloat alpha, const cfloat *a, const blasint *lda, const cfloat *b, const blasint *ldb, const cfloat beta, cfloat *c, const blasint *ldc) +void CNAME(const char *transa, const char *transb, const BLASLONG m, const BLASLONG n, const BLASLONG k, const float alpha_r, const float alpha_i, const float *a, const BLASLONG lda, const float *b, const BLASLONG ldb, const float beta_r, const float beta_i, float *c, const BLASLONG ldc) { +cfloat alpha={alpha_r,alpha_i}; +cfloat beta={beta_r,beta_i}; bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); if (!trans_a && !trans_b) { bool conja=(*transa == 'R' || *transa == 'r'); bool conjb=(*transb == 'R' || *transb == 'r'); - cgemm_sme_NN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + cgemm_sme_NN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } else if (trans_a && !trans_b) { bool conja=(*transa == 'C' || *transa == 'c'); bool conjb=(*transb == 'R' || *transb == 'r'); - cgemm_sme_TN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + cgemm_sme_TN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } else if (!trans_a && trans_b) { bool conja=(*transa == 'R' || *transa == 'r'); bool conjb=(*transb == 'C' || *transb == 'c'); - cgemm_sme_NT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + cgemm_sme_NT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } else { bool conja=(*transa == 'C' || *transa == 'c'); bool conjb=(*transb == 'C' || *transb == 'c'); - cgemm_sme_TT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + cgemm_sme_TT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } } diff --git a/kernel/arm64/sme_dgemm_kernel.c b/kernel/arm64/sme_dgemm_kernel.c index 7333c668c6..2cc02c952d 100644 --- a/kernel/arm64/sme_dgemm_kernel.c +++ b/kernel/arm64/sme_dgemm_kernel.c @@ -28,14 +28,6 @@ static inline void dgemm_sme_compute_16x16_tile(int current_K, const double *A_p { ptrdiff_t ldc_bytes = ldc * sizeof(double); - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); - - asm volatile( "smstart\n\t" // Enable all 64-bit (double precision) lanes in predicate register p0 @@ -219,17 +211,19 @@ static inline void dgemm_sme_compute_16x16_tile(int current_K, const double *A_p "cmp w15, #8\n\t" "b.ne 201b\n\t" "smstop\n\t" - : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + "msr fpsr, xzr\n\t" + : [a] "+&r"(A_ptr), [b] "+&r"(B_ptr) : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) - : "p0","x10", "x11", "w12", "x13", "x14", "w15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", "z31", "za"); - - - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + : "x10", "x11", "x12", "x13", "x14", "w15", "memory", "cc", + "v0", "v1", "v2", "v3", "v4", "v5", "v6", "v7", "v8", "v9", + "v10", "v11", "v12", "v13", "v14", "v15", "v16", "v17", "v18", + "v19", "v20", "v21", "v22", "v23", "v24", "v25", "v26", "v27", + "v28", "v29", "v30", "v31", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", + "z10", "z11", "z12", "z13", "z14", "z15", "z16", "z17", "z18", + "z19", "z20", "z21", "z22", "z23", "z24", "z25", "z26", "z27", + "z28", "z29", "z30", "z31", "za", "p0", "p1", "p2", "p3", "p4", + "p5", "p6", "p7", "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15"); } @@ -237,7 +231,7 @@ static inline void dgemm_sme_compute_16x16_tile(int current_K, const double *A_p // ======================================================================== // VERSION 0: C = alpha * A * B + beta * C (NN) // ======================================================================== -void dgemm_sme_NN(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +static void dgemm_sme_NN(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) { #if PROFILING CALI_CXX_MARK_FUNCTION; @@ -520,7 +514,7 @@ void dgemm_sme_NN(int M, int N, int K, double alpha, const double *A, int lda, c // ======================================================================== // VERSION 1: C = alpha * A^T * B + beta * C (TN) // ======================================================================== -void dgemm_sme_TN(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +static void dgemm_sme_TN(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) { if (alpha == 0.0 || K == 0) { if (beta != 1.0) { @@ -751,7 +745,7 @@ void dgemm_sme_TN(int M, int N, int K, double alpha, const double *A, int lda, c // ======================================================================== // VERSION 2: C = alpha * A * B^T + beta * C (NT) // ======================================================================== -void dgemm_sme_NT(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +static void dgemm_sme_NT(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) { if (alpha == 0.0 || K == 0) { if (beta != 1.0) { @@ -972,7 +966,7 @@ void dgemm_sme_NT(int M, int N, int K, double alpha, const double *A, int lda, c // ======================================================================== // VERSION 3: C = alpha * A^T * B^T + beta * C (TT) // ======================================================================== -void dgemm_sme_TT(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) +static void dgemm_sme_TT(int M, int N, int K, double alpha, const double *A, int lda, const double *B, int ldb, double beta, double *C, int ldc) { if (alpha == 0.0 || K == 0) { if (beta != 1.0) { @@ -1187,20 +1181,20 @@ void dgemm_sme_TT(int M, int N, int K, double alpha, const double *A, int lda, c // ======================================================================== // WRAPPER: BLAS ABI Compatible dgemm // ======================================================================== -void sme_DGEMM_KERNEL(const char *transa, const char *transb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const double *alpha, const double *a, const BLASLONG *lda, const double *b, const BLASLONG *ldb, const double *beta, double *c, const BLASLONG *ldc) +void CNAME(const char *transa, const char *transb, const BLASLONG m, const BLASLONG n, const BLASLONG k, const double *alpha, const double *a, const BLASLONG lda, const double *b, const BLASLONG ldb, const double *beta, double *c, const BLASLONG ldc) { bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); if (!trans_a && !trans_b) { - dgemm_sme_NN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + dgemm_sme_NN(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } else if (trans_a && !trans_b) { - dgemm_sme_TN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + dgemm_sme_TN(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } else if (!trans_a && trans_b) { - dgemm_sme_NT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + dgemm_sme_NT(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } else { - dgemm_sme_TT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + dgemm_sme_TT(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } } diff --git a/kernel/arm64/sme_sgemm_kernel.c b/kernel/arm64/sme_sgemm_kernel.c index 5d5925ad8c..4fcc5cb566 100644 --- a/kernel/arm64/sme_sgemm_kernel.c +++ b/kernel/arm64/sme_sgemm_kernel.c @@ -8,21 +8,7 @@ #ifndef stdmin #define stdmin(a,b) (a>b? b:a) #endif - void SMEStart() - { - asm volatile("smstart\n\t" ::: "memory"); - } - void SMEStop() - { - asm volatile("smstop\n\t" ::: "memory"); - } -// SMEGuard(const SMEGuard &) = delete; -// SMEGuard &operator=(const SMEGuard &) = delete; -/* -const int MC = 512; -const int KC = 1024; -const int NC = 2048; -*/ + #define MC 512 #define KC 1024 #define NC 2048 @@ -30,15 +16,7 @@ const int NC = 2048; static inline void sgemm_sme_compute_32x32_tile(blasint current_K, const float *A_ptr, const float *B_ptr, float *C_ptr, size_t ldc, blasint beta_mode, const float *beta_ptr) { size_t ldc_bytes = ldc * sizeof(float); - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); - - //SMEGuard __stream_guard; asm volatile("smstart\n\t" "ptrue p0.s\n\t" "cmp %w[beta_mode], #0\n\t" @@ -144,21 +122,25 @@ static inline void sgemm_sme_compute_32x32_tile(blasint current_K, const float * "cmp w15, #16\n\t" "b.ne 201b\n\t" "smstop\n\t" - : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + "msr fpsr, xzr\n\t" + : [a] "+&r"(A_ptr), [b] "+&r"(B_ptr) : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) - : "p0", "x10", "w12", "x13", "x14", "w15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z31", "za"); - - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", + : "x10", "x12", "x13", "x14", "x15", "memory", "cc", + "v0", "v1", "v2", "v3", "v4", "v5", "v6", "v7", + "v8", "v9", "v10", "v11", "v12", "v13", "v14", "v15", + "v16", "v17", "v18", "v19", "v20", "v21", "v22", "v23", + "v24", "v25", "v26", "v27", "v28", "v29", "v30", "v31", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31", + "p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", + "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15","za"); } -void sgemm_sme_NN(blasint M, blasint N, blasint K, float alpha, const float *A, blasint lda, const float *B, blasint ldb, float beta, float *C, blasint ldc) +static void sgemm_sme_NN(blasint M, blasint N, blasint K, float alpha, const float *A, blasint lda, const float *B, blasint ldb, float beta, float *C, blasint ldc) { if (alpha == 0.0f || K == 0) { if (beta != 1.0f) { @@ -277,7 +259,7 @@ void sgemm_sme_NN(blasint M, blasint N, blasint K, float alpha, const float *A, } } -void sgemm_sme_TN(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) +static void sgemm_sme_TN(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) { if (alpha == 0.0f || K == 0) { if (beta != 1.0f) { @@ -398,7 +380,7 @@ void sgemm_sme_TN(int M, int N, int K, float alpha, const float *A, int lda, con } } -void sgemm_sme_NT(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) +static void sgemm_sme_NT(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) { if (alpha == 0.0f || K == 0) { if (beta != 1.0f) { @@ -514,7 +496,7 @@ void sgemm_sme_NT(int M, int N, int K, float alpha, const float *A, int lda, con } } -void sgemm_sme_TT(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) +static void sgemm_sme_TT(int M, int N, int K, float alpha, const float *A, int lda, const float *B, int ldb, float beta, float *C, int ldc) { if (alpha == 0.0f || K == 0) { if (beta != 1.0f) { @@ -632,20 +614,20 @@ void sgemm_sme_TT(int M, int N, int K, float alpha, const float *A, int lda, con } } -/*extern "C"*/ void sme_SGEMM_KERNEL(const char *transa, const char *transb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const float *alpha, const float *a, const BLASLONG *lda, const float *b, const BLASLONG *ldb, const float *beta, float *c, const BLASLONG *ldc) +void CNAME(char *transa, char *transb, BLASLONG m, BLASLONG n, BLASLONG k, float *alpha, float *a, BLASLONG lda, float *b, BLASLONG ldb, float *beta, float *c, BLASLONG ldc) { bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); if (!trans_a && !trans_b) { - sgemm_sme_NN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + sgemm_sme_NN(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } else if (trans_a && !trans_b) { - sgemm_sme_TN(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + sgemm_sme_TN(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } else if (!trans_a && trans_b) { - sgemm_sme_NT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + sgemm_sme_NT(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } else { - sgemm_sme_TT(*m, *n, *k, *alpha, a, *lda, b, *ldb, *beta, c, *ldc); + sgemm_sme_TT(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } } diff --git a/kernel/arm64/sme_zgemm_kernel.c b/kernel/arm64/sme_zgemm_kernel.c index 805db84f8c..5e41010d8c 100644 --- a/kernel/arm64/sme_zgemm_kernel.c +++ b/kernel/arm64/sme_zgemm_kernel.c @@ -9,7 +9,7 @@ typedef double _Complex zdouble; -zdouble CDMUL(zdouble a, zdouble b,bool conja, bool conjb) { +static zdouble CDMUL(zdouble a, zdouble b,bool conja, bool conjb) { double ra=creal(a); double rb=creal(b); double ia=conja ? -cimag(a) : cimag(a); @@ -23,10 +23,8 @@ return res; } -zdouble cdzero={0.,0.}; -zdouble cdone={1.,0.}; - - +static zdouble cdzero={0.,0.}; +static zdouble cdone={1.,0.}; #define MC 128 @@ -37,15 +35,6 @@ static inline void zgemm_sme_compute_8x8_tile(int current_K, const double *A_ptr { size_t ldc_bytes = ldc * sizeof(zdouble); - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); - - - asm volatile("smstart\n\t" "ptrue p0.d\n\t" "zero {za}\n\t" @@ -142,24 +131,25 @@ static inline void zgemm_sme_compute_8x8_tile(int current_K, const double *A_ptr "add w12, w12, #1\n\t" "cmp w12, #8\n\t" "b.ne 200b\n\t" - "smstop\n\t" - : [a] "+r"(A_ptr), [b] "+r"(B_ptr) + "smstop\n\t" + "msr fpsr, xzr\n\t" + : [a] "+&r"(A_ptr), [b] "+&r"(B_ptr) : [k] "r"(current_K), [c] "r"(C_ptr), [ldc_bytes] "r"(ldc_bytes), [beta_mode] "r"(beta_mode), [beta_ptr] "r"(beta_ptr) - // Updated Clobber list to cover all the new registers utilized - : "p0", "x10", "w12", "x13", "x14", "x15", "p0", "memory", "cc", "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", "z8", "z9", "z30", "z31", "za"); - - - asm volatile("" : : :"p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", - "p8", "p9", "p10", "p11", "p12", "p13", "p14", "p15", "d8", "d9", "d10", "d11", "d12", "d13", "d14", "d15", - "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", - "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", - "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", - "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31"); - + : "x10", "x12", "x13", "x14", "x15", "memory", "cc", + "v0", "v1", "v2", "v3", "v4", "v5", "v6", "v7", + "v8", "v9", "v10", "v11", "v12", "v13", "v14", "v15", + "v16", "v17", "v18", "v19", "v20", "v21", "v22", "v23", + "v24", "v25", "v26", "v27", "v28", "v29", "v30", "v31", + "z0", "z1", "z2", "z3", "z4", "z5", "z6", "z7", + "z8", "z9", "z10", "z11", "z12", "z13", "z14", "z15", + "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", + "z24", "z25", "z26", "z27", "z28", "z29", "z30", "z31", "za", + "p0", "p1", "p2", "p3", "p4", "p5", "p6", "p7", "p8", + "p9", "p10", "p11", "p12", "p13", "p14", "p15"); } -void zgemm_sme_NN(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) +static void zgemm_sme_NN(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) { if (alpha == cdzero || K == 0) { @@ -245,7 +235,7 @@ void zgemm_sme_NN(int M, int N, int K, const zdouble alpha, const zdouble *A, in } } -void zgemm_sme_TN(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) +static void zgemm_sme_TN(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) { if (alpha == cdzero || K == 0) { if (beta != cdone) { @@ -330,7 +320,7 @@ void zgemm_sme_TN(int M, int N, int K, const zdouble alpha, const zdouble *A, in } } -void zgemm_sme_NT(blasint M, blasint N, blasint K, const zdouble alpha, const zdouble *A, blasint lda, const zdouble *b, blasint ldb, const zdouble beta, zdouble *C, blasint ldc, bool conja, bool conjb) +static void zgemm_sme_NT(blasint M, blasint N, blasint K, const zdouble alpha, const zdouble *A, blasint lda, const zdouble *b, blasint ldb, const zdouble beta, zdouble *C, blasint ldc, bool conja, bool conjb) { #if 0 if (alpha == cdzero || K == 0) { @@ -434,7 +424,7 @@ void zgemm_sme_NT(blasint M, blasint N, blasint K, const zdouble alpha, const zd } } -void zgemm_sme_TT(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) +static void zgemm_sme_TT(int M, int N, int K, const zdouble alpha, const zdouble *A, int lda, const zdouble *b, int ldb, const zdouble beta, zdouble *C, int ldc, bool conja, bool conjb) { if (alpha == cdzero || K == 0) { if (beta != cdone) { @@ -519,29 +509,32 @@ void zgemm_sme_TT(int M, int N, int K, const zdouble alpha, const zdouble *A, in } } -void sme_ZGEMM_KERNEL(const char *transa, const char *transb, const blasint *m, const blasint *n, const blasint *k, const zdouble alpha, const zdouble *a, const blasint *lda, const zdouble *b, const blasint *ldb, const zdouble beta, zdouble *c, const blasint *ldc) +void CNAME(const char *transa, const char *transb, const BLASLONG m, const BLASLONG n, const BLASLONG k, const double alpha_r, const double alpha_i, const double *a, const BLASLONG lda, const double *b, const BLASLONG ldb, const double beta_r, const double beta_i, double *c, const BLASLONG ldc) { +zdouble alpha={alpha_r,alpha_i}; +zdouble beta={beta_r,beta_i}; + bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); if (!trans_a && !trans_b) { bool conja = (*transa == 'R' || *transa == 'r'); bool conjb = (*transb == 'R' || *transb == 'r'); - zgemm_sme_NN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + zgemm_sme_NN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } else if (trans_a && !trans_b) { bool conja = (*transa == 'C' || *transa == 'c'); bool conjb = (*transb == 'R' || *transb == 'r'); - zgemm_sme_TN(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + zgemm_sme_TN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } else if (!trans_a && trans_b) { bool conja = (*transa == 'R' || *transa == 'r'); bool conjb = (*transb == 'C' || *transb == 'c'); - zgemm_sme_NT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + zgemm_sme_NT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } else { bool conja = (*transa == 'C' || *transa == 'c'); bool conjb = (*transb == 'C' || *transb == 'c'); - zgemm_sme_TT(*m, *n, *k, alpha, a, *lda, b, *ldb, beta, c, *ldc, conja, conjb); + zgemm_sme_TT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); } } From 81dd8597851da860227a99bfeba678c42fce055a Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 18:43:48 +0200 Subject: [PATCH 13/16] Move declarations of the ARM64 SME kernels to the appropriate headers --- interface/gemm.c | 48 ++++++++++++++++++------------------------------ 1 file changed, 18 insertions(+), 30 deletions(-) diff --git a/interface/gemm.c b/interface/gemm.c index 71bf12ac38..93a74959c6 100644 --- a/interface/gemm.c +++ b/interface/gemm.c @@ -45,12 +45,6 @@ #include "functable.h" #endif -#ifdef ARCH_ARM64 -void sme_SGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const float *alpha, const float *a, const BLASLONG *lda, const float *b, const BLASLONG *ldb, const float *beta, float *c, const BLASLONG *ldc); -void sme_DGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const double *alpha, const double *a, const BLASLONG *lda, const double *b, const BLASLONG *ldb, const double *beta, double *c, const BLASLONG *ldc); -void sme_CGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const float _Complex alpha, const float *a, const BLASLONG *lda, const float *b, const BLASLONG *ldb, const float _Complex beta, float *c, const BLASLONG *ldc); -void sme_ZGEMM_KERNEL(const char *ta, const char *tb, const BLASLONG *m, const BLASLONG *n, const BLASLONG *k, const double _Complex alpha, const double *a, const BLASLONG *lda, const double *b, const BLASLONG *ldb, const double _Complex beta, double *c, const BLASLONG *ldc); -#endif #ifndef COMPLEX #define SMP_THRESHOLD_MIN 65536.0 #ifdef XDOUBLE @@ -322,7 +316,7 @@ void NAME(char *TRANSA, char *TRANSB, args.alpha = (void *)alpha; args.beta = (void *)beta; - + transA = *TRANSA; transB = *TRANSB; @@ -572,8 +566,7 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 SGEMM_DIRECT(m, n, k, a, lda, b, ldb, c, ldc); return; } -else - if (order == CblasRowMajor && k==lda && n==ldb && n==ldc && TransA == CblasNoTrans && TransB == CblasNoTrans && SGEMM_DIRECT_PERFORMANT(m,n,k)) { +else if (order == CblasRowMajor && k==lda && n==ldb && n==ldc && TransA == CblasNoTrans && TransB == CblasNoTrans && SGEMM_DIRECT_PERFORMANT(m,n,k)) { SGEMM_DIRECT_ALPHA_BETA(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc); return; } @@ -595,52 +588,52 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 #endif //defined dynarch { char* TA,*TB; - if (transa & 1) + if (transa & 1) TA = "T"; - else + else TA= "N"; - if (transb & 1) + if (transb & 1) TB = "T"; - else + else TB= "N"; #ifndef COMPLEX - if (transa == 3) + if (transa == 3) TA= "T"; - if (transb == 3) + if (transb == 3) TB= "T"; FLOAT* al=(FLOAT*)args.alpha; FLOAT* be=(FLOAT*)args.beta; #ifndef DOUBLE - sme_SGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, al, args.a, &args.lda, args.b, &args.ldb, be, args.c, &args.ldc); + SME_SGEMM_KERNEL(TA,TB, args.m, args.n, args.k, al, args.a, args.lda, args.b, args.ldb, be, args.c, args.ldc); #else - sme_DGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, al, args.a, &args.lda, args.b, &args.ldb, be, args.c, &args.ldc); + SME_DGEMM_KERNEL(TA,TB, args.m, args.n, args.k, al, args.a, args.lda, args.b, args.ldb, be, args.c, args.ldc); #endif #else - if (transa == 2) + if (transa == 2) TA= "R"; - if (transb == 2) + if (transb == 2) TB= "R"; - if (transa == 3) + if (transa == 3) TA= "C"; - if (transb == 3) + if (transb == 3) TB= "C"; FLOAT* al=(FLOAT*)args.alpha; FLOAT* be=(FLOAT*)args.beta; #ifndef DOUBLE float _Complex c_al={al[0],al[1]}; float _Complex c_be={be[0],be[1]}; - sme_CGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, c_al, args.a, &args.lda, args.b, &args.ldb, c_be, args.c, &args.ldc); + SME_CGEMM_KERNEL(TA,TB, args.m, args.n, args.k, al[0],al[1], args.a, args.lda, args.b, args.ldb, be[0],be[1], args.c, args.ldc); #else double _Complex c_al={al[0],al[1]}; double _Complex c_be={be[0],be[1]}; - sme_ZGEMM_KERNEL(TA,TB, &args.m, &args.n, &args.k, c_al, args.a, &args.lda, args.b, &args.ldb, c_be, args.c, &args.ldc); + SME_ZGEMM_KERNEL(TA,TB, args.m, args.n, args.k, al[0],al[1], args.a, args.lda, args.b, args.ldb, be[0],be[1], args.c, args.ldc); #endif #endif return; } #endif //defined arm64 -#endif //defined b/hfloat16 +#endif //defined b/hfloat16 #if defined(__linux__) && defined(__x86_64__) && defined(BFLOAT16) #if defined(DYNAMIC_ARCH) @@ -742,6 +735,7 @@ double _Complex c_be={be[0],be[1]}; #if USE_SMALL_MATRIX_OPT #if !defined(COMPLEX) if(GEMM_SMALL_MATRIX_PERMIT(transa, transb, args.m, args.n, args.k, *(FLOAT *)(args.alpha), *(FLOAT *)(args.beta))){ + if(*(FLOAT *)(args.beta) == 0.0){ (GEMM_SMALL_KERNEL_B0((transb << 2) | transa))(args.m, args.n, args.k, args.a, args.lda, *(FLOAT *)(args.alpha), args.b, args.ldb, args.c, args.ldc); }else{ @@ -762,12 +756,6 @@ double _Complex c_be={be[0],be[1]}; #endif buffer = (XFLOAT *)blas_memory_alloc(0); - if (!buffer) { - info = -999; - BLASFUNC(xerbla)(ERROR_NAME, &info, sizeof(ERROR_NAME)); - return; - } - //For LOONGARCH64, applying an offset to the buffer is essential //for minimizing cache conflicts and optimizing performance. #if defined(ARCH_LOONGARCH64) && !defined(NO_AFFINITY) From f8830b66e36af655eed32b5afbb2d14ab9fde4a8 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 18:48:02 +0200 Subject: [PATCH 14/16] Add sme-f64f64 capability to VortexM4 and ARMV9SME build flags --- kernel/Makefile | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/kernel/Makefile b/kernel/Makefile index 9cab79c869..5256ef0c5d 100644 --- a/kernel/Makefile +++ b/kernel/Makefile @@ -27,7 +27,7 @@ endif ifdef TARGET_CORE ifeq ($(TARGET_CORE), ARMV9SME) - override CFLAGS += -DBUILD_KERNEL -DTABLE_NAME=gotoblas_$(TARGET_CORE) -march=armv9-a+sve2+sme + override CFLAGS += -DBUILD_KERNEL -DTABLE_NAME=gotoblas_$(TARGET_CORE) -march=armv9-a+sve2+sme+sme-f64f64 ifdef OS_WINDOWS ifeq ($(C_COMPILER), CLANG) override CFLAGS += --aarch64-stack-hazard-size=0 @@ -38,7 +38,7 @@ ifeq ($(TARGET_CORE), VORTEXM4) ifeq ($(C_COMPILER), GCC) override CFLAGS += -DBUILD_KERNEL -DTABLE_NAME=gotoblas_$(TARGET_CORE) -UHAVE_SME -march=armv8.4-a else - override CFLAGS += -DBUILD_KERNEL -DTABLE_NAME=gotoblas_$(TARGET_CORE) -march=armv8.4-a+sme + override CFLAGS += -DBUILD_KERNEL -DTABLE_NAME=gotoblas_$(TARGET_CORE) -march=armv8.4-a+sme+sme-f64f64 # ifneq ($(APPLECLANG),1) # override LDFLAGS += -lclang_rt_builtins-aarch64 # endif From 305bd67178e2d3e7d089dc091db18b861818a527 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 18:51:35 +0200 Subject: [PATCH 15/16] Add +sme-f64f64 to build flags of VortexM4 and ARMV9SME --- cmake/cc.cmake | 4 ++-- cmake/system.cmake | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/cmake/cc.cmake b/cmake/cc.cmake index 6cd5958eee..4d8735490a 100644 --- a/cmake/cc.cmake +++ b/cmake/cc.cmake @@ -315,7 +315,7 @@ if (${CORE} STREQUAL ARMV9SME) if (${CMAKE_C_COMPILER_ID} STREQUAL "NVHPC" AND NOT NO_SVE) set (CCOMMON_OPT "${CCOMMON_OPT} -tp=host") else () - set (CCOMMON_OPT "${CCOMMON_OPT} -march=armv9-a+sme") + set (CCOMMON_OPT "${CCOMMON_OPT} -march=armv9-a+sme+sme-f64f64") if (CMAKE_SYSTEM_NAME STREQUAL "Windows" AND CMAKE_C_COMPILER_ID MATCHES "Clang") set (CCOMMON_OPT "${CCOMMON_OPT} -mllvm --aarch64-stack-hazard-size=0") endif () @@ -329,7 +329,7 @@ if (${CORE} STREQUAL VORTEXM4) set (CCOMMON_OPT "${CCOMMON_OPT} -tp=host") else () if (${CMAKE_C_COMPILER_ID} STREQUAL "AppleClang") - set (CCOMMON_OPT "${CCOMMON_OPT} -march=armv8.4-a+sme -mcpu=apple-m4") + set (CCOMMON_OPT "${CCOMMON_OPT} -march=armv8.4-a+sme+sme-f64f64 -mcpu=apple-m4") else () set (CCOMMON_OPT "${CCOMMON_OPT} -march=armv8.4-a -mcpu=apple-m4") endif () diff --git a/cmake/system.cmake b/cmake/system.cmake index 720b32eb31..436f2109aa 100644 --- a/cmake/system.cmake +++ b/cmake/system.cmake @@ -369,13 +369,13 @@ if (${TARGET} STREQUAL NEOVERSEV1) endif() endif() if (${TARGET} STREQUAL ARMV9SME) - set (KERNEL_DEFINITIONS "${KERNEL_DEFINITIONS} -march=armv9-a+sme -O3") + set (KERNEL_DEFINITIONS "${KERNEL_DEFINITIONS} -march=armv9-a+sme+sme-f64f64 -O3") if (${CMAKE_SYSTEM_NAME} STREQUAL Windows AND ${CMAKE_C_COMPILER_ID} MATCHES "Clang") set (KERNEL_DEFINITIONS "${KERNEL_DEFINITIONS} -mllvm --aarch64-stack-hazard-size=0") endif() endif() if (${TARGET} STREQUAL VORTEXM4) - set (KERNEL_DEFINITIONS "${KERNEL_DEFINITIONS} -march=armv8.4-a+sme -O3") + set (KERNEL_DEFINITIONS "${KERNEL_DEFINITIONS} -march=armv8.4-a+sme+sme-f64f64 -O3") if (${CMAKE_SYSTEM_NAME} STREQUAL Windows AND ${CMAKE_C_COMPILER_ID} MATCHES "Clang") set (KERNEL_DEFINITIONS "${KERNEL_DEFINITIONS} -mllvm --aarch64-stack-hazard-size=0") endif() From 9884c480ea0f608ddc1d83182d874d9599716699 Mon Sep 17 00:00:00 2001 From: Martin Kroeker Date: Thu, 13 Aug 2026 22:20:23 +0200 Subject: [PATCH 16/16] Add casts to pacify homebrew-llvm --- kernel/arm64/sme_cgemm_kernel.c | 8 ++++---- kernel/arm64/sme_zgemm_kernel.c | 8 ++++---- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/kernel/arm64/sme_cgemm_kernel.c b/kernel/arm64/sme_cgemm_kernel.c index e613b1b343..72fa25e581 100644 --- a/kernel/arm64/sme_cgemm_kernel.c +++ b/kernel/arm64/sme_cgemm_kernel.c @@ -524,21 +524,21 @@ cfloat beta={beta_r,beta_i}; if (!trans_a && !trans_b) { bool conja=(*transa == 'R' || *transa == 'r'); bool conjb=(*transb == 'R' || *transb == 'r'); - cgemm_sme_NN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + cgemm_sme_NN(m, n, k, alpha, (const cfloat*) a, lda, (const cfloat*) b, ldb, beta, (cfloat*)c, ldc, conja, conjb); } else if (trans_a && !trans_b) { bool conja=(*transa == 'C' || *transa == 'c'); bool conjb=(*transb == 'R' || *transb == 'r'); - cgemm_sme_TN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + cgemm_sme_TN(m, n, k, alpha, (const cfloat*) a, lda, (const cfloat*) b, ldb, beta, (cfloat*)c, ldc, conja, conjb); } else if (!trans_a && trans_b) { bool conja=(*transa == 'R' || *transa == 'r'); bool conjb=(*transb == 'C' || *transb == 'c'); - cgemm_sme_NT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + cgemm_sme_NT(m, n, k, alpha, (const cfloat*) a, lda, (const cfloat*) b, ldb, beta, (cfloat*)c, ldc, conja, conjb); } else { bool conja=(*transa == 'C' || *transa == 'c'); bool conjb=(*transb == 'C' || *transb == 'c'); - cgemm_sme_TT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + cgemm_sme_TT(m, n, k, alpha, (const cfloat*) a, lda, (const cfloat*) b, ldb, beta, (cfloat*)c, ldc, conja, conjb); } } diff --git a/kernel/arm64/sme_zgemm_kernel.c b/kernel/arm64/sme_zgemm_kernel.c index 5e41010d8c..028e70b8e5 100644 --- a/kernel/arm64/sme_zgemm_kernel.c +++ b/kernel/arm64/sme_zgemm_kernel.c @@ -519,22 +519,22 @@ zdouble beta={beta_r,beta_i}; if (!trans_a && !trans_b) { bool conja = (*transa == 'R' || *transa == 'r'); bool conjb = (*transb == 'R' || *transb == 'r'); - zgemm_sme_NN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + zgemm_sme_NN(m, n, k, alpha, (const zdouble*)a, lda, (const zdouble*)b, ldb, beta, (zdouble*)c, ldc, conja, conjb); } else if (trans_a && !trans_b) { bool conja = (*transa == 'C' || *transa == 'c'); bool conjb = (*transb == 'R' || *transb == 'r'); - zgemm_sme_TN(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + zgemm_sme_TN(m, n, k, alpha, (const zdouble*)a, lda, (const zdouble*)b, ldb, beta, (zdouble*)c, ldc, conja, conjb); } else if (!trans_a && trans_b) { bool conja = (*transa == 'R' || *transa == 'r'); bool conjb = (*transb == 'C' || *transb == 'c'); - zgemm_sme_NT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + zgemm_sme_NT(m, n, k, alpha, (const zdouble*)a, lda, (const zdouble*)b, ldb, beta, (zdouble*)c, ldc, conja, conjb); } else { bool conja = (*transa == 'C' || *transa == 'c'); bool conjb = (*transb == 'C' || *transb == 'c'); - zgemm_sme_TT(m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, conja, conjb); + zgemm_sme_TT(m, n, k, alpha, (const zdouble*)a, lda, (const zdouble*)b, ldb, beta, (zdouble*)c, ldc, conja, conjb); } }