diff --git a/docs/2026.html b/docs/2026.html
index 2d64f78520..80e4326551 100644
--- a/docs/2026.html
+++ b/docs/2026.html
@@ -53,6 +53,10 @@
New features
SVE2 optimizations of class SynetInnerProduct32fProd.
NEON, SVE2 optimizations of class SynetInnerProduct16bGemmNN.
+Improving
+
+ - SVE2 optimizations of function Gemm32fNT.
+
Renaming
- Rename function SimdVectorProduct to SimdCreateMask.
diff --git a/src/Simd/SimdSve2Gemm32fNT.cpp b/src/Simd/SimdSve2Gemm32fNT.cpp
index d9ba684c1e..01e770866b 100644
--- a/src/Simd/SimdSve2Gemm32fNT.cpp
+++ b/src/Simd/SimdSve2Gemm32fNT.cpp
@@ -34,132 +34,46 @@ namespace Simd
return svaddv_f32(svptrue_b32(), value);
}
-#define SIMD_SVE2_GEMM_INIT1(row) \
- svfloat32_t c##row##0 = zero;
-
-#define SIMD_SVE2_GEMM_INIT4(row) \
- svfloat32_t c##row##0 = zero, c##row##1 = zero, c##row##2 = zero, c##row##3 = zero;
-
-#define SIMD_SVE2_GEMM_ROW1(row, offset) \
- { \
- svfloat32_t a##row = svld1_f32(mask, A##row + (offset)); \
- c##row##0 = svmla_f32_m(mask, c##row##0, a##row, b0); \
- }
-
-#define SIMD_SVE2_GEMM_ROW4(row, offset) \
- { \
- svfloat32_t a##row = svld1_f32(mask, A##row + (offset)); \
- c##row##0 = svmla_f32_m(mask, c##row##0, a##row, b0); \
- c##row##1 = svmla_f32_m(mask, c##row##1, a##row, b1); \
- c##row##2 = svmla_f32_m(mask, c##row##2, a##row, b2); \
- c##row##3 = svmla_f32_m(mask, c##row##3, a##row, b3); \
- }
-
-#define SIMD_SVE2_GEMM_STEP1_1(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- SIMD_SVE2_GEMM_ROW1(0, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP1_2(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- SIMD_SVE2_GEMM_ROW1(0, offset); \
- SIMD_SVE2_GEMM_ROW1(1, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP1_3(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- SIMD_SVE2_GEMM_ROW1(0, offset); \
- SIMD_SVE2_GEMM_ROW1(1, offset); \
- SIMD_SVE2_GEMM_ROW1(2, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP1_6(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- SIMD_SVE2_GEMM_ROW1(0, offset); \
- SIMD_SVE2_GEMM_ROW1(1, offset); \
- SIMD_SVE2_GEMM_ROW1(2, offset); \
- SIMD_SVE2_GEMM_ROW1(3, offset); \
- SIMD_SVE2_GEMM_ROW1(4, offset); \
- SIMD_SVE2_GEMM_ROW1(5, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP4_1(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- svfloat32_t b1 = svld1_f32(mask, B1 + (offset)); \
- svfloat32_t b2 = svld1_f32(mask, B2 + (offset)); \
- svfloat32_t b3 = svld1_f32(mask, B3 + (offset)); \
- SIMD_SVE2_GEMM_ROW4(0, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP4_2(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- svfloat32_t b1 = svld1_f32(mask, B1 + (offset)); \
- svfloat32_t b2 = svld1_f32(mask, B2 + (offset)); \
- svfloat32_t b3 = svld1_f32(mask, B3 + (offset)); \
- SIMD_SVE2_GEMM_ROW4(0, offset); \
- SIMD_SVE2_GEMM_ROW4(1, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP4_3(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- svfloat32_t b1 = svld1_f32(mask, B1 + (offset)); \
- svfloat32_t b2 = svld1_f32(mask, B2 + (offset)); \
- svfloat32_t b3 = svld1_f32(mask, B3 + (offset)); \
- SIMD_SVE2_GEMM_ROW4(0, offset); \
- SIMD_SVE2_GEMM_ROW4(1, offset); \
- SIMD_SVE2_GEMM_ROW4(2, offset); \
- }
-
-#define SIMD_SVE2_GEMM_STEP4_6(offset) \
- { \
- svfloat32_t b0 = svld1_f32(mask, B0 + (offset)); \
- svfloat32_t b1 = svld1_f32(mask, B1 + (offset)); \
- svfloat32_t b2 = svld1_f32(mask, B2 + (offset)); \
- svfloat32_t b3 = svld1_f32(mask, B3 + (offset)); \
- SIMD_SVE2_GEMM_ROW4(0, offset); \
- SIMD_SVE2_GEMM_ROW4(1, offset); \
- SIMD_SVE2_GEMM_ROW4(2, offset); \
- SIMD_SVE2_GEMM_ROW4(3, offset); \
- SIMD_SVE2_GEMM_ROW4(4, offset); \
- SIMD_SVE2_GEMM_ROW4(5, offset); \
+ SIMD_INLINE void Add4ExtractedSums(const svfloat32_t& sum0, const svfloat32_t& sum1,
+ const svfloat32_t& sum2, const svfloat32_t& sum3, float alpha, float* dst)
+ {
+ dst[0] += alpha * ExtractSum32f(sum0);
+ dst[1] += alpha * ExtractSum32f(sum1);
+ dst[2] += alpha * ExtractSum32f(sum2);
+ dst[3] += alpha * ExtractSum32f(sum3);
}
-#define SIMD_SVE2_GEMM_SAVE1(row) \
- C[(row) * ldc] += alpha * ExtractSum32f(c##row##0);
-
-#define SIMD_SVE2_GEMM_SAVE4(row) \
- C[(row) * ldc + 0] += alpha * ExtractSum32f(c##row##0); \
- C[(row) * ldc + 1] += alpha * ExtractSum32f(c##row##1); \
- C[(row) * ldc + 2] += alpha * ExtractSum32f(c##row##2); \
- C[(row) * ldc + 3] += alpha * ExtractSum32f(c##row##3);
-
static void Kernel1x1nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
{
size_t F = svcntw(), DF = 2 * F, k = 0;
const float* A0 = A + 0 * lda;
const float* B0 = B + 0 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT1(0);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c00b = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP1_1(k);
- SIMD_SVE2_GEMM_STEP1_1(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00b = svmla_f32_x(body, c00b, a0, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP1_1(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
}
- SIMD_SVE2_GEMM_SAVE1(0);
+ C[0] += alpha * ExtractSum32f(svadd_f32_x(body, c00, c00b));
}
static void Kernel1x4nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -171,20 +85,66 @@ namespace Simd
const float* B2 = B + 2 * ldb;
const float* B3 = B + 3 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT4(0);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c01 = svdup_n_f32(0.0f);
+ svfloat32_t c02 = svdup_n_f32(0.0f);
+ svfloat32_t c03 = svdup_n_f32(0.0f);
+ svfloat32_t c00b = svdup_n_f32(0.0f);
+ svfloat32_t c01b = svdup_n_f32(0.0f);
+ svfloat32_t c02b = svdup_n_f32(0.0f);
+ svfloat32_t c03b = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP4_1(k);
- SIMD_SVE2_GEMM_STEP4_1(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00b = svmla_f32_x(body, c00b, a0, b0);
+ b0 = svld1_f32(body, B1 + k + F);
+ c01b = svmla_f32_x(body, c01b, a0, b0);
+ b0 = svld1_f32(body, B2 + k + F);
+ c02b = svmla_f32_x(body, c02b, a0, b0);
+ b0 = svld1_f32(body, B3 + k + F);
+ c03b = svmla_f32_x(body, c03b, a0, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP4_1(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ b0 = svld1_f32(mask, B1 + k);
+ c01 = svmla_f32_m(mask, c01, a0, b0);
+ b0 = svld1_f32(mask, B2 + k);
+ c02 = svmla_f32_m(mask, c02, a0, b0);
+ b0 = svld1_f32(mask, B3 + k);
+ c03 = svmla_f32_m(mask, c03, a0, b0);
}
- SIMD_SVE2_GEMM_SAVE4(0);
+ Add4ExtractedSums(
+ svadd_f32_x(body, c00, c00b),
+ svadd_f32_x(body, c01, c01b),
+ svadd_f32_x(body, c02, c02b),
+ svadd_f32_x(body, c03, c03b),
+ alpha, C + 0 * ldc);
}
static void Kernel2x1nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -194,22 +154,42 @@ namespace Simd
const float* A1 = A + 1 * lda;
const float* B0 = B + 0 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT1(0);
- SIMD_SVE2_GEMM_INIT1(1);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c10 = svdup_n_f32(0.0f);
+ svfloat32_t c00b = svdup_n_f32(0.0f);
+ svfloat32_t c10b = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP1_2(k);
- SIMD_SVE2_GEMM_STEP1_2(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ a1 = svld1_f32(body, A1 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00b = svmla_f32_x(body, c00b, a0, b0);
+ c10b = svmla_f32_x(body, c10b, a1, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP1_2(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t a1 = svld1_f32(mask, A1 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ c10 = svmla_f32_m(mask, c10, a1, b0);
}
- SIMD_SVE2_GEMM_SAVE1(0);
- SIMD_SVE2_GEMM_SAVE1(1);
+ C[0 * ldc] += alpha * ExtractSum32f(svadd_f32_x(body, c00, c00b));
+ C[1 * ldc] += alpha * ExtractSum32f(svadd_f32_x(body, c10, c10b));
}
static void Kernel2x4nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -222,22 +202,82 @@ namespace Simd
const float* B2 = B + 2 * ldb;
const float* B3 = B + 3 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT4(0);
- SIMD_SVE2_GEMM_INIT4(1);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c01 = svdup_n_f32(0.0f);
+ svfloat32_t c02 = svdup_n_f32(0.0f);
+ svfloat32_t c03 = svdup_n_f32(0.0f);
+ svfloat32_t c10 = svdup_n_f32(0.0f);
+ svfloat32_t c11 = svdup_n_f32(0.0f);
+ svfloat32_t c12 = svdup_n_f32(0.0f);
+ svfloat32_t c13 = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP4_2(k);
- SIMD_SVE2_GEMM_STEP4_2(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ a1 = svld1_f32(body, A1 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ b0 = svld1_f32(body, B1 + k + F);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ b0 = svld1_f32(body, B2 + k + F);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ b0 = svld1_f32(body, B3 + k + F);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP4_2(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t a1 = svld1_f32(mask, A1 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ c10 = svmla_f32_m(mask, c10, a1, b0);
+ b0 = svld1_f32(mask, B1 + k);
+ c01 = svmla_f32_m(mask, c01, a0, b0);
+ c11 = svmla_f32_m(mask, c11, a1, b0);
+ b0 = svld1_f32(mask, B2 + k);
+ c02 = svmla_f32_m(mask, c02, a0, b0);
+ c12 = svmla_f32_m(mask, c12, a1, b0);
+ b0 = svld1_f32(mask, B3 + k);
+ c03 = svmla_f32_m(mask, c03, a0, b0);
+ c13 = svmla_f32_m(mask, c13, a1, b0);
}
- SIMD_SVE2_GEMM_SAVE4(0);
- SIMD_SVE2_GEMM_SAVE4(1);
+ Add4ExtractedSums(c00, c01, c02, c03, alpha, C + 0 * ldc);
+ Add4ExtractedSums(c10, c11, c12, c13, alpha, C + 1 * ldc);
}
static void Kernel3x1nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -248,24 +288,53 @@ namespace Simd
const float* A2 = A + 2 * lda;
const float* B0 = B + 0 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT1(0);
- SIMD_SVE2_GEMM_INIT1(1);
- SIMD_SVE2_GEMM_INIT1(2);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c10 = svdup_n_f32(0.0f);
+ svfloat32_t c20 = svdup_n_f32(0.0f);
+ svfloat32_t c00b = svdup_n_f32(0.0f);
+ svfloat32_t c10b = svdup_n_f32(0.0f);
+ svfloat32_t c20b = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP1_3(k);
- SIMD_SVE2_GEMM_STEP1_3(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ a1 = svld1_f32(body, A1 + k + F);
+ a2 = svld1_f32(body, A2 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00b = svmla_f32_x(body, c00b, a0, b0);
+ c10b = svmla_f32_x(body, c10b, a1, b0);
+ c20b = svmla_f32_x(body, c20b, a2, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP1_3(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t a1 = svld1_f32(mask, A1 + k);
+ svfloat32_t a2 = svld1_f32(mask, A2 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ c10 = svmla_f32_m(mask, c10, a1, b0);
+ c20 = svmla_f32_m(mask, c20, a2, b0);
}
- SIMD_SVE2_GEMM_SAVE1(0);
- SIMD_SVE2_GEMM_SAVE1(1);
- SIMD_SVE2_GEMM_SAVE1(2);
+ C[0 * ldc] += alpha * ExtractSum32f(svadd_f32_x(body, c00, c00b));
+ C[1 * ldc] += alpha * ExtractSum32f(svadd_f32_x(body, c10, c10b));
+ C[2 * ldc] += alpha * ExtractSum32f(svadd_f32_x(body, c20, c20b));
}
static void Kernel3x4nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -279,24 +348,107 @@ namespace Simd
const float* B2 = B + 2 * ldb;
const float* B3 = B + 3 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT4(0);
- SIMD_SVE2_GEMM_INIT4(1);
- SIMD_SVE2_GEMM_INIT4(2);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c01 = svdup_n_f32(0.0f);
+ svfloat32_t c02 = svdup_n_f32(0.0f);
+ svfloat32_t c03 = svdup_n_f32(0.0f);
+ svfloat32_t c10 = svdup_n_f32(0.0f);
+ svfloat32_t c11 = svdup_n_f32(0.0f);
+ svfloat32_t c12 = svdup_n_f32(0.0f);
+ svfloat32_t c13 = svdup_n_f32(0.0f);
+ svfloat32_t c20 = svdup_n_f32(0.0f);
+ svfloat32_t c21 = svdup_n_f32(0.0f);
+ svfloat32_t c22 = svdup_n_f32(0.0f);
+ svfloat32_t c23 = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP4_3(k);
- SIMD_SVE2_GEMM_STEP4_3(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ c21 = svmla_f32_x(body, c21, a2, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ c22 = svmla_f32_x(body, c22, a2, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ c23 = svmla_f32_x(body, c23, a2, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ a1 = svld1_f32(body, A1 + k + F);
+ a2 = svld1_f32(body, A2 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ b0 = svld1_f32(body, B1 + k + F);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ c21 = svmla_f32_x(body, c21, a2, b0);
+ b0 = svld1_f32(body, B2 + k + F);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ c22 = svmla_f32_x(body, c22, a2, b0);
+ b0 = svld1_f32(body, B3 + k + F);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ c23 = svmla_f32_x(body, c23, a2, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ c21 = svmla_f32_x(body, c21, a2, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ c22 = svmla_f32_x(body, c22, a2, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ c23 = svmla_f32_x(body, c23, a2, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP4_3(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t a1 = svld1_f32(mask, A1 + k);
+ svfloat32_t a2 = svld1_f32(mask, A2 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ c10 = svmla_f32_m(mask, c10, a1, b0);
+ c20 = svmla_f32_m(mask, c20, a2, b0);
+ b0 = svld1_f32(mask, B1 + k);
+ c01 = svmla_f32_m(mask, c01, a0, b0);
+ c11 = svmla_f32_m(mask, c11, a1, b0);
+ c21 = svmla_f32_m(mask, c21, a2, b0);
+ b0 = svld1_f32(mask, B2 + k);
+ c02 = svmla_f32_m(mask, c02, a0, b0);
+ c12 = svmla_f32_m(mask, c12, a1, b0);
+ c22 = svmla_f32_m(mask, c22, a2, b0);
+ b0 = svld1_f32(mask, B3 + k);
+ c03 = svmla_f32_m(mask, c03, a0, b0);
+ c13 = svmla_f32_m(mask, c13, a1, b0);
+ c23 = svmla_f32_m(mask, c23, a2, b0);
}
- SIMD_SVE2_GEMM_SAVE4(0);
- SIMD_SVE2_GEMM_SAVE4(1);
- SIMD_SVE2_GEMM_SAVE4(2);
+ Add4ExtractedSums(c00, c01, c02, c03, alpha, C + 0 * ldc);
+ Add4ExtractedSums(c10, c11, c12, c13, alpha, C + 1 * ldc);
+ Add4ExtractedSums(c20, c21, c22, c23, alpha, C + 2 * ldc);
}
static void Kernel6x1nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -310,30 +462,80 @@ namespace Simd
const float* A5 = A + 5 * lda;
const float* B0 = B + 0 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT1(0);
- SIMD_SVE2_GEMM_INIT1(1);
- SIMD_SVE2_GEMM_INIT1(2);
- SIMD_SVE2_GEMM_INIT1(3);
- SIMD_SVE2_GEMM_INIT1(4);
- SIMD_SVE2_GEMM_INIT1(5);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c10 = svdup_n_f32(0.0f);
+ svfloat32_t c20 = svdup_n_f32(0.0f);
+ svfloat32_t c30 = svdup_n_f32(0.0f);
+ svfloat32_t c40 = svdup_n_f32(0.0f);
+ svfloat32_t c50 = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP1_6(k);
- SIMD_SVE2_GEMM_STEP1_6(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t a3 = svld1_f32(body, A3 + k);
+ svfloat32_t a4 = svld1_f32(body, A4 + k);
+ svfloat32_t a5 = svld1_f32(body, A5 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ c30 = svmla_f32_x(body, c30, a3, b0);
+ c40 = svmla_f32_x(body, c40, a4, b0);
+ c50 = svmla_f32_x(body, c50, a5, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ a1 = svld1_f32(body, A1 + k + F);
+ a2 = svld1_f32(body, A2 + k + F);
+ a3 = svld1_f32(body, A3 + k + F);
+ a4 = svld1_f32(body, A4 + k + F);
+ a5 = svld1_f32(body, A5 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ c30 = svmla_f32_x(body, c30, a3, b0);
+ c40 = svmla_f32_x(body, c40, a4, b0);
+ c50 = svmla_f32_x(body, c50, a5, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t a3 = svld1_f32(body, A3 + k);
+ svfloat32_t a4 = svld1_f32(body, A4 + k);
+ svfloat32_t a5 = svld1_f32(body, A5 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ c30 = svmla_f32_x(body, c30, a3, b0);
+ c40 = svmla_f32_x(body, c40, a4, b0);
+ c50 = svmla_f32_x(body, c50, a5, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP1_6(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t a1 = svld1_f32(mask, A1 + k);
+ svfloat32_t a2 = svld1_f32(mask, A2 + k);
+ svfloat32_t a3 = svld1_f32(mask, A3 + k);
+ svfloat32_t a4 = svld1_f32(mask, A4 + k);
+ svfloat32_t a5 = svld1_f32(mask, A5 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ c10 = svmla_f32_m(mask, c10, a1, b0);
+ c20 = svmla_f32_m(mask, c20, a2, b0);
+ c30 = svmla_f32_m(mask, c30, a3, b0);
+ c40 = svmla_f32_m(mask, c40, a4, b0);
+ c50 = svmla_f32_m(mask, c50, a5, b0);
}
- SIMD_SVE2_GEMM_SAVE1(0);
- SIMD_SVE2_GEMM_SAVE1(1);
- SIMD_SVE2_GEMM_SAVE1(2);
- SIMD_SVE2_GEMM_SAVE1(3);
- SIMD_SVE2_GEMM_SAVE1(4);
- SIMD_SVE2_GEMM_SAVE1(5);
+ C[0 * ldc] += alpha * ExtractSum32f(c00);
+ C[1 * ldc] += alpha * ExtractSum32f(c10);
+ C[2 * ldc] += alpha * ExtractSum32f(c20);
+ C[3 * ldc] += alpha * ExtractSum32f(c30);
+ C[4 * ldc] += alpha * ExtractSum32f(c40);
+ C[5 * ldc] += alpha * ExtractSum32f(c50);
}
static void Kernel6x4nt(size_t K, float alpha, const float* A, size_t lda, const float* B, size_t ldb, float* C, size_t ldc)
@@ -350,30 +552,182 @@ namespace Simd
const float* B2 = B + 2 * ldb;
const float* B3 = B + 3 * ldb;
const svbool_t body = svptrue_b32();
- const svfloat32_t zero = svdup_n_f32(0.0f);
- SIMD_SVE2_GEMM_INIT4(0);
- SIMD_SVE2_GEMM_INIT4(1);
- SIMD_SVE2_GEMM_INIT4(2);
- SIMD_SVE2_GEMM_INIT4(3);
- SIMD_SVE2_GEMM_INIT4(4);
- SIMD_SVE2_GEMM_INIT4(5);
+ svfloat32_t c00 = svdup_n_f32(0.0f);
+ svfloat32_t c01 = svdup_n_f32(0.0f);
+ svfloat32_t c02 = svdup_n_f32(0.0f);
+ svfloat32_t c03 = svdup_n_f32(0.0f);
+ svfloat32_t c10 = svdup_n_f32(0.0f);
+ svfloat32_t c11 = svdup_n_f32(0.0f);
+ svfloat32_t c12 = svdup_n_f32(0.0f);
+ svfloat32_t c13 = svdup_n_f32(0.0f);
+ svfloat32_t c20 = svdup_n_f32(0.0f);
+ svfloat32_t c21 = svdup_n_f32(0.0f);
+ svfloat32_t c22 = svdup_n_f32(0.0f);
+ svfloat32_t c23 = svdup_n_f32(0.0f);
+ svfloat32_t c30 = svdup_n_f32(0.0f);
+ svfloat32_t c31 = svdup_n_f32(0.0f);
+ svfloat32_t c32 = svdup_n_f32(0.0f);
+ svfloat32_t c33 = svdup_n_f32(0.0f);
+ svfloat32_t c40 = svdup_n_f32(0.0f);
+ svfloat32_t c41 = svdup_n_f32(0.0f);
+ svfloat32_t c42 = svdup_n_f32(0.0f);
+ svfloat32_t c43 = svdup_n_f32(0.0f);
+ svfloat32_t c50 = svdup_n_f32(0.0f);
+ svfloat32_t c51 = svdup_n_f32(0.0f);
+ svfloat32_t c52 = svdup_n_f32(0.0f);
+ svfloat32_t c53 = svdup_n_f32(0.0f);
for (; k + DF <= K; k += DF)
{
- const svbool_t mask = body;
- SIMD_SVE2_GEMM_STEP4_6(k);
- SIMD_SVE2_GEMM_STEP4_6(k + F);
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t a3 = svld1_f32(body, A3 + k);
+ svfloat32_t a4 = svld1_f32(body, A4 + k);
+ svfloat32_t a5 = svld1_f32(body, A5 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ c30 = svmla_f32_x(body, c30, a3, b0);
+ c40 = svmla_f32_x(body, c40, a4, b0);
+ c50 = svmla_f32_x(body, c50, a5, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ c21 = svmla_f32_x(body, c21, a2, b0);
+ c31 = svmla_f32_x(body, c31, a3, b0);
+ c41 = svmla_f32_x(body, c41, a4, b0);
+ c51 = svmla_f32_x(body, c51, a5, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ c22 = svmla_f32_x(body, c22, a2, b0);
+ c32 = svmla_f32_x(body, c32, a3, b0);
+ c42 = svmla_f32_x(body, c42, a4, b0);
+ c52 = svmla_f32_x(body, c52, a5, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ c23 = svmla_f32_x(body, c23, a2, b0);
+ c33 = svmla_f32_x(body, c33, a3, b0);
+ c43 = svmla_f32_x(body, c43, a4, b0);
+ c53 = svmla_f32_x(body, c53, a5, b0);
+ a0 = svld1_f32(body, A0 + k + F);
+ a1 = svld1_f32(body, A1 + k + F);
+ a2 = svld1_f32(body, A2 + k + F);
+ a3 = svld1_f32(body, A3 + k + F);
+ a4 = svld1_f32(body, A4 + k + F);
+ a5 = svld1_f32(body, A5 + k + F);
+ b0 = svld1_f32(body, B0 + k + F);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ c30 = svmla_f32_x(body, c30, a3, b0);
+ c40 = svmla_f32_x(body, c40, a4, b0);
+ c50 = svmla_f32_x(body, c50, a5, b0);
+ b0 = svld1_f32(body, B1 + k + F);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ c21 = svmla_f32_x(body, c21, a2, b0);
+ c31 = svmla_f32_x(body, c31, a3, b0);
+ c41 = svmla_f32_x(body, c41, a4, b0);
+ c51 = svmla_f32_x(body, c51, a5, b0);
+ b0 = svld1_f32(body, B2 + k + F);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ c22 = svmla_f32_x(body, c22, a2, b0);
+ c32 = svmla_f32_x(body, c32, a3, b0);
+ c42 = svmla_f32_x(body, c42, a4, b0);
+ c52 = svmla_f32_x(body, c52, a5, b0);
+ b0 = svld1_f32(body, B3 + k + F);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ c23 = svmla_f32_x(body, c23, a2, b0);
+ c33 = svmla_f32_x(body, c33, a3, b0);
+ c43 = svmla_f32_x(body, c43, a4, b0);
+ c53 = svmla_f32_x(body, c53, a5, b0);
+ }
+ for (; k + F <= K; k += F)
+ {
+ svfloat32_t a0 = svld1_f32(body, A0 + k);
+ svfloat32_t a1 = svld1_f32(body, A1 + k);
+ svfloat32_t a2 = svld1_f32(body, A2 + k);
+ svfloat32_t a3 = svld1_f32(body, A3 + k);
+ svfloat32_t a4 = svld1_f32(body, A4 + k);
+ svfloat32_t a5 = svld1_f32(body, A5 + k);
+ svfloat32_t b0 = svld1_f32(body, B0 + k);
+ c00 = svmla_f32_x(body, c00, a0, b0);
+ c10 = svmla_f32_x(body, c10, a1, b0);
+ c20 = svmla_f32_x(body, c20, a2, b0);
+ c30 = svmla_f32_x(body, c30, a3, b0);
+ c40 = svmla_f32_x(body, c40, a4, b0);
+ c50 = svmla_f32_x(body, c50, a5, b0);
+ b0 = svld1_f32(body, B1 + k);
+ c01 = svmla_f32_x(body, c01, a0, b0);
+ c11 = svmla_f32_x(body, c11, a1, b0);
+ c21 = svmla_f32_x(body, c21, a2, b0);
+ c31 = svmla_f32_x(body, c31, a3, b0);
+ c41 = svmla_f32_x(body, c41, a4, b0);
+ c51 = svmla_f32_x(body, c51, a5, b0);
+ b0 = svld1_f32(body, B2 + k);
+ c02 = svmla_f32_x(body, c02, a0, b0);
+ c12 = svmla_f32_x(body, c12, a1, b0);
+ c22 = svmla_f32_x(body, c22, a2, b0);
+ c32 = svmla_f32_x(body, c32, a3, b0);
+ c42 = svmla_f32_x(body, c42, a4, b0);
+ c52 = svmla_f32_x(body, c52, a5, b0);
+ b0 = svld1_f32(body, B3 + k);
+ c03 = svmla_f32_x(body, c03, a0, b0);
+ c13 = svmla_f32_x(body, c13, a1, b0);
+ c23 = svmla_f32_x(body, c23, a2, b0);
+ c33 = svmla_f32_x(body, c33, a3, b0);
+ c43 = svmla_f32_x(body, c43, a4, b0);
+ c53 = svmla_f32_x(body, c53, a5, b0);
}
if (k < K)
{
const svbool_t mask = svwhilelt_b32(k, K);
- SIMD_SVE2_GEMM_STEP4_6(k);
+ svfloat32_t a0 = svld1_f32(mask, A0 + k);
+ svfloat32_t a1 = svld1_f32(mask, A1 + k);
+ svfloat32_t a2 = svld1_f32(mask, A2 + k);
+ svfloat32_t a3 = svld1_f32(mask, A3 + k);
+ svfloat32_t a4 = svld1_f32(mask, A4 + k);
+ svfloat32_t a5 = svld1_f32(mask, A5 + k);
+ svfloat32_t b0 = svld1_f32(mask, B0 + k);
+ c00 = svmla_f32_m(mask, c00, a0, b0);
+ c10 = svmla_f32_m(mask, c10, a1, b0);
+ c20 = svmla_f32_m(mask, c20, a2, b0);
+ c30 = svmla_f32_m(mask, c30, a3, b0);
+ c40 = svmla_f32_m(mask, c40, a4, b0);
+ c50 = svmla_f32_m(mask, c50, a5, b0);
+ b0 = svld1_f32(mask, B1 + k);
+ c01 = svmla_f32_m(mask, c01, a0, b0);
+ c11 = svmla_f32_m(mask, c11, a1, b0);
+ c21 = svmla_f32_m(mask, c21, a2, b0);
+ c31 = svmla_f32_m(mask, c31, a3, b0);
+ c41 = svmla_f32_m(mask, c41, a4, b0);
+ c51 = svmla_f32_m(mask, c51, a5, b0);
+ b0 = svld1_f32(mask, B2 + k);
+ c02 = svmla_f32_m(mask, c02, a0, b0);
+ c12 = svmla_f32_m(mask, c12, a1, b0);
+ c22 = svmla_f32_m(mask, c22, a2, b0);
+ c32 = svmla_f32_m(mask, c32, a3, b0);
+ c42 = svmla_f32_m(mask, c42, a4, b0);
+ c52 = svmla_f32_m(mask, c52, a5, b0);
+ b0 = svld1_f32(mask, B3 + k);
+ c03 = svmla_f32_m(mask, c03, a0, b0);
+ c13 = svmla_f32_m(mask, c13, a1, b0);
+ c23 = svmla_f32_m(mask, c23, a2, b0);
+ c33 = svmla_f32_m(mask, c33, a3, b0);
+ c43 = svmla_f32_m(mask, c43, a4, b0);
+ c53 = svmla_f32_m(mask, c53, a5, b0);
}
- SIMD_SVE2_GEMM_SAVE4(0);
- SIMD_SVE2_GEMM_SAVE4(1);
- SIMD_SVE2_GEMM_SAVE4(2);
- SIMD_SVE2_GEMM_SAVE4(3);
- SIMD_SVE2_GEMM_SAVE4(4);
- SIMD_SVE2_GEMM_SAVE4(5);
+ Add4ExtractedSums(c00, c01, c02, c03, alpha, C + 0 * ldc);
+ Add4ExtractedSums(c10, c11, c12, c13, alpha, C + 1 * ldc);
+ Add4ExtractedSums(c20, c21, c22, c23, alpha, C + 2 * ldc);
+ Add4ExtractedSums(c30, c31, c32, c33, alpha, C + 3 * ldc);
+ Add4ExtractedSums(c40, c41, c42, c43, alpha, C + 4 * ldc);
+ Add4ExtractedSums(c50, c51, c52, c53, alpha, C + 5 * ldc);
}
void Gemm32fNT(size_t M, size_t N, size_t K, const float* alpha, const float* A, size_t lda, const float* B, size_t ldb, const float* beta, float* C, size_t ldc)
@@ -383,21 +737,6 @@ namespace Simd
Kernel1x1nt, Kernel1x4nt, Kernel2x1nt, Kernel2x4nt, Kernel3x1nt, Kernel3x4nt, Kernel6x1nt, Kernel6x4nt);
gemmNT.Run(alpha, A, lda, B, ldb, beta, C, ldc);
}
-
-#undef SIMD_SVE2_GEMM_INIT1
-#undef SIMD_SVE2_GEMM_INIT4
-#undef SIMD_SVE2_GEMM_ROW1
-#undef SIMD_SVE2_GEMM_ROW4
-#undef SIMD_SVE2_GEMM_STEP1_1
-#undef SIMD_SVE2_GEMM_STEP1_2
-#undef SIMD_SVE2_GEMM_STEP1_3
-#undef SIMD_SVE2_GEMM_STEP1_6
-#undef SIMD_SVE2_GEMM_STEP4_1
-#undef SIMD_SVE2_GEMM_STEP4_2
-#undef SIMD_SVE2_GEMM_STEP4_3
-#undef SIMD_SVE2_GEMM_STEP4_6
-#undef SIMD_SVE2_GEMM_SAVE1
-#undef SIMD_SVE2_GEMM_SAVE4
}
#endif
}