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));