From 687076430efad40f02d18e1436f0611c80450b20 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 ZGEMM The ZGEMM 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, plus the temporaries' live range and vector-register pressure, 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: zgemm +19.7% at 256, +16.4% at 512, +16.2% 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/zgemm_kernel_8x4_zvl256b.c | 109 +++++++++++----------- 1 file changed, 57 insertions(+), 52 deletions(-) diff --git a/kernel/riscv64/zgemm_kernel_8x4_zvl256b.c b/kernel/riscv64/zgemm_kernel_8x4_zvl256b.c index ca33368f00..57fc382389 100644 --- a/kernel/riscv64/zgemm_kernel_8x4_zvl256b.c +++ b/kernel/riscv64/zgemm_kernel_8x4_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) @@ -172,60 +180,57 @@ int CNAME(BLASLONG M, BLASLONG N, BLASLONG K, FLOAT alphar, FLOAT alphai, FLOAT* B3i = B[bi+3*2+1]; bi += 4*2; - A0r = __riscv_vlse64_v_f64m1( &A[ai+0*gvl*2], sizeof(FLOAT)*2, gvl ); - A0i = __riscv_vlse64_v_f64m1( &A[ai+0*gvl*2+1], sizeof(FLOAT)*2, gvl ); - A1r = __riscv_vlse64_v_f64m1( &A[ai+1*gvl*2], sizeof(FLOAT)*2, gvl ); - A1i = __riscv_vlse64_v_f64m1( &A[ai+1*gvl*2+1], sizeof(FLOAT)*2, gvl ); + { + vfloat64m1x2_t Aseg0 = __riscv_vlseg2e64_v_f64m1x2( &A[ai+0*gvl*2], gvl ); + A0r = __riscv_vget_v_f64m1x2_f64m1( Aseg0, 0 ); + A0i = __riscv_vget_v_f64m1x2_f64m1( Aseg0, 1 ); + vfloat64m1x2_t Aseg1 = __riscv_vlseg2e64_v_f64m1x2( &A[ai+1*gvl*2], gvl ); + A1r = __riscv_vget_v_f64m1x2_f64m1( Aseg1, 0 ); + A1i = __riscv_vget_v_f64m1x2_f64m1( Aseg1, 1 ); + } ai += 8*2; - tmp0r = __riscv_vfmul_vf_f64m1( A0i, B0i, gvl); - tmp0i = __riscv_vfmul_vf_f64m1( A0r, B0i, gvl); - tmp1r = __riscv_vfmul_vf_f64m1( A1i, B0i, gvl); - tmp1i = __riscv_vfmul_vf_f64m1( A1r, B0i, gvl); - tmp2r = __riscv_vfmul_vf_f64m1( A0i, B1i, gvl); - tmp2i = __riscv_vfmul_vf_f64m1( A0r, B1i, gvl); - tmp3r = __riscv_vfmul_vf_f64m1( A1i, B1i, gvl); - tmp3i = __riscv_vfmul_vf_f64m1( A1r, B1i, gvl); - tmp0r = VFMACC_RR( tmp0r, B0r, A0r, gvl); - tmp0i = VFMACC_RI( tmp0i, B0r, A0i, gvl); - tmp1r = VFMACC_RR( tmp1r, B0r, A1r, gvl); - tmp1i = VFMACC_RI( tmp1i, B0r, A1i, gvl); - tmp2r = VFMACC_RR( tmp2r, B1r, A0r, gvl); - tmp2i = VFMACC_RI( tmp2i, B1r, A0i, gvl); - tmp3r = VFMACC_RR( tmp3r, B1r, A1r, gvl); - tmp3i = VFMACC_RI( tmp3i, B1r, A1i, 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_f64m1( A0i, B2i, gvl); - tmp0i = __riscv_vfmul_vf_f64m1( A0r, B2i, gvl); - tmp1r = __riscv_vfmul_vf_f64m1( A1i, B2i, gvl); - tmp1i = __riscv_vfmul_vf_f64m1( A1r, B2i, gvl); - tmp2r = __riscv_vfmul_vf_f64m1( A0i, B3i, gvl); - tmp2i = __riscv_vfmul_vf_f64m1( A0r, B3i, gvl); - tmp3r = __riscv_vfmul_vf_f64m1( A1i, B3i, gvl); - tmp3i = __riscv_vfmul_vf_f64m1( A1r, B3i, gvl); - tmp0r = VFMACC_RR( tmp0r, B2r, A0r, gvl); - tmp0i = VFMACC_RI( tmp0i, B2r, A0i, gvl); - tmp1r = VFMACC_RR( tmp1r, B2r, A1r, gvl); - tmp1i = VFMACC_RI( tmp1i, B2r, A1i, gvl); - tmp2r = VFMACC_RR( tmp2r, B3r, A0r, gvl); - tmp2i = VFMACC_RI( tmp2i, B3r, A0i, gvl); - tmp3r = VFMACC_RR( tmp3r, B3r, A1r, gvl); - tmp3i = VFMACC_RI( tmp3i, B3r, A1i, 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); + /* Fold the complex MAC into the accumulators: two signed FMAs + * per accumulator replace the tmp materialisation + * (vfmul + VFMACC_RR/RI) followed by vfadd. Applying the + * same op twice composes 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), so the k-step value + * added to each accumulator is identical to the original + * tmp += vfadd sequence, and the per-column accumulation + * order is unchanged. */ + 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, B0i, A1i, gvl); + ACC1r = VFMACC_RR( ACC1r, B0r, A1r, gvl); + ACC1i = FOLD_OPA( ACC1i, B0r, A1i, gvl); + ACC1i = FOLD_OPB( ACC1i, B0i, A1r, gvl); + ACC2r = VFMACC_RR( ACC2r, B1i, A0i, gvl); + ACC2r = VFMACC_RR( ACC2r, B1r, A0r, gvl); + ACC2i = FOLD_OPA( ACC2i, B1r, A0i, gvl); + ACC2i = FOLD_OPB( ACC2i, B1i, A0r, gvl); + ACC3r = VFMACC_RR( ACC3r, B1i, A1i, gvl); + ACC3r = VFMACC_RR( ACC3r, B1r, A1r, gvl); + ACC3i = FOLD_OPA( ACC3i, B1r, A1i, gvl); + ACC3i = FOLD_OPB( ACC3i, B1i, A1r, gvl); + ACC4r = VFMACC_RR( ACC4r, B2i, A0i, gvl); + ACC4r = VFMACC_RR( ACC4r, B2r, A0r, gvl); + ACC4i = FOLD_OPA( ACC4i, B2r, A0i, gvl); + ACC4i = FOLD_OPB( ACC4i, B2i, A0r, gvl); + ACC5r = VFMACC_RR( ACC5r, B2i, A1i, gvl); + ACC5r = VFMACC_RR( ACC5r, B2r, A1r, gvl); + ACC5i = FOLD_OPA( ACC5i, B2r, A1i, gvl); + ACC5i = FOLD_OPB( ACC5i, B2i, A1r, gvl); + ACC6r = VFMACC_RR( ACC6r, B3i, A0i, gvl); + ACC6r = VFMACC_RR( ACC6r, B3r, A0r, gvl); + ACC6i = FOLD_OPA( ACC6i, B3r, A0i, gvl); + ACC6i = FOLD_OPB( ACC6i, B3i, A0r, gvl); + ACC7r = VFMACC_RR( ACC7r, B3i, A1i, gvl); + ACC7r = VFMACC_RR( ACC7r, B3r, A1r, gvl); + ACC7i = FOLD_OPA( ACC7i, B3r, A1i, gvl); + ACC7i = FOLD_OPB( ACC7i, B3i, A1r, gvl); }