Commit 50e3e3e48 for llama.cpp
commit 50e3e3e480c80b9ee0ea067e475602d540dc62c8
Author: Anav Prasad <anavp@nvidia.com>
Date: Fri Oct 9 16:23:28 2026 +0000
CUDA: Remove redundant CUDA copies after SSM_SCAN (#29807)
* CUDA: fuse copy of updated state snapshots into recurrent cache with ssm_scan
* CUDA: remove redundant cuda copies with K==1 (non spec-dec) scenario as well
diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu
index 8e0c7658e..957908f65 100644
--- a/ggml/src/ggml-cuda/ggml-cuda.cu
+++ b/ggml/src/ggml-cuda/ggml-cuda.cu
@@ -2899,6 +2899,79 @@ static int ggml_cuda_try_gdn_cache_fusion(
return skip;
}
+// match ssm_scan + the strided cpy that scatters its state snapshots into the cache, so the kernel writes them and skips the cpy
+static int ggml_cuda_try_ssm_scan_cache_fusion(
+ const ggml_cgraph * cgraph, int node_idx, ggml_cuda_ssm_scan_fused_cache & fused_state_cpy) {
+ const ggml_tensor * ssm = cgraph->nodes[node_idx];
+ // the kernel skips the snapshot tail, so the scan output must not be a graph output
+ if (ssm->op != GGML_OP_SSM_SCAN || ssm->type != GGML_TYPE_F32 || (ssm->flags & GGML_TENSOR_FLAG_OUTPUT)) {
+ return 0;
+ }
+
+ const int64_t K = ggml_get_op_params_i32(ssm, 0); // snapshot slot count
+
+ const ggml_tensor * s = ssm->src[0];
+ const ggml_tensor * x = ssm->src[1];
+ const ggml_tensor * A = ssm->src[3];
+
+ const int64_t d_state = s->ne[0];
+ const int64_t D = d_state * s->ne[1] * x->ne[1]; // d_state * head_dim * n_head
+ const int64_t n_tok = x->ne[2];
+ const int64_t n_seqs = x->ne[3];
+
+ // only the mamba-2 kernels (group scan and SSD) write to the cache; mamba-1 still uses the cpy
+ if (A->nb[1] != sizeof(float) || (d_state != 96 && d_state != 128 && d_state != 256)) {
+ return 0;
+ }
+
+ // the scan reads its input rows from the cache (picked by ids), so with more than one seq a seq can read a row that another seq writes in the same launch
+ if (n_seqs != 1) {
+ return 0;
+ }
+
+ const int64_t n_written = std::min<int64_t>(n_tok, K);
+ const size_t tail_off = ggml_row_size(GGML_TYPE_F32, ggml_nelements(x));
+
+ // snapshot cpy is the first real node after the scan (skip views/no-ops)
+ const ggml_tensor * cpy = nullptr;
+ int skip = 0;
+ for (int j = node_idx + 1; j < cgraph->n_nodes && cpy == nullptr; ++j) {
+ const ggml_tensor * n = cgraph->nodes[j];
+ if (ggml_cuda_is_view_or_noop(n)) {
+ continue;
+ }
+ if (n->op != GGML_OP_CPY || (n->flags & GGML_TENSOR_FLAG_OUTPUT)) {
+ return 0;
+ }
+ cpy = n;
+ skip = j - node_idx;
+ }
+ if (cpy == nullptr) {
+ return 0;
+ }
+
+ const ggml_tensor * src = cpy->src[0]; // view of the scan snapshot tail
+ const ggml_tensor * dst = cpy->src[1]; // cache view the kernel writes to
+
+ // src must be this scan's snapshot tail (contiguous, at the tail offset)
+ if (src->op != GGML_OP_VIEW || src->view_src != ssm || src->view_offs != tail_off ||
+ !ggml_is_contiguous(src)) {
+ return 0;
+ }
+
+ // dst is the [D, n_seqs, n_written] cache view; require nb[1] == D, the per-seq stride the kernel takes from src0->nb[3]
+ const std::array<int64_t, GGML_MAX_DIMS> expected_ne = { D, n_seqs, n_written, 1 };
+ if (dst->op != GGML_OP_VIEW || dst->type != GGML_TYPE_F32 || dst->data == nullptr ||
+ !std::equal(expected_ne.begin(), expected_ne.end(), dst->ne) ||
+ dst->nb[0] != ggml_type_size(GGML_TYPE_F32) || dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) {
+ return 0;
+ }
+
+ fused_state_cpy.data = (float *) dst->data; // rollback slot 0 (newest)
+ fused_state_cpy.slot_stride = K > 1 ? (int64_t) (dst->nb[2] / sizeof(float)) : 0;
+ return skip;
+}
+
static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int node_idx, ggml_cuda_topk_moe_args & args) {
args.sigmoid = false;
args.sqrt_softplus = false;
@@ -3585,6 +3658,20 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
}
}
+ // ssm_scan -> cpy: scatter recurrent-state snapshots into the cache
+ if (node->op == GGML_OP_SSM_SCAN) {
+ ggml_cuda_ssm_scan_fused_cache fused_state_cpy;
+ const int nodes_to_skip = ggml_cuda_try_ssm_scan_cache_fusion(cgraph, i, fused_state_cpy);
+ if (nodes_to_skip > 0) {
+#ifdef GGML_CUDA_DEBUG
+ GGML_LOG_INFO("%s: fused ssm_scan snapshot copies for %s (skipped %d nodes)\n",
+ __func__, node->name, nodes_to_skip);
+#endif
+ ggml_cuda_op_ssm_scan_fused_cache(*cuda_ctx, node, fused_state_cpy);
+ return nodes_to_skip;
+ }
+ }
+
//topk-moe
if (cgraph->nodes[i]->op == GGML_OP_UNARY || cgraph->nodes[i]->op == GGML_OP_SOFT_MAX ||
cgraph->nodes[i]->op == GGML_OP_ARGSORT) {
diff --git a/ggml/src/ggml-cuda/ssm-scan.cu b/ggml/src/ggml-cuda/ssm-scan.cu
index f6a41b407..dcc199f73 100644
--- a/ggml/src/ggml-cuda/ssm-scan.cu
+++ b/ggml/src/ggml-cuda/ssm-scan.cu
@@ -149,7 +149,7 @@ __global__ void __launch_bounds__(d_state, 1)
const int src0_nb2, const int src0_nb3, const int src1_nb2, const int src1_nb3,
const int src2_nb1, const int src2_nb2, const int src3_nb1,
const int src4_nb2, const int src4_nb3, const int src5_nb2, const int src5_nb3,
- const int64_t s_off, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
+ char * s_base, const int64_t s_slot_bytes, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
const float * GGML_CUDA_RESTRICT src0 = src0_ptr;
const float * GGML_CUDA_RESTRICT src1 = src1_ptr;
const float * GGML_CUDA_RESTRICT src2 = src2_ptr;
@@ -184,7 +184,7 @@ __global__ void __launch_bounds__(d_state, 1)
const float * B_warp = (const float *) ((const char *) src4 + (seq_idx * src4_nb3) + (group_off));
const float * C_warp = (const float *) ((const char *) src5 + (seq_idx * src5_nb3) + (group_off));
float * y_warp = dst + (seq_idx * n_tok * n_head * d_head) + warp_idx;
- float * s_warp = (float *) ((char *) dst + s_off + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
+ float * s_warp = (float *) (s_base + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
// strides across n_seq_tokens
const int stride_x = src1_nb2 / sizeof(float);
@@ -227,7 +227,7 @@ __global__ void __launch_bounds__(d_state, 1)
// Slot 0 is the final state written below; slots 1..K-1 are rollback snapshots.
const int64_t slot = n_tok - 1 - i;
if (K > 1 && slot > 0 && slot < K) {
- float * s_snapshot_warp = (float *) ((char *) dst + s_off + (slot * gridDim.y + seq_idx) * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
+ float * s_snapshot_warp = (float *) ((char *) s_warp + slot * s_slot_bytes);
#pragma unroll
for (int j = 0; j < c_factor; j++) {
s_snapshot_warp[WARP_SIZE * j + lane] = state[j];
@@ -248,7 +248,11 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
const int src2_nb2, const int src3_nb1, const int src4_nb2, const int src4_nb3, const int src5_nb2,
const int src5_nb3, const int64_t s_off, const int64_t d_state, const int64_t head_dim,
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq,
- const int64_t K, cudaStream_t stream) {
+ const int64_t K, const ggml_cuda_ssm_scan_fused_cache * cache, cudaStream_t stream) {
+ // when fused, the states go straight into the recurrent cache and the dst tail is left alone
+ char * const s_base = cache ? (char *) cache->data : (char *) dst + s_off;
+ const int64_t s_slot_bytes = cache ? cache->slot_stride * (int64_t) sizeof(float) : n_seq * (int64_t) src0_nb3;
+
// NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
if (src3_nb1 == sizeof(float)) {
// Mamba-2
@@ -261,7 +265,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
ggml_cuda_kernel_launch(ssm_scan_f32_group<96/WARP_SIZE, 96>, launch_params,
src0, src1, src2, src3, src4, src5, src6, dst,
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
- src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
+ src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
} else if (d_state == 128) {
constexpr int threads = 128;
constexpr int num_warps = threads/WARP_SIZE;
@@ -271,7 +275,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
ggml_cuda_kernel_launch(ssm_scan_f32_group<128/WARP_SIZE, 128>, launch_params,
src0, src1, src2, src3, src4, src5, src6, dst,
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
- src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
+ src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
} else if (d_state == 256) { // Falcon-H1
constexpr int threads = 256;
constexpr int num_warps = threads/WARP_SIZE;
@@ -281,7 +285,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
ggml_cuda_kernel_launch(ssm_scan_f32_group<256/WARP_SIZE, 256>, launch_params,
src0, src1, src2, src3, src4, src5, src6, dst,
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
- src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
+ src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
} else {
GGML_ABORT("doesn't support d_state!=(96, 128 or 256).");
}
@@ -570,12 +574,13 @@ __global__ void ssm_ssd_scale_state_kernel(
}
// Copy initial state from src0[ids[s]] into s_cur for each sequence.
+// src0 and s_cur can alias when the state is written straight into the cache.
// Grid: (ceil(d_state * head_dim * n_head / BLOCK), n_seqs)
template <int BLOCK_SIZE>
__global__ void ssm_ssd_init_state_kernel(
- const float * __restrict__ src0, // {d_state, head_dim, n_head, n_rs}
+ const float * src0, // {d_state, head_dim, n_head, n_rs}
const int32_t * __restrict__ ids, // {n_seqs}
- float * __restrict__ s_cur, // {d_state, head_dim, n_head, n_seqs}
+ float * s_cur, // {d_state, head_dim, n_head, n_seqs}
const int state_size, // d_state * head_dim * n_head
const int64_t s0_stride_seq) { // elements between state rows
const int s = blockIdx.y;
@@ -599,7 +604,8 @@ static void ssm_scan_ssd_f32_cuda(
const int A_stride, // A (src3) stride between heads
const int B_stride_tok, const int B_stride_seq, // B (src4) strides
const int C_stride_tok, const int C_stride_seq, // C (src5) strides
- const int64_t s_off, const int64_t d_state, const int64_t head_dim,
+ float * s_cur, // state: dst state tail, or the cache when fused
+ const int64_t d_state, const int64_t head_dim,
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq) {
cudaStream_t stream = ctx.stream();
@@ -625,7 +631,6 @@ static void ssm_scan_ssd_f32_cuda(
matmul_t * X_dt = X_dt_buf.get();
matmul_t * B_weighted = B_w_buf.get();
float * C_scaled = C_s_buf.get();
- float * s_cur = (float *)((char *)dst_d + s_off); // write state directly to dst
// Step 1: softplus(dt) and parallel prefix sum over full sequence
{
@@ -780,7 +785,8 @@ static void ssm_scan_ssd_f32_cuda(
}
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
-void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
+static void ggml_cuda_op_ssm_scan_impl(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
+ const ggml_cuda_ssm_scan_fused_cache * cache) {
const struct ggml_tensor * src0 = dst->src[0]; // s
const struct ggml_tensor * src1 = dst->src[1]; // x
const struct ggml_tensor * src2 = dst->src[2]; // dt
@@ -864,12 +870,21 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
(int)(src3->nb[1] / sizeof(float)),
(int)(src4->nb[2] / sizeof(float)), (int)(src4->nb[3] / sizeof(float)),
(int)(src5->nb[2] / sizeof(float)), (int)(src5->nb[3] / sizeof(float)),
- s_off, nc, nr, nh, ng, n_t, n_s);
+ cache ? cache->data : (float *) ((char *) dst_d + s_off), nc, nr, nh, ng, n_t, n_s);
return;
}
#endif
ssm_scan_f32_cuda(src0_d, src1_d, src2_d, src3_d, src4_d, src5_d, src6_d, dst_d,
src0->nb[2], src0->nb[3], src1->nb[2], src1->nb[3], src2->nb[1], src2->nb[2],
src3->nb[1], src4->nb[2], src4->nb[3], src5->nb[2], src5->nb[3],
- s_off, nc, nr, nh, ng, n_t, n_s, K, stream);
+ s_off, nc, nr, nh, ng, n_t, n_s, K, cache, stream);
+}
+
+void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
+ ggml_cuda_op_ssm_scan_impl(ctx, dst, nullptr);
+}
+
+void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
+ ggml_cuda_ssm_scan_fused_cache cache) {
+ ggml_cuda_op_ssm_scan_impl(ctx, dst, &cache);
}
diff --git a/ggml/src/ggml-cuda/ssm-scan.cuh b/ggml/src/ggml-cuda/ssm-scan.cuh
index ee078f5eb..407a4c712 100644
--- a/ggml/src/ggml-cuda/ssm-scan.cuh
+++ b/ggml/src/ggml-cuda/ssm-scan.cuh
@@ -1,3 +1,13 @@
#include "common.cuh"
+// fused-kernel recurrent-state output; strides in elements (per-seq stride is always the state row size, set in-kernel)
+struct ggml_cuda_ssm_scan_fused_cache {
+ float * data; // rollback slot 0
+ int64_t slot_stride; // between rollback slots
+};
+
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
+
+// same op, but writes the state snapshot(s) into the cache instead of dst (see ggml_cuda_try_ssm_scan_cache_fusion)
+void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
+ ggml_cuda_ssm_scan_fused_cache cache);
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 9a8d46bfd..4e3c6942b 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -4768,6 +4768,104 @@ struct test_ssm_scan_rollback : public test_case {
}
};
+// GGML_OP_SSM_SCAN + GGML_OP_CPY (recurrent cache fusion)
+struct test_ssm_scan_cache_fusion : public test_case {
+ const ggml_type type;
+
+ const int64_t d_state;
+ const int64_t head_dim;
+ const int64_t n_head;
+ const int64_t n_group;
+ const int64_t n_seq_tokens;
+ const int64_t n_seqs;
+ const int64_t K; // snapshot slot count (1 = final state only)
+
+ ggml_tensor * cpy_node = nullptr;
+
+ std::string vars() override {
+ return VARS_TO_STR8(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, K);
+ }
+
+ test_ssm_scan_cache_fusion(ggml_type type = GGML_TYPE_F32,
+ int64_t d_state = 128, int64_t head_dim = 64, int64_t n_head = 16, int64_t n_group = 2,
+ int64_t n_seq_tokens = 4, int64_t n_seqs = 1, int64_t K = 4)
+ : type(type), d_state(d_state), head_dim(head_dim), n_head(n_head), n_group(n_group),
+ n_seq_tokens(n_seq_tokens), n_seqs(n_seqs), K(K) {}
+
+ ggml_tensor * build_graph(ggml_context * ctx) override {
+ const int64_t D = d_state * head_dim * n_head;
+ const int64_t n_written = std::min<int64_t>(n_seq_tokens, K);
+
+ // more cache rows per slot than seqs and a non-zero first row, so a wrong slot stride or offset shows up
+ const int64_t mem_size = n_seqs + 2;
+ const int64_t kv_head = 1;
+
+ ggml_tensor * s = ggml_new_tensor_4d(ctx, type, d_state, head_dim, n_head, n_seqs);
+ ggml_tensor * x = ggml_new_tensor_4d(ctx, type, head_dim, n_head, n_seq_tokens, n_seqs);
+ ggml_tensor * dt = ggml_new_tensor_3d(ctx, type, n_head, n_seq_tokens, n_seqs);
+ ggml_tensor * A = ggml_new_tensor_2d(ctx, type, 1, n_head);
+ ggml_tensor * B = ggml_new_tensor_4d(ctx, type, d_state, n_group, n_seq_tokens, n_seqs);
+ ggml_tensor * C = ggml_new_tensor_4d(ctx, type, d_state, n_group, n_seq_tokens, n_seqs);
+ ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, n_seqs);
+ ggml_set_name(A, "A");
+ ggml_set_name(ids, "ids");
+
+ ggml_tensor * out = ggml_ssm_scan(ctx, s, x, dt, A, B, C, ids, K);
+ ggml_set_name(out, "ssm_out");
+
+ // snapshot tail view [D, n_seqs, n_written]
+ ggml_tensor * src = ggml_view_3d(ctx, out,
+ D, n_seqs, n_written,
+ ggml_row_size(out->type, D),
+ ggml_row_size(out->type, D * n_seqs),
+ ggml_row_size(out->type, ggml_nelements(x)));
+
+ // recurrent cache view [D, n_seqs, n_written]
+ ggml_tensor * cache = ggml_new_tensor_2d(ctx, type, D, mem_size * n_written);
+ ggml_set_name(cache, "cache");
+ ggml_tensor * dst = ggml_view_3d(ctx, cache,
+ D, n_seqs, n_written,
+ cache->nb[1],
+ mem_size * cache->nb[1],
+ kv_head * cache->nb[1]);
+
+ ggml_tensor * cpy = ggml_cpy(ctx, src, dst);
+ ggml_set_name(cpy, "ssm_cache_cpy");
+ cpy_node = cpy;
+
+ // read the cpy output so that neither the scan nor the cpy is the graph output (cont, since the cache view is strided)
+ ggml_tensor * res = ggml_sum(ctx, ggml_cont(ctx, cpy));
+ return res;
+ }
+
+ std::string op_desc(ggml_tensor * t) override {
+ GGML_UNUSED(t);
+ return "SSM_SCAN_CACHE_FUSION";
+ }
+
+ bool run_whole_graph() override { return true; }
+ std::vector<ggml_tensor *> fusion_test_nodes() override { return { cpy_node }; }
+
+ void initialize_tensors(ggml_context * ctx) override {
+ for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) {
+ if (ggml_is_view_op(t->op)) { continue; }
+ if (strcmp(t->name, "ids") == 0) {
+ std::vector<int32_t> data(t->ne[0]);
+ for (int i = 0; i < t->ne[0]; i++) {
+ data[i] = i;
+ }
+ ggml_backend_tensor_set(t, data.data(), 0, t->ne[0] * sizeof(int32_t));
+ } else if (strcmp(t->name, "A") == 0) {
+ init_tensor_uniform(t, -1.0f, -0.5f);
+ } else if (strcmp(t->name, "cache") == 0) {
+ init_tensor_uniform(t, 0.0f, 0.0f);
+ } else {
+ init_tensor_uniform(t);
+ }
+ }
+ }
+};
+
// GGML_OP_RWKV_WKV6
struct test_rwkv_wkv6 : public test_case {
const ggml_type type;
@@ -10445,6 +10543,16 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 128, 2)); // SSD multi-chunk, no tail (exercises the chunk-to-chunk state handoff)
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 128, 2, false, /*K=*/1, /*weak_decay=*/true)); // SSD multi-chunk, carried state not numerically negligible
+ // ssm_scan + cache cpy fusion
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 4, 1, 4));
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 1, 1, 4)); // n_seq_tokens < K
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 8, 1, 3)); // n_seq_tokens > K
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 96, 64, 16, 2, 4, 1, 4));
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 256, 64, 8, 2, 4, 1, 4));
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 1, 1, 1)); // K == 1, final state only
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 4, 1, 1));
+ test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 300, 1, 1)); // K == 1, SSD path over two chunks
+
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 1, 1));
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 32, 1));
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 32, 4));