From c8c2a6f8d105b23e13c83ea47bae6f2d0b8e6808 Mon Sep 17 00:00:00 2001 From: Hugo Meiland Date: Tue, 21 Jul 2026 02:27:38 +0200 Subject: [PATCH] x60: vectorize M1 int8 GEMV kernel on RVV Replace the pure-scalar SQ8BitGemmM1Kernel_CompInt8_ScaleFp16_Impl (used for <=3 remainder rows and batch=1 token generation) with an RVV kernel that vectorizes across the 16 columns of each repacked-B group. Inner loop uses a stride-8 int8 gather (vlse8) over the IME1 tile layout, widening vwmacc accumulation, then a fused vfmacc that matches GCC's FMA-contraction of the scalar 'acc += a_scale*b_scale*isum' so results are bit-exact with the scalar path. Verified on X60 (SpaceMiT): - standalone skeleton: 9/9 cases bit-exact (max_abs 0.0), 9.36x faster - llama-bench tg64 on qwen2.5-0.5b-q8_0: 0.82 -> 5.29 t/s (6.45x) - llama-cli generation coherent (IME path, use_ime1=1) --- ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp | 48 +++++++++++++-------- 1 file changed, 29 insertions(+), 19 deletions(-) diff --git a/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp b/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp index 9190a6629228..2be41e2424c0 100644 --- a/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp +++ b/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp @@ -1148,7 +1148,7 @@ void SQ8BitGemmM4Kernel_CompInt8_ScaleFp16_Impl(size_t BlkLen, } } -// v1 scalar M1 path for i8i8: only used for <=3 remainder rows and tg (baseline ran these on RVV). +// M1 path for i8i8: used for <=3 remainder rows and tg (batch=1). // Reads A = [4B f32 scale][BlkLen int8 natural k-order] per k-block, // B = tiled q8_0x16: per 16-col group per k-block [16xfp16 scale=32B][tiles (slice sl, group g), tile=4cols x 8k as [col][k]=32B]. void SQ8BitGemmM1Kernel_CompInt8_ScaleFp16_Impl(size_t BlkLen, @@ -1163,35 +1163,45 @@ void SQ8BitGemmM1Kernel_CompInt8_ScaleFp16_Impl(size_t BlkLen, const size_t group_stride = BlockCountK * kblk_stride; const size_t a_blk_stride = 4 + BlkLen; - for (size_t n = 0; n < CountN; ++n) { - const size_t group = n / 16; - const size_t c = n % 16; - const size_t g = c / 4; - const size_t cc = c % 4; - const uint8_t * gbase = QuantBData + group * group_stride; - float acc = 0.0f; + // Vectorize across the 16 columns of a group: for fixed k they sit at + // btile + (k/8)*128 + (k%8) + c*8 (stride-8 gather). fused vfmacc matches the + // scalar acc += a_scale*b_scale*isum bit-exactly (GCC FMA-contracts the scalar). + for (size_t n0 = 0; n0 < CountN; n0 += 16) { + const size_t ncols = (CountN - n0) < 16 ? (CountN - n0) : 16; + const size_t vl = ncols; + const uint8_t * gbase = QuantBData + (n0 / 16) * group_stride; + + vfloat32m2_t vacc = __riscv_vfmv_v_f_f32m2(0.0f, vl); for (size_t kb = 0; kb < BlockCountK; ++kb) { const uint8_t * bscale_ptr = gbase + kb * kblk_stride; - _Float16 bsh; - memcpy(&bsh, bscale_ptr + c * sizeof(_Float16), sizeof(_Float16)); - const float b_scale = (float) bsh; - const int8_t * btile = (const int8_t *) (bscale_ptr + 32); + const int8_t * btile = (const int8_t *) (bscale_ptr + 32); + + float bscratch[16]; + for (size_t c = 0; c < ncols; ++c) { + _Float16 bsh; + memcpy(&bsh, bscale_ptr + c * sizeof(_Float16), sizeof(_Float16)); + bscratch[c] = (float) bsh; + } + vfloat32m2_t vbs = __riscv_vle32_v_f32m2(bscratch, vl); float a_scale; memcpy(&a_scale, QuantA + kb * a_blk_stride, sizeof(float)); const int8_t * a_int8 = (const int8_t *) (QuantA + kb * a_blk_stride + 4); - int32_t isum = 0; + vint32m2_t visum = __riscv_vmv_v_x_i32m2(0, vl); for (size_t k = 0; k < BlkLen; ++k) { - const size_t sl = k / 8; - const size_t kk = k % 8; - const size_t t = sl * 4 + g; - isum += (int32_t) a_int8[k] * (int32_t) btile[t * 32 + cc * 8 + kk]; + const int8_t * base = btile + (k / 8) * 128 + (k % 8); + vint8mf2_t vb8 = __riscv_vlse8_v_i8mf2(base, 8, vl); + vint16m1_t vb16 = __riscv_vsext_vf2_i16m1(vb8, vl); + visum = __riscv_vwmacc_vx_i32m2(visum, (int16_t) a_int8[k], vb16, vl); } - acc += a_scale * b_scale * (float) isum; + + vfloat32m2_t vf = __riscv_vfcvt_f_x_v_f32m2(visum, vl); + vfloat32m2_t vab = __riscv_vfmul_vf_f32m2(vbs, a_scale, vl); + vacc = __riscv_vfmacc_vv_f32m2(vacc, vab, vf, vl); } - C[n] = acc; + __riscv_vse32_v_f32m2(C + n0, vacc, vl); } } } // namespace