From 42194def12378e6db343849dc31adfa51f2c83d6 Mon Sep 17 00:00:00 2001 From: co-seven Date: Mon, 14 Sep 2026 09:46:43 +0000 Subject: [PATCH] ggml-spacemit: type-gate direct GEMV to q4_0, fuse GDN decode row update MUL_MAT direct path (IME reads B straight from DRAM, decode only): restrict to q4_0 and drop the q8_0 arm together with its gemm_n <= 2*NB_COLS*64 crossover. Paired A/B on K3 shows q8_0 has no winning shape: Qwen3.5-35B dense q8_0 decode GEMVs all fall under the old crossover and still measure 21.8-26.7 GB/s direct vs 28.9-30.8 staged through TCM (tg128 6.71 -> 6.91 with direct disabled), while q4_0 gains from direct (Qwen3-0.6B tg128 57.45 -> 59.05, +2.8%; pp unchanged). 35B carries no physical q4_0 weights after repack, so the gate leaves it on the staged path entirely. GDN decode: add a fused single-pass row update used for n_tokens <= 4. State rows are independent within a token, so the reference's memcpy + decay/dots pass + update pass (3 reads, 3 writes of the state per head) collapse into one load/store pair per row, with k/q/decay held in vector registers across the row loop; decode reads the input state and writes the output state directly with no intermediate copy. Kernel time at n_tokens == 1 drops 424.6 -> 216.0 us; paired A/B tg128 7.58 -> 7.93 (+4.6%), pp128 unchanged. Chunks with more tokens keep the reference two-pass form: their state working copy is L2-resident so the extra passes already run at the bandwidth floor, and the fused form loses there to register spills and the store-at-row-end reduction dependency stalling row loads on the in-order cores (measured crossover ~8 tokens; chunks are 64, or 16 with kda). Co-Authored-By: Claude Code --- ggml/src/ggml-spacemit/ime.cpp | 16 +-- ggml/src/ggml-spacemit/rvv_kernels.cpp | 131 ++++++++++++++++++++----- 2 files changed, 110 insertions(+), 37 deletions(-) diff --git a/ggml/src/ggml-spacemit/ime.cpp b/ggml/src/ggml-spacemit/ime.cpp index 86bfce3fc2b..2440d51a013 100644 --- a/ggml/src/ggml-spacemit/ime.cpp +++ b/ggml/src/ggml-spacemit/ime.cpp @@ -418,18 +418,12 @@ class tensor_traits : public ggml::spacemit::tensor_traits_base { } uint8_t * b_col_zp = block_type_has_zp() ? b_col : nullptr; - if constexpr ((std::is_same_v && INTER_SIZE == 256 && NB_COLS == 32) || - (std::is_same_v && INTER_SIZE == 32 && NB_COLS == 32)) { - // Wide Q8 output heads amortize the staged copy and stream more efficiently through TCM. - constexpr int64_t max_direct_q8_tiles = 64; - const bool direct_q8_supported = !std::is_same_v || - gemm_n <= 2 * NB_COLS * max_direct_q8_tiles; - if (gemm_m == 1 && a_row != quant_a_buffer && direct_q8_supported) { + // Direct IME-from-DRAM GEMV path, decode only (gemm_m == 1), q4_0 only: + // q8_0 always loses to the staged TCM path here. + if constexpr (std::is_same_v && INTER_SIZE == 256 && NB_COLS == 32) { + if (gemm_m == 1 && a_row != quant_a_buffer) { spacemit_kernels::rvv::memcpy1d(a_row, quant_a_buffer, gemm_workspace_size); - int64_t tile_cols = 2 * NB_COLS; - if constexpr (std::is_same_v) { - tile_cols = 4 * NB_COLS; - } + constexpr int64_t tile_cols = 4 * NB_COLS; for (int64_t ni = (int64_t) ith * tile_cols; ni < gemm_n; ni += tile_cols * nth) { const int64_t nb_real = std::min(gemm_n - ni, tile_cols); uint8_t * b_row = reinterpret_cast(w_data) + ni * row_stride_b; diff --git a/ggml/src/ggml-spacemit/rvv_kernels.cpp b/ggml/src/ggml-spacemit/rvv_kernels.cpp index 520bc2a0299..cbbe4773238 100644 --- a/ggml/src/ggml-spacemit/rvv_kernels.cpp +++ b/ggml/src/ggml-spacemit/rvv_kernels.cpp @@ -3172,10 +3172,35 @@ template void forward_concat(ggml::spacemit::context & ctx, ggml_te } } -struct rvv_gdn_decay_dots_f32 { - float state_k; - float state_q; -}; +// Fused single-pass row update: decays the state row, derives the delta-rule +// correction from the decayed row, applies it and writes the updated row with +// one load/store pair. Each state row is independent within a token, so the +// separate decay-store / update-reload passes of the reference implementation +// collapse into this. `src` and `dst` may only alias when they are the same +// row (the row is loaded before it is stored). Returns the attention output. +static inline float rvv_gdn_fused_row_f32(float * dst, + const float * src, + const vfloat32m8_t kv, + const vfloat32m8_t qv, + const vfloat32m8_t dv, + const float v_j, + const float beta, + const float kq, + const float scale, + const size_t vl) { + vfloat32m8_t s = __riscv_vle32_v_f32m8(src, vl); + s = __riscv_vfmul_vv_f32m8(s, dv, vl); + vfloat32m1_t sum_k = __riscv_vfmv_v_f_f32m1(0.0f, 1); + vfloat32m1_t sum_q = __riscv_vfmv_v_f_f32m1(0.0f, 1); + sum_k = __riscv_vfredusum_vs_f32m8_f32m1(__riscv_vfmul_vv_f32m8(s, kv, vl), sum_k, vl); + sum_q = __riscv_vfredusum_vs_f32m8_f32m1(__riscv_vfmul_vv_f32m8(s, qv, vl), sum_q, vl); + const float state_k = __riscv_vfmv_f_s_f32m1_f32(sum_k); + const float state_q = __riscv_vfmv_f_s_f32m1_f32(sum_q); + const float delta_j = (v_j - state_k) * beta; + s = __riscv_vfmacc_vf_f32m8(s, delta_j, kv, vl); + __riscv_vse32_v_f32m8(dst, s, vl); + return (state_q + delta_j * kq) * scale; +} static inline float rvv_gdn_dot_f32(const float * x, const float * y, int64_t n) { const size_t vl = __riscv_vsetvl_e32m8(n); @@ -3187,6 +3212,18 @@ static inline float rvv_gdn_dot_f32(const float * x, const float * y, int64_t n) return __riscv_vfmv_f_s_f32m1_f32(sum); } +struct rvv_gdn_decay_dots_f32 { + float state_k; + float state_q; +}; + +// Two-pass row update of the reference implementation: decay the state row, +// store it back and reduce the two dot products in the same sweep, then apply +// the delta-rule correction in a second sweep over the rows. The early store +// keeps the loads of subsequent rows pipelined on the in-order cores, which +// makes this form the faster one once the chunk has enough tokens that the +// state stays bandwidth-bound (measured crossover vs the fused form: ~8 +// tokens/chunk; prefill chunks are 64, kda chunks 16). static inline rvv_gdn_decay_dots_f32 rvv_gdn_decay_dots_f32_impl(float * state, const float * decay, float decay_scalar, @@ -3250,13 +3287,28 @@ void forward_gated_delta_net(ggml::spacemit::context & ctx, ggml_tensor * op) { const float scale = 1.0f / sqrtf((float) S_v); const bool kda = src_g->ne[0] == S_v; - // KDA needs the per-column decay vector for every state row, so keep it - // separate from the row-wise delta values that are produced in-place. - std::vector scratch((size_t) (3 * S_v + (K > 1 ? S_v * S_v : 0))); - float * delta = scratch.data(); - float * decay = delta + S_v; - float * state_q = decay + S_v; - float * state_work = K > 1 ? state_q + S_v : nullptr; + // Decode (n_tokens == 1, any n_seqs) and prefill tail chunks up to 4 tokens + // run the fused single-pass row update: the state is read/written once per + // token instead of once per phase, which halves the kernel time at + // n_tokens == 1. Larger chunks keep the reference two-pass form — its early + // state store keeps row loads pipelined, and with the chunk state L2 + // resident it already runs at the DRAM floor, so the fused form cannot win + // there (measured crossover ~8 tokens; chunks are 64, or 16 with kda). + const bool fused = n_tokens <= 4; + // KDA needs the per-column decay vector for every state row; the fused + // row update consumes it as a vector, so only this small buffer remains + // of the reference implementation's scratch space. + float kda_decay[512]; + GGML_ASSERT(S_v <= (int64_t) (sizeof(kda_decay) / sizeof(kda_decay[0]))); + // K > 1 (chunked prefill snapshots) evolves a private working copy. + std::vector state_work((fused && K > 1) ? (size_t) (S_v * S_v) : 0); + // Two-pass scratch: per-row delta and attention sums, plus its own K > 1 + // working copy of the state. + std::vector scratch(!fused ? (size_t) (3 * S_v + (K > 1 ? S_v * S_v : 0)) : 0); + float * delta = !fused ? scratch.data() : nullptr; + float * decay = !fused ? delta + S_v : nullptr; + float * state_q = !fused ? decay + S_v : nullptr; + float * state_work_tp = (!fused && K > 1) ? state_q + S_v : nullptr; const int64_t rq3 = src_v->ne[3] / src_q->ne[3]; const int64_t rk3 = src_v->ne[3] / src_k->ne[3]; @@ -3269,9 +3321,17 @@ void forward_gated_delta_net(ggml::spacemit::context & ctx, ggml_tensor * op) { const int64_t iq3 = iv3 / rq3; const int64_t ik3 = iv3 / rk3; - float * s_out = K > 1 ? state_work : state_out_base + (iv3 * H + iv1) * S_v * S_v; + float * s_out = K > 1 ? (fused ? state_work.data() : state_work_tp) + : state_out_base + (iv3 * H + iv1) * S_v * S_v; const float * s_in = state_in_base + iv3 * state_seq_stride + iv1 * S_v * S_v; - memcpy(s_out, s_in, (size_t) (S_v * S_v) * sizeof(float)); + if (!fused) { + memcpy(s_out, s_in, (size_t) (S_v * S_v) * sizeof(float)); + } + // In the fused form the first token is fused with the state copy: it + // reads the input state row and writes the updated row straight to the + // output, so decode (n_tokens == 1) never materialises an intermediate + // copy. + const float * s_cur = s_in; float * attn_data = attn_out_base + (iv3 * n_tokens * H + iv1) * S_v; for (int64_t t = 0; t < n_tokens; ++t) { @@ -3288,21 +3348,40 @@ void forward_gated_delta_net(ggml::spacemit::context & ctx, ggml_tensor * op) { const float kq = rvv_gdn_dot_f32(k, q, S_v); - float decay_scalar = 0.0f; - if (kda) { - for (int64_t i = 0; i < S_v; ++i) decay[i] = expf(g[i]); + if (fused) { + const size_t vl = __riscv_vsetvl_e32m8(S_v); + const vfloat32m8_t kv = __riscv_vle32_v_f32m8(k, vl); + const vfloat32m8_t qv = __riscv_vle32_v_f32m8(q, vl); + vfloat32m8_t dv; + if (kda) { + for (int64_t i = 0; i < S_v; ++i) kda_decay[i] = expf(g[i]); + dv = __riscv_vle32_v_f32m8(kda_decay, vl); + } else { + dv = __riscv_vfmv_v_f_f32m8(expf(g[0]), vl); + } + + for (int64_t j = 0; j < S_v; ++j) { + attn_data[j] = rvv_gdn_fused_row_f32(s_out + j * S_v, s_cur + j * S_v, + kv, qv, dv, v[j], beta, kq, scale, vl); + } + s_cur = s_out; } else { - decay_scalar = expf(g[0]); - } + float decay_scalar = 0.0f; + if (kda) { + for (int64_t i = 0; i < S_v; ++i) decay[i] = expf(g[i]); + } else { + decay_scalar = expf(g[0]); + } - for (int64_t j = 0; j < S_v; ++j) { - const auto dots = rvv_gdn_decay_dots_f32_impl(s_out + j * S_v, decay, decay_scalar, k, q, S_v, kda); - delta[j] = (v[j] - dots.state_k) * beta; - state_q[j] = dots.state_q; - } - for (int64_t j = 0; j < S_v; ++j) { - rvv_gdn_update_f32(s_out + j * S_v, k, delta[j], S_v); - attn_data[j] = (state_q[j] + delta[j] * kq) * scale; + for (int64_t j = 0; j < S_v; ++j) { + const auto dots = rvv_gdn_decay_dots_f32_impl(s_out + j * S_v, decay, decay_scalar, k, q, S_v, kda); + delta[j] = (v[j] - dots.state_k) * beta; + state_q[j] = dots.state_q; + } + for (int64_t j = 0; j < S_v; ++j) { + rvv_gdn_update_f32(s_out + j * S_v, k, delta[j], S_v); + attn_data[j] = (state_q[j] + delta[j] * kq) * scale; + } } attn_data += S_v * H;