66#include "common.h"
77#include <stdlib.h>
88#include <inttypes.h>
9- #include <math.h>
109
1110#if defined(DYNAMIC_ARCH )
1211#define COMBINE (a ,b ) a ## b
@@ -30,6 +29,7 @@ extern void SGEMM_PREPROCESS(uint64_t nbr, uint64_t nbc,\
3029 const float * restrict a , float * a_mod ) ;
3130
3231/* Function Definitions */
32+ #if !defined(TRANSA )
3333static uint64_t sve_cntw () {
3434 uint64_t cnt ;
3535 asm volatile (
@@ -39,18 +39,21 @@ static uint64_t sve_cntw() {
3939 );
4040 return cnt ;
4141}
42+ #endif
4243
4344#if defined(__ARM_FEATURE_SME ) && defined(__ARM_FEATURE_LOCALLY_STREAMING ) && defined(__clang__ ) && __clang_major__ >= 16
4445// Outer product kernel.
4546// Computes a 2SVL x 2SVL block of C, utilizing all four FP32 tiles of ZA.
4647__attribute__((always_inline )) inline void
47- kernel_2x2 (const float * A , float * B_T , const float * B , float * A_T , float * C , size_t shared_dim ,
48+ kernel_2x2 (const float * A , const float * B_T , const float * B , const float * A_T , float * C , size_t shared_dim ,
4849 size_t ldc , size_t block_rows , size_t block_cols , float alpha ,
4950 float beta , uint64_t row_idx , uint64_t col_idx )
5051 __arm_out ("za ") __arm_streaming {
5152
5253 const uint64_t svl = svcntw ();
54+ #if defined(TRANSA )
5355 size_t ldb = ldc ;
56+ #endif
5457 // Predicate set-up
5558 svbool_t pg = svptrue_b32 ();
5659 svbool_t pg_a_0 = svwhilelt_b32_u64 (0 , block_rows );
@@ -63,26 +66,29 @@ kernel_2x2(const float *A, float *B_T, const float *B, float *A_T, float *C, siz
6366#define pg_c_1 pg_b_1
6467
6568 svzero_za ();
66- svfloat32_t beta_vec = svdup_f32 (beta );
67-
68- // Load C to ZA
69- for (size_t i = 0 ; i < MIN (svl , block_rows ); i ++ ) {
70- svfloat32_t row_c_0 = svld1 (pg_c_0 , & C [i * ldc ]);
71- row_c_0 = svmul_x (pg , beta_vec , row_c_0 );
72- svwrite_hor_za32_f32_m (/*tile*/ 0 , /*slice*/ i , pg_c_0 , row_c_0 );
73-
74- svfloat32_t row_c_1 = svld1 (pg_c_1 , & C [i * ldc + svl ]);
75- row_c_1 = svmul_x (pg , beta_vec , row_c_1 );
76- svwrite_hor_za32_f32_m (/*tile*/ 1 , /*slice*/ i , pg_c_1 , row_c_1 );
77- }
78- for (size_t i = svl ; i < block_rows ; i ++ ) {
79- svfloat32_t row_c_0 = svld1 (pg_c_0 , & C [i * ldc ]);
80- row_c_0 = svmul_x (pg , beta_vec , row_c_0 );
81- svwrite_hor_za32_f32_m (/*tile*/ 2 , /*slice*/ i , pg_c_0 , row_c_0 );
82-
83- svfloat32_t row_c_1 = svld1 (pg_c_1 , & C [i * ldc + svl ]);
84- row_c_1 = svmul_x (pg , beta_vec , row_c_1 );
85- svwrite_hor_za32_f32_m (/*tile*/ 3 , /*slice*/ i , pg_c_1 , row_c_1 );
69+ // beta == 0 must not read C; ZA is already initialized to zero.
70+ if (beta != 0.0f ) {
71+ svfloat32_t beta_vec = svdup_f32 (beta );
72+
73+ // Load C to ZA
74+ for (size_t i = 0 ; i < MIN (svl , block_rows ); i ++ ) {
75+ svfloat32_t row_c_0 = svld1 (pg_c_0 , & C [i * ldc ]);
76+ row_c_0 = svmul_x (pg , beta_vec , row_c_0 );
77+ svwrite_hor_za32_f32_m (/*tile*/ 0 , /*slice*/ i , pg_c_0 , row_c_0 );
78+
79+ svfloat32_t row_c_1 = svld1 (pg_c_1 , & C [i * ldc + svl ]);
80+ row_c_1 = svmul_x (pg , beta_vec , row_c_1 );
81+ svwrite_hor_za32_f32_m (/*tile*/ 1 , /*slice*/ i , pg_c_1 , row_c_1 );
82+ }
83+ for (size_t i = svl ; i < block_rows ; i ++ ) {
84+ svfloat32_t row_c_0 = svld1 (pg_c_0 , & C [i * ldc ]);
85+ row_c_0 = svmul_x (pg , beta_vec , row_c_0 );
86+ svwrite_hor_za32_f32_m (/*tile*/ 2 , /*slice*/ i , pg_c_0 , row_c_0 );
87+
88+ svfloat32_t row_c_1 = svld1 (pg_c_1 , & C [i * ldc + svl ]);
89+ row_c_1 = svmul_x (pg , beta_vec , row_c_1 );
90+ svwrite_hor_za32_f32_m (/*tile*/ 3 , /*slice*/ i , pg_c_1 , row_c_1 );
91+ }
8692 }
8793
8894 svfloat32_t alpha_vec = svdup_f32 (alpha );
@@ -250,12 +256,19 @@ void CNAME (BLASLONG N, BLASLONG K, float alpha, float * __restrict A, \
250256 BLASLONG strideA , float * __restrict B , BLASLONG strideB , \
251257 float beta , float * __restrict R , BLASLONG strideR )
252258{
259+ if (alpha == 0.0f || K == 0 ) {
260+ if (beta == 1.0f )
261+ return ;
262+ ssyr2k_direct_sme1_2VLx2VL (N , 0 , & alpha , A , B , & beta , R );
263+ return ;
264+ }
265+
253266#if !defined(TRANSA )
254267 uint64_t n_mod , vl_elms ;
255268
256269 vl_elms = sve_cntw ();
257270
258- n_mod = ceil (( double ) N /( double ) vl_elms ) * vl_elms ;
271+ n_mod = ((( uint64_t ) N + vl_elms - 1 ) / vl_elms ) * vl_elms ;
259272
260273 float * A_mod = (float * ) malloc (n_mod * K * sizeof (float ));
261274 float * B_mod = (float * ) malloc (n_mod * K * sizeof (float ));
0 commit comments