Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 29 additions & 19 deletions ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand Down