diff --git a/CONTRIBUTORS.md b/CONTRIBUTORS.md index b6a1884791..5618286fbe 100644 --- a/CONTRIBUTORS.md +++ b/CONTRIBUTORS.md @@ -288,3 +288,6 @@ hheei * Hugo Meiland * [2026-08-09] Add Cortex-A72 DGEMM 6x8 microkernel and blocking + +* Michael Tesch + * [2026-09-29] SME2 SGEMM/DGEMM kernel for Apple M4; SYMM, SYRK, SYR2K, TRMM and TRSM on the SME GEMM kernel diff --git a/README.md b/README.md index d651c42b36..d9d6ecc271 100644 --- a/README.md +++ b/README.md @@ -187,7 +187,7 @@ Please read `GotoBLAS_01Readme.txt` for older CPU models already supported by th - **Neoverse V3**: preliminary support - **Neoverse V3AE**: preliminary support - **Apple Vortex**: preliminary support based on ThunderX2/3 -- **Apple VortexM4**: preliminary support based on ThunderX2/3, SME kernels for SGEMM,SSYMM,STRMM,SSYRK,SSYR2K +- **Apple VortexM4**: preliminary support based on ThunderX2/3, SME kernels for SGEMM,SSYMM,STRMM,SSYRK,SSYR2K; on SME2 cores with a 512-bit streaming vector length (tested on M4) an SME2 kernel for SGEMM and DGEMM, which S/D SYMM, SYRK, SYR2K, TRMM and TRSM also use - **A64FX**: preliminary support, optimized Level-3 BLAS - **ARMV8SVE**: any ARMV8 cpu with SVE extensions - **ARMV9SME**: any ARMV9 cpu with SVE and SME extensions diff --git a/cpuid_arm64.c b/cpuid_arm64.c index 874af7800a..c11de5bf2d 100644 --- a/cpuid_arm64.c +++ b/cpuid_arm64.c @@ -408,7 +408,13 @@ int detect(void) if (value64 == 3660830781) return CPU_VORTEX; //A15/M2 if (value64 == 2271604202) return CPU_VORTEX; //A16/M3 if (value64 == 1867590060) return CPU_VORTEXM4; //M4 + if (value64 == 399882554) return CPU_VORTEXM4; //M4 Pro/Max if (value64 == 492472296) return CPU_VORTEXM4; //M5 + { /* later Apple cores with SME: same kernels until they get their own entry */ + int sme = 0; + size_t len = sizeof(sme); + if (sysctlbyname("hw.optional.arm.FEAT_SME", &sme, &len, NULL, 0) == 0 && sme) return CPU_VORTEXM4; + } #else #ifdef OS_WINDOWS HKEY reghandle; diff --git a/interface/gemm.c b/interface/gemm.c index 93a74959c6..1788706675 100644 --- a/interface/gemm.c +++ b/interface/gemm.c @@ -586,6 +586,12 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 #endif ) #endif //defined dynarch +/* below this the NEON paths win over the cost of entering streaming mode (Apple M4: about 20^3, 16^3 in fp64) */ +#if defined(DOUBLE) && !defined(COMPLEX) +if ((double)args.m * (double)args.n * (double)args.k >= 4000.) +#else +if ((double)args.m * (double)args.n * (double)args.k >= 8000.) +#endif { char* TA,*TB; if (transa & 1) diff --git a/interface/symm.c b/interface/symm.c index c1ab196080..ef00b44d2a 100644 --- a/interface/symm.c +++ b/interface/symm.c @@ -121,6 +121,10 @@ extern char* gotoblas_corename(void); #endif #endif +#if defined(ARCH_ARM64) +#include "../kernel/arm64/sme_level3.h" +#endif + static int (*symm[])(blas_arg_t *, BLASLONG *, BLASLONG *, FLOAT *, FLOAT *, BLASLONG) = { #ifndef GEMM3M #ifndef HEMM @@ -398,6 +402,10 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 #endif #endif +#endif + +#ifdef SME_LEVEL3 /* arm64 SME: recursive blocking on the SME GEMM kernel */ + if (s3_symm_hook(side, uplo, &args)) return; #endif IDEBUG_START; diff --git a/interface/syr2k.c b/interface/syr2k.c index 4b005d0880..35480c391d 100644 --- a/interface/syr2k.c +++ b/interface/syr2k.c @@ -71,6 +71,10 @@ #endif #endif +#if defined(ARCH_ARM64) +#include "../kernel/arm64/sme_level3.h" +#endif + static int (*syr2k[])(blas_arg_t *, BLASLONG *, BLASLONG *, FLOAT *, FLOAT *, BLASLONG) = { #ifndef HEMM SYR2K_UN, SYR2K_UC, SYR2K_LN, SYR2K_LC, @@ -388,6 +392,10 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 #endif +#ifdef SME_LEVEL3 /* arm64 SME: recursive blocking on the SME GEMM kernel */ + if (s3_syrk_hook(1, uplo, trans, &args)) return; +#endif + IDEBUG_START; FUNCTION_PROFILE_START(); diff --git a/interface/syrk.c b/interface/syrk.c index bfce536a06..c466ab53f4 100644 --- a/interface/syrk.c +++ b/interface/syrk.c @@ -77,6 +77,10 @@ #define GEMM_MULTITHREAD_THRESHOLD 4 #endif +#if defined(ARCH_ARM64) +#include "../kernel/arm64/sme_level3.h" +#endif + static int (*syrk[])(blas_arg_t *, BLASLONG *, BLASLONG *, FLOAT *, FLOAT *, BLASLONG) = { #ifndef HEMM SYRK_UN, SYRK_UC, SYRK_LN, SYRK_LC, @@ -371,6 +375,10 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 #endif +#ifdef SME_LEVEL3 /* arm64 SME: recursive blocking on the SME GEMM kernel */ + if (s3_syrk_hook(0, uplo, trans, &args)) return; +#endif + IDEBUG_START; FUNCTION_PROFILE_START(); diff --git a/interface/trsm.c b/interface/trsm.c index 1395e88f59..25dd2352be 100644 --- a/interface/trsm.c +++ b/interface/trsm.c @@ -91,6 +91,10 @@ extern char* gotoblas_corename(void); #endif +#if defined(ARCH_ARM64) +#include "../kernel/arm64/sme_level3.h" +#endif + static int (*trsm[])(blas_arg_t *, BLASLONG *, BLASLONG *, FLOAT *, FLOAT *, BLASLONG) = { #ifndef TRMM TRSM_LNUU, TRSM_LNUN, TRSM_LNLU, TRSM_LNLN, @@ -392,6 +396,10 @@ if (strcmp(gotoblas_corename(), "armv9sme") == 0 #endif if ((args.m == 0) || (args.n == 0)) return; +#ifdef SME_LEVEL3 /* arm64 SME: recursive blocking on the SME GEMM kernel */ + if (s3_trxm_hook(side, uplo, trans, unit, &args)) return; +#endif + IDEBUG_START; FUNCTION_PROFILE_START(); diff --git a/kernel/arm64/sgemm_direct_performant.c b/kernel/arm64/sgemm_direct_performant.c index 1e9dbf0c3b..f5f2c001d4 100644 --- a/kernel/arm64/sgemm_direct_performant.c +++ b/kernel/arm64/sgemm_direct_performant.c @@ -1,4 +1,10 @@ #include "common.h" +#if defined(__clang__) && defined(__ARM_FEATURE_SME) +#if (defined(__apple_build_version__) && __clang_major__ >= 17) || (!defined(__apple_build_version__) && __clang_major__ >= 18) +#define HAVE_SME2_GEMM 1 +#include "sme2_gemm_detect.h" +#endif +#endif /* helper for the direct sgemm code adapted from Arjan van der Ven's x86_64 version */ int CNAME(BLASLONG M, BLASLONG N, BLASLONG K) @@ -7,6 +13,11 @@ int CNAME(BLASLONG M, BLASLONG N, BLASLONG K) return 0; unsigned long long mnk = (unsigned long long)M * (unsigned long long)N * (unsigned long long)K; +#ifdef HAVE_SME2_GEMM + /* the SME2 kernel behind SME_SGEMM_KERNEL is faster from about 64^3 (Apple M4) */ + if (mnk >= 64ULL * 64ULL * 64ULL && s2_usable()) + return 0; +#endif /* benchmark performance on M4 peaks around 512 and crosses the graph of the NEON SGEMM at about 3100 */ if (mnk >= 3100ULL * 3100ULL * 3100ULL) return 0; diff --git a/kernel/arm64/sme2_gemm_detect.h b/kernel/arm64/sme2_gemm_detect.h new file mode 100644 index 0000000000..3e4d857cac --- /dev/null +++ b/kernel/arm64/sme2_gemm_detect.h @@ -0,0 +1,59 @@ +/*************************************************************************** +Copyright (c) 2026, 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. +*****************************************************************************/ + +/* Runtime check for the SME2 GEMM of sme2_gemm_impl.h; also used to keep the SME1 direct sgemm out of its way. */ +#ifndef SME2_GEMM_DETECT_H +#define SME2_GEMM_DETECT_H + +#include +#if defined(__APPLE__) +#include +#elif defined(__linux__) +#include +#endif + +/* SME2 with a 512-bit streaming vector length (the kernels assume 16 fp32 lanes). */ +static inline int s2_usable(void) { + static int ok = -1; + if (ok < 0) { + int sme2 = 0; +#if defined(__APPLE__) + int v = 0; + size_t len = sizeof(v); + sme2 = sysctlbyname("hw.optional.arm.FEAT_SME2", &v, &len, NULL, 0) == 0 && v; +#elif defined(__linux__) + sme2 = (getauxval(AT_HWCAP2) & (1UL << 37)) != 0; /* HWCAP2_SME2 */ +#endif + ok = sme2 && svcntsw() == 16; + } + return ok; +} + +/* The kernel keeps M, N and K in int; larger problems (INTERFACE64) stay on the other kernels. */ +#define S2_FITS(m, n, k) ((m) < (1L << 30) && (n) < (1L << 30) && (k) < (1L << 30)) + +#endif diff --git a/kernel/arm64/sme2_gemm_impl.h b/kernel/arm64/sme2_gemm_impl.h new file mode 100644 index 0000000000..0bcd73615a --- /dev/null +++ b/kernel/arm64/sme2_gemm_impl.h @@ -0,0 +1,1103 @@ +/*************************************************************************** +Copyright (c) 2026, 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. +*****************************************************************************/ + +/* + * SME2 real GEMM for a 512-bit streaming vector length, included by sme_sgemm_kernel.c and sme_dgemm_kernel.c. + * + * The design follows C. Deng, W. Yang, J. Fang, D. Dong, "Demystifying ARM SME to Optimize General Matrix + * Multiplications", arXiv:2512.21473, with the changes measured in MTGEMM-A (https://github.com/tesch1/mtgemm-a): + * - Goto loop nest on a row-major core problem; column-major C = op(A) op(B) is solved as C^T = op(B)^T op(A)^T. + * - mc, nc, kc from the paper's analytical model (TLB bound on kc, L2 bound, maximal compute-to-memory ratio). + * - A is packed into panels of VL rows with the transposition done in ZA (rows in, columns out). + * - B is packed "online": the first micro-kernel call of a block reads B from the source and writes the packed + * panel as a side effect, so there is no separate packing pass. + * - Micro-kernel of 1 x NT tiles (fp32: 16 x 64, fp64: 8 x 64) with SME2 multi-vector loads, a 2 x (NT/2) kernel + * for column tails of half a panel or more, and an edge kernel with all tiles along M. + * - Four-row ZA moves (MOVA vg4) for the C tiles and for the transposition of A, core-side L2 prefetch. + * - One thread per SME unit (Apple: per performance cluster) when the problem is large enough. + */ + +#include +#include +#include "sme2_gemm_detect.h" +#include +#include +#if defined(__APPLE__) +#include +#endif + +/* The file is built with the SME flags of its target (no SME2); the functions here enable SME2 themselves. */ +#pragma clang attribute push(__attribute__((target("sme2"))), apply_to = function) + +#define S2_S __arm_streaming +#define S2_ZA __arm_inout("za") +#define S2_INL static inline __attribute__((always_inline)) +#define S2_NOINL static __attribute__((noinline)) +#define S2_UNROLL _Pragma("GCC unroll 8") + +#ifdef DOUBLE +#define S2_VL 8 +#define S2_NT 8 +typedef svfloat64_t s2_v; +typedef svfloat64x4_t s2_v4; +#define S2_PT() svptrue_b64() +#define S2_PW(i, n) svwhilelt_b64_s64((int64_t)(i), (int64_t)(n)) +#define S2_CT() svptrue_c64() +#define S2_CW(i, n) svwhilelt_c64_s64((int64_t)(i), (int64_t)(n), 4) +#define S2_LD1(p, q) svld1_f64(p, q) +#define S2_ST1(p, q, v) svst1_f64(p, q, v) +#define S2_LD4(c, q) svld1_f64_x4(c, q) +#define S2_ST4(c, q, v) svst1_f64_x4(c, q, v) +#define S2_GET4(v, i) svget4_f64(v, i) +#define S2_CREATE4(a, b, c, d) svcreate4_f64(a, b, c, d) +#define S2_ZEROV() svdup_n_f64(0) +#define S2_MUL(v, a) svmul_n_f64_x(svptrue_b64(), v, a) +#define S2_MOPA(t, p, a, b) svmopa_za64_f64_m(t, p, p, a, b) +#define S2_RDH(t, s) svread_hor_za64_f64_m(svundef_f64(), svptrue_b64(), t, s) +#define S2_WRH(t, s, v) svwrite_hor_za64_f64_m(t, s, svptrue_b64(), v) +#define S2_RDV4(t, s) svread_ver_za64_f64_vg4(t, s) +#define S2_RDH4(t, s) svread_hor_za64_f64_vg4(t, s) +#define S2_WRH4(t, s, v) svwrite_hor_za64_f64_vg4(t, s, v) +#else +#define S2_VL 16 +#define S2_NT 4 +typedef svfloat32_t s2_v; +typedef svfloat32x4_t s2_v4; +#define S2_PT() svptrue_b32() +#define S2_PW(i, n) svwhilelt_b32_s64((int64_t)(i), (int64_t)(n)) +#define S2_CT() svptrue_c32() +#define S2_CW(i, n) svwhilelt_c32_s64((int64_t)(i), (int64_t)(n), 4) +#define S2_LD1(p, q) svld1_f32(p, q) +#define S2_ST1(p, q, v) svst1_f32(p, q, v) +#define S2_LD4(c, q) svld1_f32_x4(c, q) +#define S2_ST4(c, q, v) svst1_f32_x4(c, q, v) +#define S2_GET4(v, i) svget4_f32(v, i) +#define S2_CREATE4(a, b, c, d) svcreate4_f32(a, b, c, d) +#define S2_ZEROV() svdup_n_f32(0) +#define S2_MUL(v, a) svmul_n_f32_x(svptrue_b32(), v, a) +#define S2_MOPA(t, p, a, b) svmopa_za32_f32_m(t, p, p, a, b) +#define S2_RDH(t, s) svread_hor_za32_f32_m(svundef_f32(), svptrue_b32(), t, s) +#define S2_WRH(t, s, v) svwrite_hor_za32_f32_m(t, s, svptrue_b32(), v) +#define S2_RDV4(t, s) svread_ver_za32_f32_vg4(t, s) +#define S2_RDH4(t, s) svread_hor_za32_f32_vg4(t, s) +#define S2_WRH4(t, s, v) svwrite_hor_za32_f32_vg4(t, s, v) +#endif + +#define S2_CH (S2_NT * S2_VL) /* columns of A moved through ZA per transposition pass */ +#define S2_NR (S2_NT * S2_VL) /* width of the main B panel */ +#define S2_TN2 (S2_NT / 2) /* tiles along N of the half-width tail kernel */ +#define S2_W2 (S2_TN2 * S2_VL) +#define S2_PF_ROWS_B 8 /* B rows ahead for the core prefetch of B strips */ + +/* Blocking model parameters (Apple M4: SME load bandwidth holds up to 8 MB, 16 KB pages). */ +#ifndef S2_L2_BYTES +#define S2_L2_BYTES (8L << 20) +#endif +#ifndef S2_PAGE_BYTES +#define S2_PAGE_BYTES 16384L +#endif +#ifndef S2_TLB_ENTRIES +#define S2_TLB_ENTRIES 160 +#endif + +enum { S2_ZERO = 0, S2_LOAD = 1, S2_SCALE = 2 }; + +/* Tile numbers of the ZA intrinsics must be literals; with a constant t these switches fold away when inlined. */ +S2_INL void s2_mopa(int t, svbool_t p, s2_v a, s2_v b) S2_S S2_ZA { + switch (t) { + case 0: S2_MOPA(0, p, a, b); break; + case 1: S2_MOPA(1, p, a, b); break; + case 2: S2_MOPA(2, p, a, b); break; + case 3: S2_MOPA(3, p, a, b); break; +#ifdef DOUBLE + case 4: S2_MOPA(4, p, a, b); break; + case 5: S2_MOPA(5, p, a, b); break; + case 6: S2_MOPA(6, p, a, b); break; + case 7: S2_MOPA(7, p, a, b); break; +#endif + } +} +S2_INL s2_v s2_rdh(int t, uint32_t s) S2_S S2_ZA { + switch (t) { + case 0: return S2_RDH(0, s); + case 1: return S2_RDH(1, s); + case 2: return S2_RDH(2, s); +#ifdef DOUBLE + case 3: return S2_RDH(3, s); + case 4: return S2_RDH(4, s); + case 5: return S2_RDH(5, s); + case 6: return S2_RDH(6, s); + default: return S2_RDH(7, s); +#else + default: return S2_RDH(3, s); +#endif + } +} +S2_INL void s2_wrh(int t, uint32_t s, s2_v v) S2_S S2_ZA { + switch (t) { + case 0: S2_WRH(0, s, v); break; + case 1: S2_WRH(1, s, v); break; + case 2: S2_WRH(2, s, v); break; + case 3: S2_WRH(3, s, v); break; +#ifdef DOUBLE + case 4: S2_WRH(4, s, v); break; + case 5: S2_WRH(5, s, v); break; + case 6: S2_WRH(6, s, v); break; + case 7: S2_WRH(7, s, v); break; +#endif + } +} +S2_INL s2_v4 s2_rdv4(int t, uint32_t s) S2_S S2_ZA { + switch (t) { + case 0: return S2_RDV4(0, s); + case 1: return S2_RDV4(1, s); + case 2: return S2_RDV4(2, s); +#ifdef DOUBLE + case 3: return S2_RDV4(3, s); + case 4: return S2_RDV4(4, s); + case 5: return S2_RDV4(5, s); + case 6: return S2_RDV4(6, s); + default: return S2_RDV4(7, s); +#else + default: return S2_RDV4(3, s); +#endif + } +} +/* Four consecutive horizontal slices of tile t at once (MOVA vg4). */ +S2_INL s2_v4 s2_rdh4(int t, uint32_t s) S2_S S2_ZA { + switch (t) { + case 0: return S2_RDH4(0, s); + case 1: return S2_RDH4(1, s); + case 2: return S2_RDH4(2, s); +#ifdef DOUBLE + case 3: return S2_RDH4(3, s); + case 4: return S2_RDH4(4, s); + case 5: return S2_RDH4(5, s); + case 6: return S2_RDH4(6, s); + default: return S2_RDH4(7, s); +#else + default: return S2_RDH4(3, s); +#endif + } +} +S2_INL void s2_wrh4(int t, uint32_t s, s2_v4 v) S2_S S2_ZA { + switch (t) { + case 0: S2_WRH4(0, s, v); break; + case 1: S2_WRH4(1, s, v); break; + case 2: S2_WRH4(2, s, v); break; + case 3: S2_WRH4(3, s, v); break; +#ifdef DOUBLE + case 4: S2_WRH4(4, s, v); break; + case 5: S2_WRH4(5, s, v); break; + case 6: S2_WRH4(6, s, v); break; + case 7: S2_WRH4(7, s, v); break; +#endif + } +} +S2_INL s2_v s2_get4(s2_v4 v, int i) S2_S { + switch (i) { + case 0: return S2_GET4(v, 0); + case 1: return S2_GET4(v, 1); + case 2: return S2_GET4(v, 2); + default: return S2_GET4(v, 3); + } +} +S2_INL s2_v4 s2_scale4(s2_v4 v, FLOAT a) S2_S { + return S2_CREATE4(S2_MUL(S2_GET4(v, 0), a), S2_MUL(S2_GET4(v, 1), a), S2_MUL(S2_GET4(v, 2), a), S2_MUL(S2_GET4(v, 3), a)); +} +/* Vector i of the pair (x0, x1) of four-vector groups. */ +S2_INL s2_v s2_pick(s2_v4 x0, s2_v4 x1, int i) S2_S { return i < 4 ? s2_get4(x0, i) : s2_get4(x1, i - 4); } + +S2_INL int64_t s2_clip(int64_t n, int64_t hi) __arm_streaming_compatible { return n < 0 ? 0 : (n > hi ? hi : n); } +S2_INL s2_v4 s2_ld4(const FLOAT *p) S2_S { return S2_LD4(S2_CT(), p); } +S2_INL s2_v4 s2_ld4n(const FLOAT *p, int64_t n) S2_S { return S2_LD4(S2_CW(0, n), p); } +S2_INL void s2_st4(FLOAT *p, s2_v4 v) S2_S { S2_ST4(S2_CT(), p, v); } +S2_INL void s2_st4n(FLOAT *p, s2_v4 v, int64_t n) S2_S { S2_ST4(S2_CW(0, n), p, v); } + +/* Core-side prefetch into L2 (prfm pldl2keep); the SME unit reads through L2. */ +S2_INL void s2_pf_l2(const void *p, int bytes) __arm_streaming_compatible { + for (int l = 0; l < bytes; l += 128) __builtin_prefetch((const char *)p + l, 0, 2); +} + +/* Four rows x 64 columns into horizontal slices r..r+3 of every tile: strided-register x4 loads put the four rows + of one tile into z(4t)..z(4t+3), so one MOVA vg4 per tile suffices. */ +S2_INL void s2_rows_in4(const FLOAT *src, BLASLONG ld, uint32_t r) S2_S S2_ZA { + const FLOAT *p1 = src + ld, *p2 = src + 2 * ld, *p3 = src + 3 * ld; +#ifdef DOUBLE + __asm__ volatile( + "ptrue pn8.d\n" + "ld1d {z16.d, z20.d, z24.d, z28.d}, pn8/z, [%[a0]]\n" + "ld1d {z17.d, z21.d, z25.d, z29.d}, pn8/z, [%[a1]]\n" + "ld1d {z18.d, z22.d, z26.d, z30.d}, pn8/z, [%[a2]]\n" + "ld1d {z19.d, z23.d, z27.d, z31.d}, pn8/z, [%[a3]]\n" + "mova za0h.d[%w[r], 0:3], {z16.d - z19.d}\n" + "mova za1h.d[%w[r], 0:3], {z20.d - z23.d}\n" + "mova za2h.d[%w[r], 0:3], {z24.d - z27.d}\n" + "mova za3h.d[%w[r], 0:3], {z28.d - z31.d}\n" + "ld1d {z16.d, z20.d, z24.d, z28.d}, pn8/z, [%[a0], #4, mul vl]\n" + "ld1d {z17.d, z21.d, z25.d, z29.d}, pn8/z, [%[a1], #4, mul vl]\n" + "ld1d {z18.d, z22.d, z26.d, z30.d}, pn8/z, [%[a2], #4, mul vl]\n" + "ld1d {z19.d, z23.d, z27.d, z31.d}, pn8/z, [%[a3], #4, mul vl]\n" + "mova za4h.d[%w[r], 0:3], {z16.d - z19.d}\n" + "mova za5h.d[%w[r], 0:3], {z20.d - z23.d}\n" + "mova za6h.d[%w[r], 0:3], {z24.d - z27.d}\n" + "mova za7h.d[%w[r], 0:3], {z28.d - z31.d}\n" + : + : [a0] "r"(src), [a1] "r"(p1), [a2] "r"(p2), [a3] "r"(p3), [r] "Ucj"(r) + : "p8", "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", "z24", "z25", "z26", "z27", "z28", "z29", + "z30", "z31", "memory"); +#else + __asm__ volatile( + "ptrue pn8.s\n" + "ld1w {z16.s, z20.s, z24.s, z28.s}, pn8/z, [%[a0]]\n" + "ld1w {z17.s, z21.s, z25.s, z29.s}, pn8/z, [%[a1]]\n" + "ld1w {z18.s, z22.s, z26.s, z30.s}, pn8/z, [%[a2]]\n" + "ld1w {z19.s, z23.s, z27.s, z31.s}, pn8/z, [%[a3]]\n" + "mova za0h.s[%w[r], 0:3], {z16.s - z19.s}\n" + "mova za1h.s[%w[r], 0:3], {z20.s - z23.s}\n" + "mova za2h.s[%w[r], 0:3], {z24.s - z27.s}\n" + "mova za3h.s[%w[r], 0:3], {z28.s - z31.s}\n" + : + : [a0] "r"(src), [a1] "r"(p1), [a2] "r"(p2), [a3] "r"(p3), [r] "Ucj"(r) + : "p8", "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", "z24", "z25", "z26", "z27", "z28", "z29", + "z30", "z31", "memory"); +#endif +} + +/* Inverse of s2_rows_in4: horizontal slices r..r+3 of every tile to four rows of 64 columns. */ +S2_INL void s2_rows_out4(FLOAT *dst, BLASLONG ld, uint32_t r) S2_S S2_ZA { + FLOAT *p1 = dst + ld, *p2 = dst + 2 * ld, *p3 = dst + 3 * ld; +#ifdef DOUBLE + __asm__ volatile( + "ptrue pn8.d\n" + "mova {z16.d - z19.d}, za0h.d[%w[r], 0:3]\n" + "mova {z20.d - z23.d}, za1h.d[%w[r], 0:3]\n" + "mova {z24.d - z27.d}, za2h.d[%w[r], 0:3]\n" + "mova {z28.d - z31.d}, za3h.d[%w[r], 0:3]\n" + "st1d {z16.d, z20.d, z24.d, z28.d}, pn8, [%[a0]]\n" + "st1d {z17.d, z21.d, z25.d, z29.d}, pn8, [%[a1]]\n" + "st1d {z18.d, z22.d, z26.d, z30.d}, pn8, [%[a2]]\n" + "st1d {z19.d, z23.d, z27.d, z31.d}, pn8, [%[a3]]\n" + "mova {z16.d - z19.d}, za4h.d[%w[r], 0:3]\n" + "mova {z20.d - z23.d}, za5h.d[%w[r], 0:3]\n" + "mova {z24.d - z27.d}, za6h.d[%w[r], 0:3]\n" + "mova {z28.d - z31.d}, za7h.d[%w[r], 0:3]\n" + "st1d {z16.d, z20.d, z24.d, z28.d}, pn8, [%[a0], #4, mul vl]\n" + "st1d {z17.d, z21.d, z25.d, z29.d}, pn8, [%[a1], #4, mul vl]\n" + "st1d {z18.d, z22.d, z26.d, z30.d}, pn8, [%[a2], #4, mul vl]\n" + "st1d {z19.d, z23.d, z27.d, z31.d}, pn8, [%[a3], #4, mul vl]\n" + : + : [a0] "r"(dst), [a1] "r"(p1), [a2] "r"(p2), [a3] "r"(p3), [r] "Ucj"(r) + : "p8", "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", "z24", "z25", "z26", "z27", "z28", "z29", + "z30", "z31", "memory"); +#else + __asm__ volatile( + "ptrue pn8.s\n" + "mova {z16.s - z19.s}, za0h.s[%w[r], 0:3]\n" + "mova {z20.s - z23.s}, za1h.s[%w[r], 0:3]\n" + "mova {z24.s - z27.s}, za2h.s[%w[r], 0:3]\n" + "mova {z28.s - z31.s}, za3h.s[%w[r], 0:3]\n" + "st1w {z16.s, z20.s, z24.s, z28.s}, pn8, [%[a0]]\n" + "st1w {z17.s, z21.s, z25.s, z29.s}, pn8, [%[a1]]\n" + "st1w {z18.s, z22.s, z26.s, z30.s}, pn8, [%[a2]]\n" + "st1w {z19.s, z23.s, z27.s, z31.s}, pn8, [%[a3]]\n" + : + : [a0] "r"(dst), [a1] "r"(p1), [a2] "r"(p2), [a3] "r"(p3), [r] "Ucj"(r) + : "p8", "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", "z24", "z25", "z26", "z27", "z28", "z29", + "z30", "z31", "memory"); +#endif +} + +#ifndef DOUBLE +/* 2 x 2 fp32 tiles (h = tile row): four C rows of 32 columns via strided x2 loads, one MOVA vg4 per tile. */ +#define S2_X2_ARGS \ + : : [a0] "r"(src), [a1] "r"(p1), [a2] "r"(p2), [a3] "r"(p3), [r] "Ucj"(r) \ + : "p8", "z16", "z17", "z18", "z19", "z24", "z25", "z26", "z27", "memory" +#define S2_X2_LD \ + "ptrue pn8.s\n ld1w {z16.s, z24.s}, pn8/z, [%[a0]]\n ld1w {z17.s, z25.s}, pn8/z, [%[a1]]\n" \ + "ld1w {z18.s, z26.s}, pn8/z, [%[a2]]\n ld1w {z19.s, z27.s}, pn8/z, [%[a3]]\n" +#define S2_X2_ST \ + "st1w {z16.s, z24.s}, pn8, [%[a0]]\n st1w {z17.s, z25.s}, pn8, [%[a1]]\n" \ + "st1w {z18.s, z26.s}, pn8, [%[a2]]\n st1w {z19.s, z27.s}, pn8, [%[a3]]\n" +S2_INL void s2_rows_in4_x2(const float *src, BLASLONG ld, uint32_t r, int h) S2_S S2_ZA { + const float *p1 = src + ld, *p2 = src + 2 * ld, *p3 = src + 3 * ld; + if (h == 0) + __asm__ volatile(S2_X2_LD "mova za0h.s[%w[r], 0:3], {z16.s - z19.s}\n mova za1h.s[%w[r], 0:3], {z24.s - z27.s}\n" + S2_X2_ARGS); + else + __asm__ volatile(S2_X2_LD "mova za2h.s[%w[r], 0:3], {z16.s - z19.s}\n mova za3h.s[%w[r], 0:3], {z24.s - z27.s}\n" + S2_X2_ARGS); +} +S2_INL void s2_rows_out4_x2(float *dst, BLASLONG ld, uint32_t r, int h) S2_S S2_ZA { + const float *src = dst; + float *p1 = dst + ld, *p2 = dst + 2 * ld, *p3 = dst + 3 * ld; + if (h == 0) + __asm__ volatile("ptrue pn8.s\n mova {z16.s - z19.s}, za0h.s[%w[r], 0:3]\n mova {z24.s - z27.s}, za1h.s[%w[r], 0:3]\n" + S2_X2_ST S2_X2_ARGS); + else + __asm__ volatile("ptrue pn8.s\n mova {z16.s - z19.s}, za2h.s[%w[r], 0:3]\n mova {z24.s - z27.s}, za3h.s[%w[r], 0:3]\n" + S2_X2_ST S2_X2_ARGS); +} +#endif + +#ifdef DOUBLE +/* fp64 half-width block (h = tile row): four C rows of 32 columns into tiles 4h..4h+3, one MOVA vg4 per tile. */ +S2_INL void s2_rows_in4_h(const double *src, BLASLONG ld, uint32_t r, int h) S2_S S2_ZA { + const double *p1 = src + ld, *p2 = src + 2 * ld, *p3 = src + 3 * ld; +#define S2_H_LD \ + "ptrue pn8.d\n ld1d {z16.d, z20.d, z24.d, z28.d}, pn8/z, [%[a0]]\n" \ + "ld1d {z17.d, z21.d, z25.d, z29.d}, pn8/z, [%[a1]]\n ld1d {z18.d, z22.d, z26.d, z30.d}, pn8/z, [%[a2]]\n" \ + "ld1d {z19.d, z23.d, z27.d, z31.d}, pn8/z, [%[a3]]\n" +#define S2_H_ARGS \ + : : [a0] "r"(src), [a1] "r"(p1), [a2] "r"(p2), [a3] "r"(p3), [r] "Ucj"(r) \ + : "p8", "z16", "z17", "z18", "z19", "z20", "z21", "z22", "z23", "z24", "z25", "z26", "z27", "z28", "z29", \ + "z30", "z31", "memory" + if (h == 0) + __asm__ volatile(S2_H_LD "mova za0h.d[%w[r], 0:3], {z16.d - z19.d}\n mova za1h.d[%w[r], 0:3], {z20.d - z23.d}\n" + "mova za2h.d[%w[r], 0:3], {z24.d - z27.d}\n mova za3h.d[%w[r], 0:3], {z28.d - z31.d}\n" S2_H_ARGS); + else + __asm__ volatile(S2_H_LD "mova za4h.d[%w[r], 0:3], {z16.d - z19.d}\n mova za5h.d[%w[r], 0:3], {z20.d - z23.d}\n" + "mova za6h.d[%w[r], 0:3], {z24.d - z27.d}\n mova za7h.d[%w[r], 0:3], {z28.d - z31.d}\n" S2_H_ARGS); +#undef S2_H_LD +#undef S2_H_ARGS +} +#endif + +/* C block of TM x TN tiles into ZA: rows < mrows, columns < ncols; tile (p, q) holds rows p*VL.., columns q*VL.. */ +S2_INL void s2_c_load(const int TM, const int TN, FLOAT *C, BLASLONG ldc, int mrows, int ncols, int cmode, + FLOAT beta) S2_S S2_ZA { + if (cmode == S2_ZERO) { /* rows and columns outside the block hold stale values but are never stored */ + svzero_za(); + return; + } + int s0 = 0; + if (TM == 1 && TN == S2_NT && cmode == S2_LOAD && ncols == TN * S2_VL) + for (; s0 + 4 <= mrows; s0 += 4) s2_rows_in4(C + (BLASLONG)s0 * ldc, ldc, s0); +#ifndef DOUBLE + if (TM == 2 && TN == 2 && cmode == S2_LOAD && ncols == 2 * S2_VL && mrows == 2 * S2_VL) { + for (uint32_t r = 0; r < S2_VL; r += 4) s2_rows_in4_x2(C + (BLASLONG)r * ldc, ldc, r, 0); + for (uint32_t r = 0; r < S2_VL; r += 4) s2_rows_in4_x2(C + (BLASLONG)(S2_VL + r) * ldc, ldc, r, 1); + return; + } +#else + if (TN == 4 && cmode == S2_LOAD && ncols == 4 * S2_VL && mrows == TM * S2_VL) { + S2_UNROLL for (int P = 0; P < TM; ++P) + for (uint32_t r = 0; r < S2_VL; r += 4) s2_rows_in4_h(C + (BLASLONG)(P * S2_VL + r) * ldc, ldc, r, P); + return; + } +#endif + const int sc = cmode == S2_SCALE; + S2_UNROLL for (int P = 0; P < TM; ++P) { + const int prow = mrows - P * S2_VL < S2_VL ? mrows - P * S2_VL : S2_VL; + int s = P == 0 ? s0 : 0; + for (; s < prow; ++s) { + FLOAT *row = C + (BLASLONG)(P * S2_VL + s) * ldc; + if (TN % 4 == 0) { + S2_UNROLL for (int G = 0; G < TN / 4; ++G) { + s2_v4 v = ncols >= (G + 1) * 4 * S2_VL ? s2_ld4(row + G * 4 * S2_VL) + : s2_ld4n(row + G * 4 * S2_VL, s2_clip(ncols - G * 4 * S2_VL, 4 * S2_VL)); + S2_UNROLL for (int Q = 0; Q < 4; ++Q) { + s2_v x = s2_get4(v, Q); + if (sc) x = S2_MUL(x, beta); + s2_wrh(P * TN + G * 4 + Q, s, x); + } + } + } else { + S2_UNROLL for (int Q = 0; Q < TN; ++Q) { + s2_v x = S2_LD1(S2_PW(Q * S2_VL, ncols), row + Q * S2_VL); + if (sc) x = S2_MUL(x, beta); + s2_wrh(P * TN + Q, s, x); + } + } + } + } +} + +S2_INL void s2_c_store(const int TM, const int TN, FLOAT *C, BLASLONG ldc, int mrows, int ncols) S2_S S2_ZA { + int s0 = 0; + if (TM == 1 && TN == S2_NT && ncols == TN * S2_VL) + for (; s0 + 4 <= mrows; s0 += 4) s2_rows_out4(C + (BLASLONG)s0 * ldc, ldc, s0); +#ifndef DOUBLE + if (TM == 2 && TN == 2 && ncols == 2 * S2_VL && mrows == 2 * S2_VL) { + for (uint32_t r = 0; r < S2_VL; r += 4) s2_rows_out4_x2(C + (BLASLONG)r * ldc, ldc, r, 0); + for (uint32_t r = 0; r < S2_VL; r += 4) s2_rows_out4_x2(C + (BLASLONG)(S2_VL + r) * ldc, ldc, r, 1); + return; + } +#endif + S2_UNROLL for (int P = 0; P < TM; ++P) { + const int prow = mrows - P * S2_VL < S2_VL ? mrows - P * S2_VL : S2_VL; + int s = P == 0 ? s0 : 0; + for (; s + 4 <= prow; s += 4) { /* four rows per ZA move, stores predicated for partial widths */ + FLOAT *row = C + (BLASLONG)(P * S2_VL + s) * ldc; + if (TN % 4 == 0) { + S2_UNROLL for (int G = 0; G < TN / 4; ++G) { + const int t = P * TN + G * 4; + const s2_v4 q0 = s2_rdh4(t, s), q1 = s2_rdh4(t + 1, s), q2 = s2_rdh4(t + 2, s), q3 = s2_rdh4(t + 3, s); + S2_UNROLL for (int u = 0; u < 4; ++u) { + const s2_v4 v = S2_CREATE4(s2_get4(q0, u), s2_get4(q1, u), s2_get4(q2, u), s2_get4(q3, u)); + if (ncols >= (G + 1) * 4 * S2_VL) s2_st4(row + u * ldc + G * 4 * S2_VL, v); + else s2_st4n(row + u * ldc + G * 4 * S2_VL, v, s2_clip(ncols - G * 4 * S2_VL, 4 * S2_VL)); + } + } + } else { + S2_UNROLL for (int Q = 0; Q < TN; ++Q) { + const svbool_t pg = S2_PW(Q * S2_VL, ncols); + const s2_v4 q = s2_rdh4(P * TN + Q, s); + S2_UNROLL for (int u = 0; u < 4; ++u) S2_ST1(pg, row + u * ldc + Q * S2_VL, s2_get4(q, u)); + } + } + } + for (; s < prow; ++s) { + FLOAT *row = C + (BLASLONG)(P * S2_VL + s) * ldc; + if (TN % 4 == 0) { + S2_UNROLL for (int G = 0; G < TN / 4; ++G) { + const int t = P * TN + G * 4; + s2_v4 v = S2_CREATE4(s2_rdh(t, s), s2_rdh(t + 1, s), s2_rdh(t + 2, s), s2_rdh(t + 3, s)); + if (ncols >= (G + 1) * 4 * S2_VL) s2_st4(row + G * 4 * S2_VL, v); + else s2_st4n(row + G * 4 * S2_VL, v, s2_clip(ncols - G * 4 * S2_VL, 4 * S2_VL)); + } + } else { + S2_UNROLL for (int Q = 0; Q < TN; ++Q) S2_ST1(S2_PW(Q * S2_VL, ncols), row + Q * S2_VL, s2_rdh(P * TN + Q, s)); + } + } + } +} + +/* A row of B for TN >= 4 (up to eight vectors): from the packed panel, or from the source and written to it. */ +S2_INL void s2_b_row(const int TN, const int ONLINE, const FLOAT *src, FLOAT *dst, int ncols, s2_v4 *b0, + s2_v4 *b1) S2_S { + if (ONLINE) { + *b0 = ncols >= 4 * S2_VL ? s2_ld4(src) : s2_ld4n(src, s2_clip(ncols, 4 * S2_VL)); + s2_st4(dst, *b0); + if (TN == 8) { + *b1 = ncols >= 8 * S2_VL ? s2_ld4(src + 4 * S2_VL) : s2_ld4n(src + 4 * S2_VL, s2_clip(ncols - 4 * S2_VL, 4 * S2_VL)); + s2_st4(dst + 4 * S2_VL, *b1); + } else { + *b1 = *b0; + } + } else { + *b0 = s2_ld4(dst); + *b1 = TN == 8 ? s2_ld4(dst + 4 * S2_VL) : *b0; + } +} + +/* + * Micro-kernel (paper Alg. 1): TM x TN tiles. Ar holds TM packed panels of VL rows (stride astr), Br is the + * kb x (TN*VL) packed panel of B. ONLINE: B comes from the source Bs and is written to Br as a side effect. + */ +S2_INL void s2_kernel(const int TM, const int TN, const int ONLINE, int kb, const FLOAT *Ar, BLASLONG astr, FLOAT *Br, + const FLOAT *Bs, BLASLONG ldb, int ncols, FLOAT *C, BLASLONG ldc, int mrows, int cmode, + FLOAT beta, int pfb) S2_S S2_ZA { + const int NR = TN * S2_VL; + s2_c_load(TM, TN, C, ldc, mrows, ncols, cmode, beta); + const svbool_t pt = S2_PT(); + int k = 0; + if (TN >= 4) { + for (; k + 4 <= kb; k += 4) { + if (ONLINE && pfb && k + S2_PF_ROWS_B + 4 <= kb) + for (int u = 0; u < 4; ++u) s2_pf_l2(Bs + (BLASLONG)(k + S2_PF_ROWS_B + u) * ldb, ncols * (int)sizeof(FLOAT)); + const s2_v4 a0 = s2_ld4(Ar + (BLASLONG)k * S2_VL); + const s2_v4 a1 = TM > 1 ? s2_ld4(Ar + astr + (BLASLONG)k * S2_VL) : a0; + S2_UNROLL for (int U = 0; U < 4; ++U) { + s2_v4 b0, b1; + s2_b_row(TN, ONLINE, Bs + (BLASLONG)(k + U) * ldb, Br + (BLASLONG)(k + U) * NR, ncols, &b0, &b1); + S2_UNROLL for (int P = 0; P < TM; ++P) { + const s2_v a = s2_get4(P == 0 ? a0 : a1, U); + S2_UNROLL for (int Q = 0; Q < TN; ++Q) s2_mopa(P * TN + Q, pt, a, s2_pick(b0, b1, Q)); + } + } + } + for (; k < kb; ++k) { + s2_v4 b0, b1; + s2_b_row(TN, ONLINE, Bs + (BLASLONG)k * ldb, Br + (BLASLONG)k * NR, ncols, &b0, &b1); + S2_UNROLL for (int P = 0; P < TM; ++P) { + const s2_v a = S2_LD1(pt, Ar + P * astr + (BLASLONG)k * S2_VL); + S2_UNROLL for (int Q = 0; Q < TN; ++Q) s2_mopa(P * TN + Q, pt, a, s2_pick(b0, b1, Q)); + } + } + } else { + for (; k + 4 <= kb; k += 4) { + if (ONLINE && pfb && k + S2_PF_ROWS_B + 4 <= kb) + for (int u = 0; u < 4; ++u) s2_pf_l2(Bs + (BLASLONG)(k + S2_PF_ROWS_B + u) * ldb, ncols * (int)sizeof(FLOAT)); + s2_v4 bb0, bb1; + FLOAT *dst = Br + (BLASLONG)k * NR; + if (ONLINE) { + const svbool_t p0 = S2_PW(0, ncols), p1 = S2_PW(S2_VL, ncols); + const FLOAT *s = Bs + (BLASLONG)k * ldb; + if (TN == 1) { + bb0 = S2_CREATE4(S2_LD1(p0, s), S2_LD1(p0, s + ldb), S2_LD1(p0, s + 2 * ldb), S2_LD1(p0, s + 3 * ldb)); + bb1 = bb0; + } else { + bb0 = S2_CREATE4(S2_LD1(p0, s), S2_LD1(p1, s + S2_VL), S2_LD1(p0, s + ldb), S2_LD1(p1, s + ldb + S2_VL)); + bb1 = S2_CREATE4(S2_LD1(p0, s + 2 * ldb), S2_LD1(p1, s + 2 * ldb + S2_VL), S2_LD1(p0, s + 3 * ldb), + S2_LD1(p1, s + 3 * ldb + S2_VL)); + } + s2_st4(dst, bb0); + if (TN == 2) s2_st4(dst + 4 * S2_VL, bb1); + } else { + bb0 = s2_ld4(dst); + bb1 = TN == 2 ? s2_ld4(dst + 4 * S2_VL) : bb0; + } + /* Up to four A panels at a time, FMOPAs ordered so that consecutive ones hit different tiles. */ + S2_UNROLL for (int H = 0; H < (TM + 3) / 4; ++H) { + const int P0 = H * 4, NP = TM - P0 < 4 ? TM - P0 : 4; + const s2_v4 a0 = s2_ld4(Ar + P0 * astr + (BLASLONG)k * S2_VL); + const s2_v4 a1 = NP > 1 ? s2_ld4(Ar + (P0 + 1) * astr + (BLASLONG)k * S2_VL) : a0; + const s2_v4 a2 = NP > 2 ? s2_ld4(Ar + (P0 + 2) * astr + (BLASLONG)k * S2_VL) : a0; + const s2_v4 a3 = NP > 3 ? s2_ld4(Ar + (P0 + 3) * astr + (BLASLONG)k * S2_VL) : a0; + S2_UNROLL for (int U = 0; U < 4; ++U) { + S2_UNROLL for (int PP = 0; PP < NP; ++PP) { + const s2_v a = s2_get4(PP == 0 ? a0 : PP == 1 ? a1 : PP == 2 ? a2 : a3, U); + S2_UNROLL for (int Q = 0; Q < TN; ++Q) s2_mopa((P0 + PP) * TN + Q, pt, a, s2_pick(bb0, bb1, U * TN + Q)); + } + } + } + } + for (; k < kb; ++k) { + FLOAT *dst = Br + (BLASLONG)k * NR; + s2_v b0, b1; + if (ONLINE) { + const FLOAT *s = Bs + (BLASLONG)k * ldb; + b0 = S2_LD1(S2_PW(0, ncols), s); + S2_ST1(pt, dst, b0); + if (TN == 2) { + b1 = S2_LD1(S2_PW(S2_VL, ncols), s + S2_VL); + S2_ST1(pt, dst + S2_VL, b1); + } else { + b1 = b0; + } + } else { + b0 = S2_LD1(pt, dst); + b1 = TN == 2 ? S2_LD1(pt, dst + S2_VL) : b0; + } + S2_UNROLL for (int P = 0; P < TM; ++P) { + const s2_v a = S2_LD1(pt, Ar + P * astr + (BLASLONG)k * S2_VL); + s2_mopa(P * TN, pt, a, b0); + if (TN == 2) s2_mopa(P * TN + 1, pt, a, b1); + } + } + } + s2_c_store(TM, TN, C, ldc, mrows, ncols); +} + +/* One noinline function per kernel shape; the table index is TM - 1. */ +typedef void (*s2_kfn)(int, const FLOAT *, BLASLONG, FLOAT *, const FLOAT *, BLASLONG, int, FLOAT *, BLASLONG, int, + int, FLOAT, int) S2_S S2_ZA; +#define S2_DEF_KERNEL(TM, TN, ON) \ + S2_NOINL void s2_k_##TM##_##TN##_##ON(int kb, const FLOAT *Ar, BLASLONG astr, FLOAT *Br, const FLOAT *Bs, \ + BLASLONG ldb, int ncols, FLOAT *C, BLASLONG ldc, int mrows, int cmode, \ + FLOAT beta, int pfb) S2_S S2_ZA { \ + s2_kernel(TM, TN, ON, kb, Ar, astr, Br, Bs, ldb, ncols, C, ldc, mrows, cmode, beta, pfb); \ + } +#define S2_DEF_BOTH(TM, TN) S2_DEF_KERNEL(TM, TN, 0) S2_DEF_KERNEL(TM, TN, 1) +#ifdef DOUBLE +S2_DEF_BOTH(1, 8) +S2_DEF_BOTH(1, 4) +S2_DEF_BOTH(2, 4) +#else +S2_DEF_BOTH(1, 4) +S2_DEF_BOTH(1, 2) +S2_DEF_BOTH(2, 2) +#endif +S2_DEF_BOTH(1, 1) +S2_DEF_BOTH(2, 1) +S2_DEF_BOTH(3, 1) +S2_DEF_BOTH(4, 1) +#ifdef DOUBLE +S2_DEF_BOTH(5, 1) +S2_DEF_BOTH(6, 1) +S2_DEF_BOTH(7, 1) +S2_DEF_BOTH(8, 1) +#endif + +/* + * Transposition in ZA (paper Fig. 6): `rows` (<= VL) rows of kb contiguous elements (row stride ld) go in through + * horizontal slices of all tiles, CH columns at a time, and leave through vertical slices: dst[k * w + r]. + * w == VL gives the A panels; w > VL writes one VL-column group of a wider B panel. + */ +S2_INL void s2_pack_t(int rows, int kb, const FLOAT *src, BLASLONG ld, FLOAT *dst, const int w, FLOAT alpha, + const int scale, int pf) S2_S S2_ZA { + const svbool_t pt = S2_PT(); + if (rows <= 0) { + for (int k = 0; k < kb; ++k) S2_ST1(pt, dst + (BLASLONG)k * w, S2_ZEROV()); + return; + } + for (int k0 = 0; k0 < kb; k0 += S2_CH) { + const int kn = kb - k0 < S2_CH ? kb - k0 : S2_CH; + if (rows < S2_VL) svzero_za(); + int r = 0; + if (kn >= S2_CH) + for (; r + 4 <= rows; r += 4) { + const FLOAT *q = src + (BLASLONG)r * ld + k0; + if (pf && k0 + S2_CH < kb) + for (int u = 0; u < 4; ++u) s2_pf_l2(q + u * ld + S2_CH, S2_CH * (int)sizeof(FLOAT)); + s2_rows_in4(q, ld, r); + } + for (; r < rows; ++r) { + const FLOAT *q = src + (BLASLONG)r * ld + k0; + if (pf && k0 + S2_CH < kb) s2_pf_l2(q + S2_CH, S2_CH * (int)sizeof(FLOAT)); + S2_UNROLL for (int G = 0; G < S2_NT / 4; ++G) { + s2_v4 v = kn >= S2_CH ? s2_ld4(q + G * 4 * S2_VL) : s2_ld4n(q + G * 4 * S2_VL, s2_clip(kn - G * 4 * S2_VL, 4 * S2_VL)); + S2_UNROLL for (int Q = 0; Q < 4; ++Q) { + s2_wrh(G * 4 + Q, r, s2_get4(v, Q)); + } + } + } + S2_UNROLL for (int t = 0; t < S2_NT; ++t) { + const int kt = kn - t * S2_VL; + for (int c = 0; c < S2_VL && c < kt; c += 4) { + s2_v4 v = s2_rdv4(t, c); + if (scale) /* alpha on the way out, so that the four-row path serves scaled packing too */ + v = S2_CREATE4(S2_MUL(S2_GET4(v, 0), alpha), S2_MUL(S2_GET4(v, 1), alpha), S2_MUL(S2_GET4(v, 2), alpha), + S2_MUL(S2_GET4(v, 3), alpha)); + FLOAT *d = dst + (BLASLONG)(k0 + t * S2_VL + c) * w; + if (w == S2_VL) { + if (kt - c >= 4) s2_st4(d, v); + else s2_st4n(d, v, (int64_t)(kt - c) * S2_VL); + } else { + S2_UNROLL for (int u = 0; u < 4; ++u) + if (c + u < kt) S2_ST1(pt, d + (BLASLONG)u * w, s2_get4(v, u)); + } + } + } + } +} + +/* Specializations (as MTGEMM-A's templates): A panels, A panels scaled by alpha, and B panel groups. */ +S2_NOINL void s2_pack_a_rows(int rows, int kb, const FLOAT *src, BLASLONG ld, FLOAT *dst) S2_S S2_ZA { + s2_pack_t(rows, kb, src, ld, dst, S2_VL, (FLOAT)1, 0, 1); +} +S2_NOINL void s2_pack_a_rows_scaled(int rows, int kb, const FLOAT *src, BLASLONG ld, FLOAT *dst, FLOAT alpha) S2_S S2_ZA { + s2_pack_t(rows, kb, src, ld, dst, S2_VL, alpha, 1, 1); +} +S2_NOINL void s2_pack_b_group(int rows, int kb, const FLOAT *src, BLASLONG ld, FLOAT *dst, int w) S2_S S2_ZA { + if (w == S2_NR) s2_pack_t(rows, kb, src, ld, dst, S2_NR, (FLOAT)1, 0, 1); + else if (w == S2_W2) s2_pack_t(rows, kb, src, ld, dst, S2_W2, (FLOAT)1, 0, 1); + else s2_pack_t(rows, kb, src, ld, dst, S2_VL, (FLOAT)1, 0, 1); +} + +/* + * A with contiguous columns (element (r, k) at r + k * ld) into panels of VL rows: a copy, four panels and four + * depth steps at a time (four x4 loads, one x4 store per panel), since single dependent loads run at the SME + * unit's load latency. + */ +S2_INL void s2_pack_a_cols_t(int mb, int kb, const FLOAT *A, BLASLONG ld, FLOAT alpha, FLOAT *Ac, const int scale) + S2_S { + const svbool_t pt = S2_PT(); + for (int p0 = 0; p0 < mb; p0 += 4 * S2_VL) { + const int64_t rows = s2_clip(mb - p0, 4 * S2_VL); + const int np = (int)((rows + S2_VL - 1) / S2_VL); + const svcount_t pc = S2_CW(0, rows); + const FLOAT *s = A + p0; + FLOAT *d = Ac + (BLASLONG)p0 * kb; + int k = 0; + for (; k + 4 <= kb; k += 4) { + if (k + 8 + 4 <= kb) + for (int u = 0; u < 4; ++u) s2_pf_l2(s + (BLASLONG)(k + 8 + u) * ld, (int)rows * (int)sizeof(FLOAT)); + s2_v4 x0 = S2_LD4(pc, s + (BLASLONG)k * ld), x1 = S2_LD4(pc, s + (BLASLONG)(k + 1) * ld); + s2_v4 x2 = S2_LD4(pc, s + (BLASLONG)(k + 2) * ld), x3 = S2_LD4(pc, s + (BLASLONG)(k + 3) * ld); + S2_UNROLL for (int q = 0; q < 4; ++q) { + if (q >= np) break; + s2_v4 v = S2_CREATE4(s2_get4(x0, q), s2_get4(x1, q), s2_get4(x2, q), s2_get4(x3, q)); + if (scale) + v = S2_CREATE4(S2_MUL(S2_GET4(v, 0), alpha), S2_MUL(S2_GET4(v, 1), alpha), S2_MUL(S2_GET4(v, 2), alpha), + S2_MUL(S2_GET4(v, 3), alpha)); + s2_st4(d + (BLASLONG)q * S2_VL * kb + (BLASLONG)k * S2_VL, v); + } + } + for (; k < kb; ++k) { + const s2_v4 x = S2_LD4(pc, s + (BLASLONG)k * ld); + S2_UNROLL for (int q = 0; q < 4; ++q) { + if (q >= np) break; + s2_v v = s2_get4(x, q); + if (scale) v = S2_MUL(v, alpha); + S2_ST1(pt, d + (BLASLONG)q * S2_VL * kb + (BLASLONG)k * S2_VL, v); + } + } + } +} +S2_NOINL void s2_pack_a_cols(int mb, int kb, const FLOAT *A, BLASLONG ld, FLOAT alpha, FLOAT *Ac) + S2_S { + if (alpha != (FLOAT)1) s2_pack_a_cols_t(mb, kb, A, ld, alpha, Ac, 1); + else s2_pack_a_cols_t(mb, kb, A, ld, alpha, Ac, 0); +} + +/* B panel of `ncols` (<= w) columns from a B with contiguous columns (element (k, n) at k + n * ld). */ +S2_INL void s2_pack_b_cols(int kb, int ncols, int w, const FLOAT *Bs, BLASLONG ld, FLOAT *Bd) S2_S S2_ZA { + for (int g = 0; g < w; g += S2_VL) + s2_pack_b_group((int)s2_clip(ncols - g, S2_VL), kb, Bs + (BLASLONG)g * ld, ld, Bd + g, w); +} + +/* ---- blocking model (paper eqs. 1-3) ---- */ + +typedef struct { int mc, nc, kc; } s2_blocking; + +S2_INL int s2_round_up(int x, int m) { return (x + m - 1) / m * m; } + +static int s2_model_kc_max(int mr, int nr) { + int best = 16; + for (int kc = 16; kc <= 1 << 16; kc += 16) { + const int ta = (int)(((long)mr * kc * sizeof(FLOAT) + S2_PAGE_BYTES - 1) / S2_PAGE_BYTES) + 1; + const int tb = (int)(((long)nr * kc * sizeof(FLOAT) + S2_PAGE_BYTES - 1) / S2_PAGE_BYTES) + 1; + if (ta + 2 * tb + mr < S2_TLB_ENTRIES) best = kc; + else break; + } + return best; +} + +static int s2_balance(int total, int block, int step) { + if (block >= total) return s2_round_up(total, step); + const int nblk = (total + block - 1) / block; + const int b = s2_round_up((total + nblk - 1) / nblk, step); + return b < block ? b : block; +} + +static s2_blocking s2_model(int M, int N, int K, int mr, int nr) { + const long budget = S2_L2_BYTES / (long)sizeof(FLOAT); + int kcap = s2_model_kc_max(mr, nr); + if (kcap > s2_round_up(K, 16)) kcap = s2_round_up(K, 16); + const int mcap = s2_round_up(M, mr), ncap = s2_round_up(N, nr); + double best = -1; + s2_blocking r = {mr, nr, 16}; + for (int kc = 16; kc <= kcap; kc += 16) + for (int mc = mr; mc <= mcap; mc += mr) { + const long room = budget - (long)mc * kc; + if (room <= 0) break; + long nc = room / (2L * kc + 2L * mc); + nc = nc / nr * nr; + if (nc > ncap) nc = ncap; + if (nc < nr) break; + const double cmr = 2.0 * mc * nc * kc / ((double)mc * kc + (double)kc * nc + 2.0 * mc * nc); + if (cmr > best * (1 + 1e-9)) { + best = cmr; + r.mc = mc; + r.nc = (int)nc; + r.kc = kc; + } + } + r.kc = s2_balance(K, r.kc, 16); + r.mc = s2_balance(M, r.mc, mr); + r.nc = s2_balance(N, r.nc, nr); + return r; +} + +#if defined(_MSC_VER) && !defined(__clang__) +#define S2_TLS __declspec(thread) +#else +#define S2_TLS _Thread_local +#endif + +/* The search takes up to tens of microseconds for tall problems, so each thread remembers its last shapes. */ +static s2_blocking s2_blocking_for(int M, int N, int K) { + static S2_TLS struct { int M, N, K; s2_blocking b; } cache[8]; + static S2_TLS int next; + for (int i = 0; i < 8; ++i) + if (cache[i].M == M && cache[i].N == N && cache[i].K == K) return cache[i].b; + s2_blocking b = s2_model(M, N, K, S2_VL, S2_NR); + b.mc = s2_round_up(b.mc < S2_VL ? S2_VL : b.mc, S2_VL); + b.nc = s2_round_up(b.nc < S2_NR ? S2_NR : b.nc, S2_NR); + if (b.kc < 1) b.kc = 1; + cache[next].M = M; + cache[next].N = N; + cache[next].K = K; + cache[next].b = b; + next = (next + 1) & 7; + return b; +} + +/* ---- driver ---- */ + +/* + * Row-major core problem C = alpha A B + beta C, C of M x N (row stride ldc). + * a_cols: A has contiguous columns (element (i, k) at i + k * lda), else contiguous rows (i * lda + k). + * b_cols: B has contiguous columns (element (k, j) at k + j * ldb), else contiguous rows (k * ldb + j). + */ +typedef struct { + int M, N, K; + FLOAT alpha, beta; + const FLOAT *A, *B; + FLOAT *C; + BLASLONG lda, ldb, ldc; + int a_cols, b_cols; +} s2_job; + +static const s2_kfn s2_k_main[2] = { +#ifdef DOUBLE + s2_k_1_8_0, s2_k_1_8_1}; +static const s2_kfn s2_k_mid[2][2] = {{s2_k_1_4_0, s2_k_2_4_0}, {s2_k_1_4_1, s2_k_2_4_1}}; +static const s2_kfn s2_k_edge[2][S2_NT] = { + {s2_k_1_1_0, s2_k_2_1_0, s2_k_3_1_0, s2_k_4_1_0, s2_k_5_1_0, s2_k_6_1_0, s2_k_7_1_0, s2_k_8_1_0}, + {s2_k_1_1_1, s2_k_2_1_1, s2_k_3_1_1, s2_k_4_1_1, s2_k_5_1_1, s2_k_6_1_1, s2_k_7_1_1, s2_k_8_1_1}}; +#else + s2_k_1_4_0, s2_k_1_4_1}; +static const s2_kfn s2_k_mid[2][2] = {{s2_k_1_2_0, s2_k_2_2_0}, {s2_k_1_2_1, s2_k_2_2_1}}; +static const s2_kfn s2_k_edge[2][S2_NT] = {{s2_k_1_1_0, s2_k_2_1_0, s2_k_3_1_0, s2_k_4_1_0}, + {s2_k_1_1_1, s2_k_2_1_1, s2_k_3_1_1, s2_k_4_1_1}}; +#endif + +__arm_new("za") __arm_locally_streaming static void s2_drive(const s2_job *jb, s2_blocking blk, FLOAT *Ac, FLOAT *Bc) { + const int M = jb->M, N = jb->N, K = jb->K, mc = blk.mc, nc = blk.nc, kc = blk.kc; + const BLASLONG lda = jb->lda, ldb = jb->ldb, ldc = jb->ldc; + const int online = !jb->b_cols; + for (int i = 0; i < M; i += mc) { + const int mb = M - i < mc ? M - i : mc; + for (int k = 0; k < K; k += kc) { + const int kb = K - k < kc ? K - k : kc; + if (jb->a_cols) { + s2_pack_a_cols(mb, kb, jb->A + i + (BLASLONG)k * lda, lda, jb->alpha, Ac); + } else { + const FLOAT *Ap = jb->A + (BLASLONG)i * lda + k; + for (int p0 = 0; p0 < mb; p0 += S2_VL) { + const int rows = mb - p0 < S2_VL ? mb - p0 : S2_VL; + if (jb->alpha != (FLOAT)1) + s2_pack_a_rows_scaled(rows, kb, Ap + (BLASLONG)p0 * lda, lda, Ac + (BLASLONG)p0 * kb, jb->alpha); + else + s2_pack_a_rows(rows, kb, Ap + (BLASLONG)p0 * lda, lda, Ac + (BLASLONG)p0 * kb); + } + } + const int cmode = k > 0 ? S2_LOAD : (jb->beta == (FLOAT)0 ? S2_ZERO : (jb->beta == (FLOAT)1 ? S2_LOAD : S2_SCALE)); + for (int j = 0; j < N; j += nc) { + const int nb = N - j < nc ? N - j : nc, nmain = nb / S2_NR * S2_NR; + const FLOAT *Bs = jb->b_cols ? jb->B + (BLASLONG)j * ldb + k : jb->B + (BLASLONG)k * ldb + j; + FLOAT *Cb = jb->C + (BLASLONG)i * ldc + j; + /* Column tails wider than one vector go to the half-width kernel with two tile rows (predicated when + narrower than it), the last VL or fewer columns to the edge kernel. */ + const int ecol = nb - nmain > S2_VL ? nmain + (nb - nmain - 1) / S2_W2 * S2_W2 + ((nb - nmain - 1) % S2_W2 >= S2_VL ? S2_W2 : 0) : nmain; + if (!online) { + for (int jj = 0; jj < nmain; jj += S2_NR) + s2_pack_b_cols(kb, S2_NR, S2_NR, Bs + (BLASLONG)jj * ldb, ldb, Bc + (BLASLONG)jj * kb); + for (int jj = nmain; jj < ecol; jj += S2_W2) + s2_pack_b_cols(kb, nb - jj < S2_W2 ? nb - jj : S2_W2, S2_W2, Bs + (BLASLONG)jj * ldb, ldb, Bc + (BLASLONG)jj * kb); + for (int jj = ecol; jj < nb; jj += S2_VL) + s2_pack_b_cols(kb, nb - jj < S2_VL ? nb - jj : S2_VL, S2_VL, Bs + (BLASLONG)jj * ldb, ldb, Bc + (BLASLONG)jj * kb); + } + for (int ii = 0; ii < mb; ii += S2_VL) { + const int rows = mb - ii < S2_VL ? mb - ii : S2_VL; + for (int jj = 0; jj < nmain; jj += S2_NR) { + FLOAT *Cp = Cb + (BLASLONG)ii * ldc + jj; + FLOAT *Br = Bc + (BLASLONG)jj * kb; + const FLOAT *Cn = jj + S2_NR < nmain ? Cp + S2_NR : Cb + (BLASLONG)(ii + S2_VL) * ldc; + if (jj + S2_NR < nmain || ii + S2_VL < mb) /* next C tile towards L2 while this one computes */ + for (int r = 0; r < rows; ++r) s2_pf_l2(Cn + (BLASLONG)r * ldc, S2_NR * (int)sizeof(FLOAT)); + const int on = online && ii == 0; + s2_k_main[on](kb, Ac + (BLASLONG)ii * kb, (BLASLONG)S2_VL * kb, Br, on ? Bs + jj : NULL, ldb, S2_NR, Cp, + ldc, rows, cmode, jb->beta, 1); + } + } + for (int jj = nmain; jj < ecol; jj += S2_W2) { + const int cols = nb - jj < S2_W2 ? nb - jj : S2_W2; + for (int ii = 0; ii < mb; ii += 2 * S2_VL) { + const int rows = mb - ii < 2 * S2_VL ? mb - ii : 2 * S2_VL, on = online && ii == 0; + s2_k_mid[on][(rows + S2_VL - 1) / S2_VL - 1](kb, Ac + (BLASLONG)ii * kb, (BLASLONG)S2_VL * kb, + Bc + (BLASLONG)jj * kb, on ? Bs + jj : NULL, ldb, cols, + Cb + (BLASLONG)ii * ldc + jj, ldc, rows, cmode, jb->beta, 1); + } + } + if (ecol < nb) /* edge kernel: all tiles stacked along M, one VL-wide column panel */ + for (int ii = 0; ii < mb; ii += S2_NT * S2_VL) { + const int rows = mb - ii < S2_NT * S2_VL ? mb - ii : S2_NT * S2_VL; + for (int jj = ecol; jj < nb; jj += S2_VL) { + const int cols = nb - jj < S2_VL ? nb - jj : S2_VL, on = online && ii == 0; + s2_k_edge[on][(rows + S2_VL - 1) / S2_VL - 1](kb, Ac + (BLASLONG)ii * kb, (BLASLONG)S2_VL * kb, + Bc + (BLASLONG)jj * kb, on ? Bs + jj : NULL, ldb, cols, + Cb + (BLASLONG)ii * ldc + jj, ldc, rows, cmode, jb->beta, 1); + } + } + } + } + } +} + +#define S2_THREAD_WORK (1L << 20) /* largest per-thread packing buffer; larger problems use blas_memory_alloc */ + +/* Bytes of packed A and B for blocking b (the model keeps this near S2_L2_BYTES, well inside BUFFER_SIZE). */ +static size_t s2_work_bytes(s2_blocking b, size_t *a_bytes) { + *a_bytes = ((size_t)s2_round_up(b.mc, 4 * S2_VL) * b.kc * sizeof(FLOAT) + 127) & ~(size_t)127; + return *a_bytes + (size_t)b.nc * b.kc * sizeof(FLOAT) + 128; +} + +/* + * Per-thread packing buffer for problems up to S2_THREAD_WORK bytes, grown on demand and freed when the thread + * exits: the per-call cost of blas_memory_alloc, and of packing on the stack next to the core's own data, is + * significant against a small GEMM (fp32 20^3 column-major: 568 ns on the stack, 336 ns here). + */ +static pthread_key_t s2_work_key; +static pthread_once_t s2_work_once = PTHREAD_ONCE_INIT; +static void s2_work_free(void *p) { free(p); } +static void s2_work_key_init(void) { pthread_key_create(&s2_work_key, s2_work_free); } +static void *s2_thread_work(size_t bytes) { + pthread_once(&s2_work_once, s2_work_key_init); + size_t *p = (size_t *)pthread_getspecific(s2_work_key); + if (p == NULL || p[0] < bytes) { + size_t cap = p ? p[0] : 0; + while (cap < bytes) cap = cap ? 2 * cap : 65536; + free(p); + void *q = NULL; + if (posix_memalign(&q, 16384, cap + 128)) q = NULL; + p = (size_t *)q; + if (p) p[0] = cap; + pthread_setspecific(s2_work_key, p); + if (!p) return NULL; + } + return (char *)p + 128; +} + +/* Runs one job on a workspace of at least s2_work_bytes() bytes (NULL: a per-thread or blas_memory_alloc buffer). */ +static void s2_run(const s2_job *jb, void *work) { + if (jb->M <= 0 || jb->N <= 0) return; + const s2_blocking b = s2_blocking_for(jb->M, jb->N, jb->K); + size_t a_bytes; + const size_t bytes = s2_work_bytes(b, &a_bytes); + void *buffer = NULL, *heap = NULL; + if (work == NULL && bytes <= S2_THREAD_WORK) work = s2_thread_work(bytes); + if (work == NULL && (buffer = blas_memory_alloc(0)) != NULL) work = (char *)buffer + GEMM_OFFSET_A; + if (work == NULL && (work = heap = malloc(bytes + 128)) == NULL) return; + char *base = (char *)(((uintptr_t)work + 127) & ~(uintptr_t)127); + s2_drive(jb, b, (FLOAT *)base, (FLOAT *)(base + a_bytes)); + if (buffer) blas_memory_free(buffer); + free(heap); +} + +/* ---- runtime checks, threads, entry ---- */ + +/* + * SME units that threads can use: Apple M4 and M5 have one per cluster of the fastest cores. M5 Pro and M5 Max, + * whose second tier is "Performance" rather than "Efficiency", have two per "Super" cluster (third-party + * measurement, not verified here). Elsewhere assume one. + */ +static int s2_units(void) { + static int n = 0; + if (n == 0) { + int u = 1; +#if defined(__APPLE__) + int cpus = 0, per = 0; + char name[32] = ""; + size_t len = sizeof(int); + if (sysctlbyname("hw.perflevel0.physicalcpu", &cpus, &len, NULL, 0) == 0 && + (len = sizeof(int), sysctlbyname("hw.perflevel0.cpusperl2", &per, &len, NULL, 0) == 0) && per > 0 && + cpus / per > 1) + u = cpus / per; + len = sizeof(name) - 1; + if (sysctlbyname("hw.perflevel1.name", name, &len, NULL, 0) == 0 && strcmp(name, "Performance") == 0) u *= 2; +#endif + n = u; + } + return n; +} + +#ifdef SMP +static int s2_routine(blas_arg_t *args, BLASLONG *range_m, BLASLONG *range_n, FLOAT *sa, FLOAT *sb, BLASLONG pos) { + s2_run((const s2_job *)args->common, sa); + return 0; +} +#endif + +#define S2_MAX_PARTS 4 + +/* Column-major C = alpha op(A) op(B) + beta C, solved as the row-major C^T = alpha op(B)^T op(A)^T + beta C^T. */ +static void s2_gemm(int trans_a, int trans_b, BLASLONG m, BLASLONG n, BLASLONG k, FLOAT alpha, const FLOAT *a, + BLASLONG lda, const FLOAT *b, BLASLONG ldb, FLOAT beta, FLOAT *c, BLASLONG ldc) { + if (m <= 0 || n <= 0) return; + if (k <= 0 || alpha == (FLOAT)0) { + for (BLASLONG j = 0; j < n; ++j) + for (BLASLONG i = 0; i < m; ++i) c[i + j * ldc] = beta == (FLOAT)0 ? (FLOAT)0 : beta * c[i + j * ldc]; + return; + } + s2_job jb; + jb.M = (int)n; + jb.N = (int)m; + jb.K = (int)k; + jb.alpha = alpha; + jb.beta = beta; + jb.A = b; + jb.lda = ldb; + jb.a_cols = trans_b; + jb.B = a; + jb.ldb = lda; + jb.b_cols = trans_a; + jb.C = c; + jb.ldc = ldc; + + int parts = 1; +#ifdef SMP + /* One thread per SME unit from 2^22 multiply-adds, each part at least four panels wide along the split. */ + const int big = jb.M > jb.N ? jb.M : jb.N; + if ((double)jb.M * jb.N * jb.K >= (double)(1 << 22)) { + parts = s2_units(); + const int avail = num_cpu_avail(3); + if (parts > avail) parts = avail; + if (parts > big / (4 * S2_VL)) parts = big / (4 * S2_VL); + if (parts > S2_MAX_PARTS) parts = S2_MAX_PARTS; + if (parts < 1) parts = 1; + } + if (parts > 1) { + s2_job pj[S2_MAX_PARTS]; + blas_arg_t args[S2_MAX_PARTS]; + blas_queue_t queue[S2_MAX_PARTS]; + const int split_m = jb.M >= jb.N; + const int total = split_m ? jb.M : jb.N; + const int step = s2_round_up((total + parts - 1) / parts, 4 * S2_VL); + int used = 0; + for (int p = 0, off = 0; p < parts && off < total; ++p, off += step) { + const int len = total - off < step ? total - off : step; + pj[p] = jb; + if (split_m) { + pj[p].M = len; + pj[p].A = jb.a_cols ? jb.A + off : jb.A + (BLASLONG)off * jb.lda; + pj[p].C = jb.C + (BLASLONG)off * jb.ldc; + } else { + pj[p].N = len; + pj[p].B = jb.b_cols ? jb.B + (BLASLONG)off * jb.ldb : jb.B + off; + pj[p].C = jb.C + off; + } + memset(&args[p], 0, sizeof(args[p])); + args[p].common = &pj[p]; + memset(&queue[p], 0, sizeof(queue[p])); +#ifdef DOUBLE + queue[p].mode = BLAS_DOUBLE | BLAS_REAL; +#else + queue[p].mode = BLAS_SINGLE | BLAS_REAL; +#endif + queue[p].routine = (void *)s2_routine; + queue[p].args = &args[p]; + queue[p].next = &queue[p + 1]; + used = p + 1; + } + queue[used - 1].next = NULL; + void *buffer = blas_memory_alloc(0); /* the calling thread runs queue[0] on the workspace it passes */ + if (buffer != NULL) { + queue[0].sa = (char *)buffer + GEMM_OFFSET_A; + exec_blas(used, queue); + blas_memory_free(buffer); + return; + } + } +#endif + s2_run(&jb, NULL); +} + +#pragma clang attribute pop diff --git a/kernel/arm64/sme_dgemm_kernel.c b/kernel/arm64/sme_dgemm_kernel.c index 2cc02c952d..ec53fdb7ec 100644 --- a/kernel/arm64/sme_dgemm_kernel.c +++ b/kernel/arm64/sme_dgemm_kernel.c @@ -6,6 +6,12 @@ #include #include #include "common.h" +#if defined(__clang__) && defined(__ARM_FEATURE_SME) +#if (defined(__apple_build_version__) && __clang_major__ >= 17) || (!defined(__apple_build_version__) && __clang_major__ >= 18) +#define HAVE_SME2_GEMM 1 +#include "sme2_gemm_impl.h" +#endif +#endif #ifndef stdmin #define stdmin(a,b) (a>b? b:a) #endif @@ -1185,6 +1191,14 @@ void CNAME(const char *transa, const char *transb, const BLASLONG m, const BLASL { bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); +#if defined(HAVE_SME2_GEMM) + /* below 5000 multiply-adds (smaller than 4000 take NEON) the existing kernel is faster when it needs no + edge tiles (M and N multiples of 16) */ + if (s2_usable() && S2_FITS(m, n, k) && !((double)m * (double)n * (double)k < 5000. && m % 16 == 0 && n % 16 == 0)) { + s2_gemm(trans_a, trans_b, m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); + return; + } +#endif if (!trans_a && !trans_b) { dgemm_sme_NN(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); } diff --git a/kernel/arm64/sme_level3.h b/kernel/arm64/sme_level3.h new file mode 100644 index 0000000000..e217c9cf13 --- /dev/null +++ b/kernel/arm64/sme_level3.h @@ -0,0 +1,418 @@ +/*************************************************************************** +Copyright (c) 2026, 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. +*****************************************************************************/ + +/* + * SYMM, SYRK, SYR2K, TRMM and TRSM on arm64 SME targets, on top of the SME GEMM kernel (SME_[SD]GEMM_KERNEL). + * + * The level-3 driver runs these routines on the NEON GEMM kernel of the target. Here each routine splits its + * symmetric or triangular dimension in two, recursively: the off-diagonal blocks are GEMM calls on the SME + * kernel, and only blocks of at most S3_NB rows or columns remain. Those run as one GEMM on a full copy of + * the block (SYMM, SYRK, SYR2K, TRMM), or, for TRSM, as substitution in C on blocks of at most S3_NB_TRSM, so + * that every solve keeps the arithmetic of substitution (the generic TRSM kernels of these targets run at 10-16 + * GFLOPS on such thin blocks). All arguments are in the column-major form that the interface files build. + * + * Included by interface/{symm,syrk,syr2k,trsm}.c; they call the s3_*_hook functions at the end of this file, + * which return 1 when they have handled the call. + */ + +#if !defined(COMPLEX) && !defined(XDOUBLE) && !defined(BFLOAT16) && !defined(HFLOAT16) && defined(ARCH_ARM64) && \ + (defined(USE_SGEMM_KERNEL_DIRECT) || defined(DYNAMIC_ARCH)) +#define SME_LEVEL3 1 + +#include +#include + +#define S3_NB 64 /* largest diagonal block; also the split granularity */ +#define S3_MIN_WORK 2e5 /* multiply-adds below which the level-3 driver keeps the call */ +#define S3_BCHUNK 4096 /* columns (rows) of B per temporary copy in the TRMM base case */ +#ifndef S3_NB_TRSM +#define S3_NB_TRSM 32 /* TRSM blocks solved by substitution */ +#endif + +#ifdef DYNAMIC_ARCH +extern char *gotoblas_corename(void); +#endif + +/* Whether to take this path: large enough, and the SME GEMM kernel is what gemm.c calls on this core. */ +static inline int s3_enabled(double work) { + if (work < S3_MIN_WORK) return 0; +#if defined(DYNAMIC_ARCH) + const char *c = gotoblas_corename(); + if (strcmp(c, "armv9sme") != 0 +#if defined(__clang__) + && strcmp(c, "vortexm4") != 0 +#endif + ) + return 0; +#endif + return 1; +} + +/* Column-major C = alpha op(A) op(B) + beta C on the SME kernel. */ +static inline void s3_gemm(int ta, int tb, BLASLONG m, BLASLONG n, BLASLONG k, FLOAT alpha, FLOAT *a, BLASLONG lda, + FLOAT *b, BLASLONG ldb, FLOAT beta, FLOAT *c, BLASLONG ldc) { + if (m <= 0 || n <= 0) return; + char *TA = ta ? "T" : "N", *TB = tb ? "T" : "N"; +#ifdef DOUBLE + SME_DGEMM_KERNEL(TA, TB, m, n, k, &alpha, a, lda, b, ldb, &beta, c, ldc); +#else + SME_SGEMM_KERNEL(TA, TB, m, n, k, &alpha, a, lda, b, ldb, &beta, c, ldc); +#endif +} + +/* First part of a split of t > S3_NB: about half, a multiple of S3_NB. */ +static inline BLASLONG s3_split(BLASLONG t) { return (t / 2 + S3_NB - 1) / S3_NB * S3_NB; } + +/* Address of element (i, j) of op(A), where op(A) is A (tr 0) or its transpose (tr 1). */ +#define S3_OP(a, lda, tr, i, j) ((tr) ? (a) + (j) + (BLASLONG)(i) * (lda) : (a) + (i) + (BLASLONG)(j) * (lda)) + +/* ---- SYMM: C = alpha A B + beta C (side 0) or alpha B A + beta C (side 1), A symmetric ---- */ + +/* The t x t diagonal block of a symmetric matrix with both triangles filled. */ +static inline void s3_sym_full(int lower, BLASLONG t, FLOAT *a, BLASLONG lda, FLOAT *f) { + for (BLASLONG j = 0; j < t; ++j) + for (BLASLONG i = 0; i < t; ++i) f[i + j * t] = ((i >= j) == lower) ? a[i + j * lda] : a[j + i * lda]; +} + +static inline void s3_symm_rec(int side, int lower, BLASLONG m, BLASLONG n, FLOAT alpha, FLOAT *a, BLASLONG lda, FLOAT *b, + BLASLONG ldb, FLOAT beta, FLOAT *c, BLASLONG ldc) { + const BLASLONG t = side ? n : m; + if (t <= S3_NB) { + FLOAT f[S3_NB * S3_NB]; + s3_sym_full(lower, t, a, lda, f); + if (!side) s3_gemm(0, 0, m, n, m, alpha, f, t, b, ldb, beta, c, ldc); + else s3_gemm(0, 0, m, n, n, alpha, b, ldb, f, t, beta, c, ldc); + return; + } + const BLASLONG t1 = s3_split(t), t2 = t - t1; + FLOAT *a22 = a + t1 + t1 * lda; + /* A21 is stored for lower; for upper it is the transpose of the stored A12 */ + FLOAT *a21 = lower ? a + t1 : a + t1 * lda; + const int tr21 = !lower; + if (!side) { /* C1 = A11 B1 + A21^T B2, C2 = A21 B1 + A22 B2 */ + s3_symm_rec(0, lower, t1, n, alpha, a, lda, b, ldb, beta, c, ldc); + s3_gemm(!tr21, 0, t1, n, t2, alpha, a21, lda, b + t1, ldb, 1, c, ldc); + s3_symm_rec(0, lower, t2, n, alpha, a22, lda, b + t1, ldb, beta, c + t1, ldc); + s3_gemm(tr21, 0, t2, n, t1, alpha, a21, lda, b, ldb, 1, c + t1, ldc); + } else { /* C1 = B1 A11 + B2 A21, C2 = B1 A21^T + B2 A22 */ + s3_symm_rec(1, lower, m, t1, alpha, a, lda, b, ldb, beta, c, ldc); + s3_gemm(0, tr21, m, t1, t2, alpha, b + t1 * ldb, ldb, a21, lda, 1, c, ldc); + s3_symm_rec(1, lower, m, t2, alpha, a22, lda, b + t1 * ldb, ldb, beta, c + t1 * ldc, ldc); + s3_gemm(0, !tr21, m, t2, t1, alpha, b, ldb, a21, lda, 1, c + t1 * ldc, ldc); + } +} + +/* ---- SYRK (two 0) and SYR2K (two 1) on the lower or upper triangle of the n x n C ---- */ + +/* Rows r0.. of op(X): X is n x k for tr 0, k x n for tr 1. */ +#define S3_ROWS(x, ld, tr, r0) ((tr) ? (x) + (BLASLONG)(r0) * (ld) : (x) + (r0)) + +/* C = alpha (op(A) op(B)^T + [two] op(B) op(A)^T) + beta C; SYRK passes B = A. */ +static inline void s3_syrk_rec(int two, int lower, int tr, BLASLONG n, BLASLONG k, FLOAT alpha, FLOAT *a, BLASLONG lda, + FLOAT *b, BLASLONG ldb, FLOAT beta, FLOAT *c, BLASLONG ldc) { + if (n <= S3_NB) { + FLOAT f[S3_NB * S3_NB]; + s3_gemm(tr, !tr, n, n, k, alpha, a, lda, b, ldb, 0, f, n); + if (two) s3_gemm(tr, !tr, n, n, k, alpha, b, ldb, a, lda, 1, f, n); + for (BLASLONG j = 0; j < n; ++j) + for (BLASLONG i = lower ? j : 0; i < (lower ? n : j + 1); ++i) + c[i + j * ldc] = f[i + j * n] + (beta == 0 ? 0 : beta * c[i + j * ldc]); + return; + } + const BLASLONG n1 = s3_split(n), n2 = n - n1; + FLOAT *a2 = S3_ROWS(a, lda, tr, n1), *b2 = S3_ROWS(b, ldb, tr, n1); + s3_syrk_rec(two, lower, tr, n1, k, alpha, a, lda, b, ldb, beta, c, ldc); + if (lower) { /* C21 = alpha (A2 B1^T + [two] B2 A1^T) + beta C21 */ + s3_gemm(tr, !tr, n2, n1, k, alpha, a2, lda, b, ldb, beta, c + n1, ldc); + if (two) s3_gemm(tr, !tr, n2, n1, k, alpha, b2, ldb, a, lda, 1, c + n1, ldc); + } else { /* C12 = alpha (A1 B2^T + [two] B1 A2^T) + beta C12 */ + s3_gemm(tr, !tr, n1, n2, k, alpha, a, lda, b2, ldb, beta, c + n1 * ldc, ldc); + if (two) s3_gemm(tr, !tr, n1, n2, k, alpha, b, ldb, a2, lda, 1, c + n1 * ldc, ldc); + } + s3_syrk_rec(two, lower, tr, n2, k, alpha, a2, lda, b2, ldb, beta, c + n1 + n1 * ldc, ldc); +} + +/* ---- TRMM: B = alpha op(A) B (side 0) or alpha B op(A) (side 1), A triangular ---- */ + +/* The t x t diagonal block of op(A) as a full matrix: zeros outside the triangle, ones on a unit diagonal. */ +static inline void s3_tri_full(int lower_op, int tr, int unit, BLASLONG t, FLOAT *a, BLASLONG lda, FLOAT *f) { + for (BLASLONG j = 0; j < t; ++j) + for (BLASLONG i = 0; i < t; ++i) + f[i + j * t] = i == j ? (unit ? 1 : *S3_OP(a, lda, tr, i, i)) + : (((i > j) == lower_op) ? *S3_OP(a, lda, tr, i, j) : 0); +} + +/* lower_op: op(A) is lower triangular; work holds S3_NB * S3_BCHUNK elements. */ +static inline void s3_trmm_rec(int side, int lower_op, int tr, int unit, BLASLONG m, BLASLONG n, FLOAT alpha, FLOAT *a, + BLASLONG lda, FLOAT *b, BLASLONG ldb, FLOAT *work) { + const BLASLONG t = side ? n : m; + if (t <= S3_NB) { + FLOAT f[S3_NB * S3_NB]; + s3_tri_full(lower_op, tr, unit, t, a, lda, f); + if (!side) { /* B = alpha F B, S3_BCHUNK columns at a time through a copy */ + for (BLASLONG j0 = 0; j0 < n; j0 += S3_BCHUNK) { + const BLASLONG nc = n - j0 < S3_BCHUNK ? n - j0 : S3_BCHUNK; + for (BLASLONG j = 0; j < nc; ++j) memcpy(work + j * t, b + (j0 + j) * ldb, t * sizeof(FLOAT)); + s3_gemm(0, 0, t, nc, t, alpha, f, t, work, t, 0, b + j0 * ldb, ldb); + } + } else { /* B = alpha B F, S3_BCHUNK rows at a time through a copy */ + for (BLASLONG i0 = 0; i0 < m; i0 += S3_BCHUNK) { + const BLASLONG mc = m - i0 < S3_BCHUNK ? m - i0 : S3_BCHUNK; + for (BLASLONG j = 0; j < t; ++j) memcpy(work + j * mc, b + i0 + j * ldb, mc * sizeof(FLOAT)); + s3_gemm(0, 0, mc, t, t, alpha, work, mc, f, t, 0, b + i0, ldb); + } + } + return; + } + const BLASLONG t1 = s3_split(t), t2 = t - t1; + FLOAT *a22 = a + t1 + t1 * lda; + FLOAT *a21 = S3_OP(a, lda, tr, t1, 0), *a12 = S3_OP(a, lda, tr, 0, t1); + if (!side) { + FLOAT *b2 = b + t1; + if (lower_op) { /* B2 = A22 B2 + A21 B1, then B1 = A11 B1 */ + s3_trmm_rec(0, 1, tr, unit, t2, n, alpha, a22, lda, b2, ldb, work); + s3_gemm(tr, 0, t2, n, t1, alpha, a21, lda, b, ldb, 1, b2, ldb); + s3_trmm_rec(0, 1, tr, unit, t1, n, alpha, a, lda, b, ldb, work); + } else { /* B1 = A11 B1 + A12 B2, then B2 = A22 B2 */ + s3_trmm_rec(0, 0, tr, unit, t1, n, alpha, a, lda, b, ldb, work); + s3_gemm(tr, 0, t1, n, t2, alpha, a12, lda, b2, ldb, 1, b, ldb); + s3_trmm_rec(0, 0, tr, unit, t2, n, alpha, a22, lda, b2, ldb, work); + } + } else { + FLOAT *b2 = b + t1 * ldb; + if (lower_op) { /* B1 = B1 A11 + B2 A21, then B2 = B2 A22 */ + s3_trmm_rec(1, 1, tr, unit, m, t1, alpha, a, lda, b, ldb, work); + s3_gemm(0, tr, m, t1, t2, alpha, b2, ldb, a21, lda, 1, b, ldb); + s3_trmm_rec(1, 1, tr, unit, m, t2, alpha, a22, lda, b2, ldb, work); + } else { /* B2 = B2 A22 + B1 A12, then B1 = B1 A11 */ + s3_trmm_rec(1, 0, tr, unit, m, t2, alpha, a22, lda, b2, ldb, work); + s3_gemm(0, tr, m, t2, t1, alpha, b, ldb, a12, lda, 1, b2, ldb); + s3_trmm_rec(1, 0, tr, unit, m, t1, alpha, a, lda, b, ldb, work); + } + } +} + +/* ---- TRSM: op(A) X = alpha B (side 0) or X op(A) = alpha B (side 1), X overwrites B ---- */ + +/* + * Left-side base solve: x[p][c] = alpha B(row of solve step p, j0 + c) for 16 columns, and back. Rows of B are + * strided, so blocks of rows x columns are moved with NEON transposes (4 x 4 fp32, 2 x 2 fp64) instead of one + * strided scalar per element, which cost as much as the solve itself. Step p is row p (fwd) or t - 1 - p. + */ +#ifdef DOUBLE +#define S3_TB 2 +#else +#define S3_TB 4 +#endif +static inline void s3_tile_io(int load, int fwd, BLASLONG t, BLASLONG nc, FLOAT alpha, FLOAT *b, BLASLONG ldb, + FLOAT x[][16]) { + const BLASLONG cfull = nc / S3_TB * S3_TB, pfull = t / S3_TB * S3_TB; + for (BLASLONG c0 = 0; c0 < cfull; c0 += S3_TB) + for (BLASLONG p0 = 0; p0 < pfull; p0 += S3_TB) { + /* rows rb..rb+TB-1 of B hold steps p0..p0+TB-1 (ascending if fwd, descending otherwise) */ + const BLASLONG rb = fwd ? p0 : t - S3_TB - p0; + FLOAT *col = b + rb + c0 * ldb; +#ifdef DOUBLE + if (load) { + const float64x2_t v0 = vld1q_f64(col), v1 = vld1q_f64(col + ldb), a = vdupq_n_f64(alpha); + const float64x2_t r0 = vmulq_f64(vzip1q_f64(v0, v1), a), r1 = vmulq_f64(vzip2q_f64(v0, v1), a); + vst1q_f64(&x[fwd ? p0 : p0 + 1][c0], r0); + vst1q_f64(&x[fwd ? p0 + 1 : p0][c0], r1); + } else { + const float64x2_t r0 = vld1q_f64(&x[fwd ? p0 : p0 + 1][c0]), r1 = vld1q_f64(&x[fwd ? p0 + 1 : p0][c0]); + vst1q_f64(col, vzip1q_f64(r0, r1)); + vst1q_f64(col + ldb, vzip2q_f64(r0, r1)); + } +#else + float32x4_t v0, v1, v2, v3; + if (load) { + v0 = vld1q_f32(col); v1 = vld1q_f32(col + ldb); v2 = vld1q_f32(col + 2 * ldb); v3 = vld1q_f32(col + 3 * ldb); + } else { /* rows of the tile, in B's row order */ + v0 = vld1q_f32(&x[fwd ? p0 : p0 + 3][c0]); v1 = vld1q_f32(&x[fwd ? p0 + 1 : p0 + 2][c0]); + v2 = vld1q_f32(&x[fwd ? p0 + 2 : p0 + 1][c0]); v3 = vld1q_f32(&x[fwd ? p0 + 3 : p0][c0]); + } + const float32x4x2_t t01 = vtrnq_f32(v0, v1), t23 = vtrnq_f32(v2, v3); + float32x4_t r0 = vcombine_f32(vget_low_f32(t01.val[0]), vget_low_f32(t23.val[0])); + float32x4_t r1 = vcombine_f32(vget_low_f32(t01.val[1]), vget_low_f32(t23.val[1])); + float32x4_t r2 = vcombine_f32(vget_high_f32(t01.val[0]), vget_high_f32(t23.val[0])); + float32x4_t r3 = vcombine_f32(vget_high_f32(t01.val[1]), vget_high_f32(t23.val[1])); + if (load) { + r0 = vmulq_n_f32(r0, alpha); r1 = vmulq_n_f32(r1, alpha); r2 = vmulq_n_f32(r2, alpha); r3 = vmulq_n_f32(r3, alpha); + vst1q_f32(&x[fwd ? p0 : p0 + 3][c0], r0); vst1q_f32(&x[fwd ? p0 + 1 : p0 + 2][c0], r1); + vst1q_f32(&x[fwd ? p0 + 2 : p0 + 1][c0], r2); vst1q_f32(&x[fwd ? p0 + 3 : p0][c0], r3); + } else { + vst1q_f32(col, r0); vst1q_f32(col + ldb, r1); vst1q_f32(col + 2 * ldb, r2); vst1q_f32(col + 3 * ldb, r3); + } +#endif + } + /* the rest element by element: columns beyond the last full group, and steps beyond the last full group */ + for (BLASLONG c = 0; c < 16; ++c) + for (BLASLONG p = (c < cfull ? pfull : 0); p < t; ++p) { + FLOAT *e = b + (fwd ? p : t - 1 - p) + c * ldb; + if (load) x[p][c] = c < nc ? alpha * *e : 0; + else if (c < nc) *e = x[p][c]; + } +} + +/* + * Substitution on a block of at most S3_NB_TRSM, in dot form: each result accumulates in registers over the + * results solved before it, 16 columns (left side) or 16 rows (right side) at a time, so that each vector FMA + * needs one load. Indices run in solve order; op(A) is first copied in that order. The diagonal is applied as a + * reciprocal, as the OpenBLAS TRSM kernels do; alpha comes first. + */ +static inline void s3_trsm_base(int side, int lower_op, int tr, int unit, BLASLONG m, BLASLONG n, FLOAT alpha, + FLOAT *a, BLASLONG lda, FLOAT *b, BLASLONG ldb) { + const BLASLONG t = side ? n : m; + FLOAT rinv[S3_NB_TRSM], ta[S3_NB_TRSM * S3_NB_TRSM]; + /* solve order: left lower and right upper forward, the others backward */ + const int fwd = side ? !lower_op : lower_op; +#define S3_ORD(x) (fwd ? (x) : t - 1 - (x)) + for (BLASLONG p = 0; p < t; ++p) { + const BLASLONG q = S3_ORD(p); + rinv[p] = unit ? 1 : 1 / *S3_OP(a, lda, tr, q, q); + for (BLASLONG r = 0; r < p; ++r) /* coefficient of solved result r in result p */ + ta[p * t + r] = side ? *S3_OP(a, lda, tr, S3_ORD(r), q) : *S3_OP(a, lda, tr, q, S3_ORD(r)); + } + if (!side) { + for (BLASLONG j0 = 0; j0 < n; j0 += 16) { /* 16 columns of B at a time, solved in a local tile */ + const BLASLONG nc = n - j0 < 16 ? n - j0 : 16; + FLOAT x[S3_NB_TRSM][16]; + s3_tile_io(1, fwd, t, nc, alpha, b + j0 * ldb, ldb, x); + for (BLASLONG p = 0; p < t; ++p) { + FLOAT acc[16]; + for (int c = 0; c < 16; ++c) acc[c] = x[p][c]; + for (BLASLONG r = 0; r < p; ++r) { + const FLOAT l = ta[p * t + r]; + for (int c = 0; c < 16; ++c) acc[c] -= l * x[r][c]; + } + for (int c = 0; c < 16; ++c) x[p][c] = acc[c] * rinv[p]; + } + s3_tile_io(0, fwd, t, nc, alpha, b + j0 * ldb, ldb, x); + } + } else { + for (BLASLONG i0 = 0; i0 < m; i0 += 16) { /* 16 rows of B at a time, solved in a local tile */ + const BLASLONG mr = m - i0 < 16 ? m - i0 : 16; + FLOAT x[S3_NB_TRSM][16]; + for (BLASLONG p = 0; p < t; ++p) { + const FLOAT *bj = b + i0 + S3_ORD(p) * ldb; + FLOAT acc[16]; + if (mr == 16) + for (int r = 0; r < 16; ++r) acc[r] = alpha * bj[r]; + else + for (int r = 0; r < 16; ++r) acc[r] = r < mr ? alpha * bj[r] : 0; + for (BLASLONG q = 0; q < p; ++q) { + const FLOAT c = ta[p * t + q]; + for (int r = 0; r < 16; ++r) acc[r] -= c * x[q][r]; + } + for (int r = 0; r < 16; ++r) x[p][r] = acc[r] * rinv[p]; + } + for (BLASLONG p = 0; p < t; ++p) { + FLOAT *bj = b + i0 + S3_ORD(p) * ldb; + for (BLASLONG r = 0; r < mr; ++r) bj[r] = x[p][r]; + } + } + } +#undef S3_ORD +} + +static inline void s3_trsm_rec(int side, int lower_op, int tr, int unit, BLASLONG m, BLASLONG n, FLOAT alpha, + FLOAT *a, BLASLONG lda, FLOAT *b, BLASLONG ldb) { + const BLASLONG t = side ? n : m; + if (t <= S3_NB_TRSM) { + s3_trsm_base(side, lower_op, tr, unit, m, n, alpha, a, lda, b, ldb); + return; + } + const BLASLONG t1 = (t / 2 + S3_NB_TRSM - 1) / S3_NB_TRSM * S3_NB_TRSM, t2 = t - t1; + FLOAT *a22 = a + t1 + t1 * lda; + FLOAT *a21 = S3_OP(a, lda, tr, t1, 0), *a12 = S3_OP(a, lda, tr, 0, t1); + if (!side) { + FLOAT *b2 = b + t1; + if (lower_op) { /* X1 = A11 \ alpha B1; B2 = alpha B2 - A21 X1; X2 = A22 \ B2 */ + s3_trsm_rec(0, 1, tr, unit, t1, n, alpha, a, lda, b, ldb); + s3_gemm(tr, 0, t2, n, t1, -1, a21, lda, b, ldb, alpha, b2, ldb); + s3_trsm_rec(0, 1, tr, unit, t2, n, 1, a22, lda, b2, ldb); + } else { /* X2 = A22 \ alpha B2; B1 = alpha B1 - A12 X2; X1 = A11 \ B1 */ + s3_trsm_rec(0, 0, tr, unit, t2, n, alpha, a22, lda, b2, ldb); + s3_gemm(tr, 0, t1, n, t2, -1, a12, lda, b2, ldb, alpha, b, ldb); + s3_trsm_rec(0, 0, tr, unit, t1, n, 1, a, lda, b, ldb); + } + } else { + FLOAT *b2 = b + t1 * ldb; + if (lower_op) { /* X2 = alpha B2 / A22; B1 = alpha B1 - X2 A21; X1 = B1 / A11 */ + s3_trsm_rec(1, 1, tr, unit, m, t2, alpha, a22, lda, b2, ldb); + s3_gemm(0, tr, m, t1, t2, -1, b2, ldb, a21, lda, alpha, b, ldb); + s3_trsm_rec(1, 1, tr, unit, m, t1, 1, a, lda, b, ldb); + } else { /* X1 = alpha B1 / A11; B2 = alpha B2 - X1 A12; X2 = B2 / A22 */ + s3_trsm_rec(1, 0, tr, unit, m, t1, alpha, a, lda, b, ldb); + s3_gemm(0, tr, m, t2, t1, -1, b, ldb, a12, lda, alpha, b2, ldb); + s3_trsm_rec(1, 0, tr, unit, m, t2, 1, a22, lda, b2, ldb); + } + } +} + +/* ---- entry points for the interface files (arguments as they build them); 1: the call was handled here ---- */ + +static inline int s3_symm_hook(int side, int uplo, blas_arg_t *args) { + if (!s3_enabled(side ? (double)args->m * args->n * args->n : (double)args->m * args->m * args->n)) return 0; + if (!side) + s3_symm_rec(0, uplo, args->m, args->n, *(FLOAT *)args->alpha, (FLOAT *)args->a, args->lda, (FLOAT *)args->b, + args->ldb, *(FLOAT *)args->beta, (FLOAT *)args->c, args->ldc); + else /* for the right side args->a is the general matrix and args->b the symmetric one */ + s3_symm_rec(1, uplo, args->m, args->n, *(FLOAT *)args->alpha, (FLOAT *)args->b, args->ldb, (FLOAT *)args->a, + args->lda, *(FLOAT *)args->beta, (FLOAT *)args->c, args->ldc); + return 1; +} + +/* SYRK (two 0) or SYR2K (two 1) */ +static inline int s3_syrk_hook(int two, int uplo, int trans, blas_arg_t *args) { + if (args->k <= 0 || !s3_enabled((double)args->n * args->n * args->k * (two ? 1.0 : 0.5))) return 0; + s3_syrk_rec(two, uplo, trans & 1, args->n, args->k, *(FLOAT *)args->alpha, (FLOAT *)args->a, args->lda, + two ? (FLOAT *)args->b : (FLOAT *)args->a, two ? args->ldb : args->lda, *(FLOAT *)args->beta, + (FLOAT *)args->c, args->ldc); + return 1; +} + +/* TRSM, or TRMM when interface/trsm.c is compiled with TRMM; alpha is in args->beta as trsm.c stores it */ +static inline int s3_trxm_hook(int side, int uplo, int trans, int unit, blas_arg_t *args) { + if ((side ? args->n : args->m) <= S3_NB_TRSM || + !s3_enabled(side ? (double)args->m * args->n * args->n : (double)args->m * args->m * args->n)) + return 0; + const int lower_op = (uplo == 1) != (trans & 1); +#ifndef TRMM + s3_trsm_rec(side, lower_op, trans & 1, unit == 0, args->m, args->n, *(FLOAT *)args->beta, (FLOAT *)args->a, + args->lda, (FLOAT *)args->b, args->ldb); +#else + FLOAT *work = (FLOAT *)blas_memory_alloc(0); + if (!work) return 0; + s3_trmm_rec(side, lower_op, trans & 1, unit == 0, args->m, args->n, *(FLOAT *)args->beta, (FLOAT *)args->a, + args->lda, (FLOAT *)args->b, args->ldb, work); + blas_memory_free(work); +#endif + return 1; +} + +#endif diff --git a/kernel/arm64/sme_sgemm_kernel.c b/kernel/arm64/sme_sgemm_kernel.c index 4fcc5cb566..3db4dd3f8e 100644 --- a/kernel/arm64/sme_sgemm_kernel.c +++ b/kernel/arm64/sme_sgemm_kernel.c @@ -5,6 +5,12 @@ #include #include #include "common.h" +#if defined(__clang__) && defined(__ARM_FEATURE_SME) +#if (defined(__apple_build_version__) && __clang_major__ >= 17) || (!defined(__apple_build_version__) && __clang_major__ >= 18) +#define HAVE_SME2_GEMM 1 +#include "sme2_gemm_impl.h" +#endif +#endif #ifndef stdmin #define stdmin(a,b) (a>b? b:a) #endif @@ -618,6 +624,12 @@ void CNAME(char *transa, char *transb, BLASLONG m, BLASLONG n, BLASLONG k, f { bool trans_a = (*transa == 'T' || *transa == 't' || *transa == 'C' || *transa == 'c'); bool trans_b = (*transb == 'T' || *transb == 't' || *transb == 'C' || *transb == 'c'); +#if defined(HAVE_SME2_GEMM) + if (s2_usable() && S2_FITS(m, n, k)) { + s2_gemm(trans_a, trans_b, m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); + return; + } +#endif if (!trans_a && !trans_b) { sgemm_sme_NN(m, n, k, *alpha, a, lda, b, ldb, *beta, c, ldc); }