From 928321246807c6fd8bb525e3c4d0195c1e82a676 Mon Sep 17 00:00:00 2001 From: 6eanut Date: Thu, 17 Sep 2026 01:51:55 +0000 Subject: [PATCH] kernel/riscv64: fold the complex MAC into the accumulators in CGEMM The CGEMM micro-kernel built every complex product into a temporary and then added it to the accumulator: tmp = vfmul(Ai, Bi); tmp = VFMACC(tmp, Br, Ar); ACC = vfadd(ACC, tmp) Three vector instructions per accumulator and k-step. The two FMAs write the same destination, so applying them straight to ACC yields the same per-k-step sum for two instructions instead: ACCr = VFMACC_RR( ACCr, Bi, Ai); ACCr = VFMACC_RR( ACCr, Br, Ar); ACCi = FOLD_OPA( ACCi, Br, Ai); ACCi = FOLD_OPB( ACCi, Bi, Ar); VFMACC_RR/VFMACC_RI already encode the +/- conjugation of the variant being compiled, but the imaginary fold needs an op sequence that has no existing macro, so it is spelled out per variant as FOLD_OPA/FOLD_OPB next to them. Two FMAs compose to a plain accumulation because vfmsac(x,u,v)=u*v-x and vfmsac(vfmsac(a,u,v),w,y)=a+w*y-u*v, and likewise for vfmacc/vfnmsac. The value added per k-step is therefore unchanged and the per-column accumulation order is preserved. Only the FP rounding order differs, by a few ULP; measured against an exact double reference the patched kernel has the same error as the unpatched one. A packed A element has its real and imaginary halves adjacent, so one strided segment-2 load replaces the two separate strided loads. Measured on a SpaceMiT X100 (2.2 GHz, single thread, ZVL256B), median of three runs: cblas_cgemm +44.3% at 256, +45.0% at 512, +42.3% at 1024. Tests: make tests returns 0; 125/125 utest and 1473/1473 extension tests pass; the CBLAS L1/L2/L3 suites show no new failures. Co-authored-by: Yuansheng Co-authored-by: Ning Tian Signed-off-by: jiakai xu --- kernel/riscv64/cgemm_kernel_8x8_zvl256b.c | 95 +++++++++++------------ 1 file changed, 45 insertions(+), 50 deletions(-) diff --git a/kernel/riscv64/cgemm_kernel_8x8_zvl256b.c b/kernel/riscv64/cgemm_kernel_8x8_zvl256b.c index 7980c029a4..1b0499cc1a 100644 --- a/kernel/riscv64/cgemm_kernel_8x8_zvl256b.c +++ b/kernel/riscv64/cgemm_kernel_8x8_zvl256b.c @@ -49,6 +49,8 @@ AUTOGENERATED KERNEL #define S3 1 #define VFMACC_RR __riscv_vfmsac #define VFMACC_RI __riscv_vfmacc + #define FOLD_OPA __riscv_vfmacc + #define FOLD_OPB __riscv_vfmacc #endif #if defined(NR) || defined(NC) || defined(TR) || defined(TC) #define S0 1 @@ -57,6 +59,8 @@ AUTOGENERATED KERNEL #define S3 -1 #define VFMACC_RR __riscv_vfmacc #define VFMACC_RI __riscv_vfmsac + #define FOLD_OPA __riscv_vfmacc + #define FOLD_OPB __riscv_vfnmsac #endif #if defined(RN) || defined(RT) || defined(CN) || defined(CT) #define S0 1 @@ -65,6 +69,8 @@ AUTOGENERATED KERNEL #define S3 1 #define VFMACC_RR __riscv_vfmacc #define VFMACC_RI __riscv_vfnmsac + #define FOLD_OPA __riscv_vfnmsac + #define FOLD_OPB __riscv_vfmacc #endif #if defined(RR) || defined(RC) || defined(CR) || defined(CC) #define S0 1 @@ -73,6 +79,8 @@ AUTOGENERATED KERNEL #define S3 -1 #define VFMACC_RR __riscv_vfmsac #define VFMACC_RI __riscv_vfnmacc + #define FOLD_OPA __riscv_vfnmsac + #define FOLD_OPB __riscv_vfnmsac #endif int CNAME(BLASLONG M, BLASLONG N, BLASLONG K, FLOAT alphar, FLOAT alphai, FLOAT* A, FLOAT* B, FLOAT* C, BLASLONG ldc) @@ -186,58 +194,45 @@ int CNAME(BLASLONG M, BLASLONG N, BLASLONG K, FLOAT alphar, FLOAT alphai, FLOAT* B7i = B[bi+7*2+1]; bi += 8*2; - A0r = __riscv_vlse32_v_f32m1( &A[ai+0*gvl*2], sizeof(FLOAT)*2, gvl ); - A0i = __riscv_vlse32_v_f32m1( &A[ai+0*gvl*2+1], sizeof(FLOAT)*2, gvl ); + { + vfloat32m1x2_t Aseg = __riscv_vlseg2e32_v_f32m1x2( &A[ai+0*gvl*2], gvl ); + A0r = __riscv_vget_v_f32m1x2_f32m1( Aseg, 0 ); + A0i = __riscv_vget_v_f32m1x2_f32m1( Aseg, 1 ); + } ai += 8*2; - tmp0r = __riscv_vfmul_vf_f32m1( A0i, B0i, gvl); - tmp0i = __riscv_vfmul_vf_f32m1( A0r, B0i, gvl); - tmp1r = __riscv_vfmul_vf_f32m1( A0i, B1i, gvl); - tmp1i = __riscv_vfmul_vf_f32m1( A0r, B1i, gvl); - tmp2r = __riscv_vfmul_vf_f32m1( A0i, B2i, gvl); - tmp2i = __riscv_vfmul_vf_f32m1( A0r, B2i, gvl); - tmp3r = __riscv_vfmul_vf_f32m1( A0i, B3i, gvl); - tmp3i = __riscv_vfmul_vf_f32m1( A0r, B3i, gvl); - tmp0r = VFMACC_RR( tmp0r, B0r, A0r, gvl); - tmp0i = VFMACC_RI( tmp0i, B0r, A0i, gvl); - tmp1r = VFMACC_RR( tmp1r, B1r, A0r, gvl); - tmp1i = VFMACC_RI( tmp1i, B1r, A0i, gvl); - tmp2r = VFMACC_RR( tmp2r, B2r, A0r, gvl); - tmp2i = VFMACC_RI( tmp2i, B2r, A0i, gvl); - tmp3r = VFMACC_RR( tmp3r, B3r, A0r, gvl); - tmp3i = VFMACC_RI( tmp3i, B3r, A0i, gvl); - ACC0r = __riscv_vfadd( ACC0r, tmp0r, gvl); - ACC0i = __riscv_vfadd( ACC0i, tmp0i, gvl); - ACC1r = __riscv_vfadd( ACC1r, tmp1r, gvl); - ACC1i = __riscv_vfadd( ACC1i, tmp1i, gvl); - ACC2r = __riscv_vfadd( ACC2r, tmp2r, gvl); - ACC2i = __riscv_vfadd( ACC2i, tmp2i, gvl); - ACC3r = __riscv_vfadd( ACC3r, tmp3r, gvl); - ACC3i = __riscv_vfadd( ACC3i, tmp3i, gvl); - tmp0r = __riscv_vfmul_vf_f32m1( A0i, B4i, gvl); - tmp0i = __riscv_vfmul_vf_f32m1( A0r, B4i, gvl); - tmp1r = __riscv_vfmul_vf_f32m1( A0i, B5i, gvl); - tmp1i = __riscv_vfmul_vf_f32m1( A0r, B5i, gvl); - tmp2r = __riscv_vfmul_vf_f32m1( A0i, B6i, gvl); - tmp2i = __riscv_vfmul_vf_f32m1( A0r, B6i, gvl); - tmp3r = __riscv_vfmul_vf_f32m1( A0i, B7i, gvl); - tmp3i = __riscv_vfmul_vf_f32m1( A0r, B7i, gvl); - tmp0r = VFMACC_RR( tmp0r, B4r, A0r, gvl); - tmp0i = VFMACC_RI( tmp0i, B4r, A0i, gvl); - tmp1r = VFMACC_RR( tmp1r, B5r, A0r, gvl); - tmp1i = VFMACC_RI( tmp1i, B5r, A0i, gvl); - tmp2r = VFMACC_RR( tmp2r, B6r, A0r, gvl); - tmp2i = VFMACC_RI( tmp2i, B6r, A0i, gvl); - tmp3r = VFMACC_RR( tmp3r, B7r, A0r, gvl); - tmp3i = VFMACC_RI( tmp3i, B7r, A0i, gvl); - ACC4r = __riscv_vfadd( ACC4r, tmp0r, gvl); - ACC4i = __riscv_vfadd( ACC4i, tmp0i, gvl); - ACC5r = __riscv_vfadd( ACC5r, tmp1r, gvl); - ACC5i = __riscv_vfadd( ACC5i, tmp1i, gvl); - ACC6r = __riscv_vfadd( ACC6r, tmp2r, gvl); - ACC6i = __riscv_vfadd( ACC6i, tmp2i, gvl); - ACC7r = __riscv_vfadd( ACC7r, tmp3r, gvl); - ACC7i = __riscv_vfadd( ACC7i, tmp3i, gvl); + ACC0r = VFMACC_RR( ACC0r, B0i, A0i, gvl); + ACC0r = VFMACC_RR( ACC0r, B0r, A0r, gvl); + ACC0i = FOLD_OPA( ACC0i, B0r, A0i, gvl); + ACC0i = FOLD_OPB( ACC0i, B0i, A0r, gvl); + ACC1r = VFMACC_RR( ACC1r, B1i, A0i, gvl); + ACC1r = VFMACC_RR( ACC1r, B1r, A0r, gvl); + ACC1i = FOLD_OPA( ACC1i, B1r, A0i, gvl); + ACC1i = FOLD_OPB( ACC1i, B1i, A0r, gvl); + ACC2r = VFMACC_RR( ACC2r, B2i, A0i, gvl); + ACC2r = VFMACC_RR( ACC2r, B2r, A0r, gvl); + ACC2i = FOLD_OPA( ACC2i, B2r, A0i, gvl); + ACC2i = FOLD_OPB( ACC2i, B2i, A0r, gvl); + ACC3r = VFMACC_RR( ACC3r, B3i, A0i, gvl); + ACC3r = VFMACC_RR( ACC3r, B3r, A0r, gvl); + ACC3i = FOLD_OPA( ACC3i, B3r, A0i, gvl); + ACC3i = FOLD_OPB( ACC3i, B3i, A0r, gvl); + ACC4r = VFMACC_RR( ACC4r, B4i, A0i, gvl); + ACC4r = VFMACC_RR( ACC4r, B4r, A0r, gvl); + ACC4i = FOLD_OPA( ACC4i, B4r, A0i, gvl); + ACC4i = FOLD_OPB( ACC4i, B4i, A0r, gvl); + ACC5r = VFMACC_RR( ACC5r, B5i, A0i, gvl); + ACC5r = VFMACC_RR( ACC5r, B5r, A0r, gvl); + ACC5i = FOLD_OPA( ACC5i, B5r, A0i, gvl); + ACC5i = FOLD_OPB( ACC5i, B5i, A0r, gvl); + ACC6r = VFMACC_RR( ACC6r, B6i, A0i, gvl); + ACC6r = VFMACC_RR( ACC6r, B6r, A0r, gvl); + ACC6i = FOLD_OPA( ACC6i, B6r, A0i, gvl); + ACC6i = FOLD_OPB( ACC6i, B6i, A0r, gvl); + ACC7r = VFMACC_RR( ACC7r, B7i, A0i, gvl); + ACC7r = VFMACC_RR( ACC7r, B7r, A0r, gvl); + ACC7i = FOLD_OPA( ACC7i, B7r, A0i, gvl); + ACC7i = FOLD_OPB( ACC7i, B7i, A0r, gvl); }