Skip to content
Closed
Show file tree
Hide file tree
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
1 change: 0 additions & 1 deletion ggml/src/ggml-cpu/arch-fallback.h
Original file line number Diff line number Diff line change
Expand Up @@ -203,7 +203,6 @@
#elif defined(__riscv)
// quants.c
#define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0
#define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0
// repack.cpp
#define ggml_quantize_mat_q8_0_4x1_generic ggml_quantize_mat_q8_0_4x1
#define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4
Expand Down
49 changes: 49 additions & 0 deletions ggml/src/ggml-cpu/arch/riscv/quants.c
Original file line number Diff line number Diff line change
Expand Up @@ -274,6 +274,55 @@ void ggml_vec_dot_q4_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const voi
#endif
}

void ggml_vec_dot_q1_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) {
#if defined(__riscv_v)
const int qk = QK1_0;
const int nb = n / qk;
// A single q8_0 block always contains 32 int8 values. With e8,m4 this
// maps to vl=32 on any spec-compliant RVV target (minimum VLEN is 64-bit).
const size_t vl = __riscv_vsetvl_e8m4(32);

assert(n % qk == 0);
assert(vl == 32);
assert(nrc == 1);
UNUSED(nrc);
UNUSED(bx);
UNUSED(by);
UNUSED(bs);

const block_q1_0 * GGML_RESTRICT x = vx;
const block_q8_0 * GGML_RESTRICT y = vy;
const vint16m1_t v_zero_sum = __riscv_vmv_v_x_i16m1(0, 1);

float sumf = 0.0f;

for (int ib = 0; ib < nb; ++ib) {
const float d0 = GGML_CPU_FP16_TO_FP32(x[ib].d);
float sumi = 0.0f;

for (int k = 0; k < 4; ++k) {
const block_q8_0 * GGML_RESTRICT yb = &y[ib * 4 + k];
const float d1 = GGML_CPU_FP16_TO_FP32(yb->d);
const uint8_t * GGML_RESTRICT bits = &x[ib].qs[k * 4];
vbool2_t m_positive = __riscv_vlm_v_b2(bits, vl);
vint8m4_t v_q8 = __riscv_vle8_v_i8m4(yb->qs, vl);
vint8m4_t v_q8_neg = __riscv_vsub_vv_i8m4(__riscv_vmv_v_x_i8m4(0, vl), v_q8, vl);
vint8m4_t v_q8_signed = __riscv_vmerge_vvm_i8m4(v_q8_neg, v_q8, m_positive, vl);
vint16m1_t v_sum = __riscv_vwredsum_vs_i8m4_i16m1(v_q8_signed, v_zero_sum, vl);
const int sumi_block = __riscv_vmv_x_s_i16m1_i16(v_sum);

sumi += d1 * sumi_block;
}

sumf += d0 * sumi;
}

*s = sumf;
#else
ggml_vec_dot_q1_0_q8_0_generic(n, s, bs, vx, bx, vy, by, nrc);
#endif
}

void ggml_vec_dot_q4_1_q8_1(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) {
#if defined(__riscv_v)
const int qk = QK8_1;
Expand Down
Loading