]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
ggml : recurrent state rollback for ggml_ssm_scan (#26623)
authorlnigam <redacted>
Fri, 14 Aug 2026 14:20:40 +0000 (19:50 +0530)
committerGitHub <redacted>
Fri, 14 Aug 2026 14:20:40 +0000 (17:20 +0300)
* Initial changes for Recurrent state rollback for nemotron for cpu and cuda

* Removing CPU RS rollback. Will enable it in subsequent PRs

* addition of test case

* Removing assert and calling runtime API to check if op is supported

* removing extra API and updating the call sites for K

* replace static cuda detection to runtime fused_op api

* address review comments and fallback when SSM rollback not supprted

* Adding changes for supporting RS-rollback in CPU. Also added test-backend-ops for cpu and cuda

* removing memory manipulation as rs rollback is now supported in CPU

* removing the static probe which is not needed now

* correcting the format

* address review comments

* enabling test for all the backends, unsupported backends will fallback to CPU

* Apply suggestions from code review

Co-authored-by: Georgi Gerganov <redacted>
* choose different graph based on the result of fused_ssm_op is supported or not and also handled memory->n_rs_seq >1 case incase of op is not supported

* Support K > 1 in ssm_scan for all backends

* Fix CI Issues

---------

Co-authored-by: Georgi Gerganov <redacted>
Co-authored-by: Gaurav Garg <redacted>
25 files changed:
ggml/include/ggml.h
ggml/src/ggml-cpu/ggml-cpu.cpp
ggml/src/ggml-cpu/ops.cpp
ggml/src/ggml-cuda/ggml-cuda.cu
ggml/src/ggml-cuda/ssm-scan.cu
ggml/src/ggml-et/et-kernels/src/ssm_scan_f32.c
ggml/src/ggml-et/ggml-et-ops.cpp
ggml/src/ggml-et/ggml-et-ops.h
ggml/src/ggml-metal/ggml-metal-device.m
ggml/src/ggml-metal/ggml-metal-impl.h
ggml/src/ggml-metal/ggml-metal-ops.cpp
ggml/src/ggml-metal/ggml-metal.metal
ggml/src/ggml-sycl/ssm_scan.cpp
ggml/src/ggml-vulkan/ggml-vulkan.cpp
ggml/src/ggml-vulkan/vulkan-shaders/ssm_scan.comp
ggml/src/ggml-webgpu/ggml-webgpu.cpp
ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl
ggml/src/ggml.c
src/llama-arch.cpp
src/llama-context.cpp
src/llama-model-loader.cpp
src/models/mamba-base.cpp
src/models/plamo2.cpp
tests/CMakeLists.txt
tests/test-backend-ops.cpp

index 5cb49d0ee482c84176f9ffcce2dc76ad3601c77e..c2ccd9725396a04177e04f9e224cc4d689f33ae8 100644 (file)
@@ -2459,7 +2459,8 @@ extern "C" {
             struct ggml_tensor  * A,
             struct ggml_tensor  * B,
             struct ggml_tensor  * C,
-            struct ggml_tensor  * ids);
+            struct ggml_tensor  * ids,
+            int64_t               K);
 
     // partition into non-overlapping windows with padding if needed
     // example:
index c0c9aa3cf09c573e308b32d3c90f058c709ce997..8cece71f186f9c99d5f5a110241e6e3b890c1581 100644 (file)
@@ -472,6 +472,8 @@ static bool ggml_backend_cpu_device_supports_op(ggml_backend_dev_t dev, const st
                 src1->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
         case GGML_OP_CONV_2D:
             return ggml_is_contiguous(op->src[0]);
+        case GGML_OP_SSM_SCAN:
+            return ggml_get_op_params_i32(op, 0) == 1 || op->src[3]->ne[0] == 1;
         default:
             return true;
     }
index 25bb7438389dd646a2f202d9273aba1a98610376..001e1ae85a89969fbd02b14c887d6872ee0e61a0 100644 (file)
@@ -9644,11 +9644,13 @@ static void ggml_compute_forward_ssm_scan_f32(
     const int64_t ng = src4->ne[1];
     const int64_t nt = src1->ne[2]; // number of tokens per sequence
     const int64_t ns = src1->ne[3]; // number of sequences in the batch
+    const int64_t K  = ggml_get_op_params_i32(dst, 0);
 
     // can't use ggml_nbytes because src1 is not necessarily contiguous
     const int64_t s_off = ggml_nelements(src1) * ggml_element_size(src1);
 
-    GGML_ASSERT(ggml_nelements(src1) + nc*nr*nh*ns == ggml_nelements(dst));
+    GGML_ASSERT(K >= 1);
+    GGML_ASSERT(ggml_nelements(src1) + K*nc*nr*nh*ns == ggml_nelements(dst));
     GGML_ASSERT(src0->nb[0] == sizeof(float));
     GGML_ASSERT(src1->nb[0] == sizeof(float));
     GGML_ASSERT(src2->nb[0] == sizeof(float));
@@ -9657,6 +9659,7 @@ static void ggml_compute_forward_ssm_scan_f32(
     GGML_ASSERT(src5->nb[0] == sizeof(float));
     GGML_ASSERT(src6->nb[0] == sizeof(int32_t));
     GGML_ASSERT(nh % ng == 0);
+    GGML_ASSERT(src3->ne[0] == 1 || K == 1);
 
     // heads per thread
     const int dh = (nh + nth - 1)/nth;
@@ -9831,6 +9834,13 @@ static void ggml_compute_forward_ssm_scan_f32(
                     }
                 }
             }
+            const int64_t slot = nt - 1 - i2;
+            if (K > 1 && slot > 0 && slot < K) {
+                float * s_snapshot = (float *) ((char *) dst->data + s_off + (slot*ns + i3)*(src0->nb[3]));
+                for (int h = ih0; h < ih1; ++h) {
+                    memcpy((char *) s_snapshot + h*src0->nb[2], (char *) s + h*src0->nb[2], src0->nb[2]);
+                }
+            }
             // use the output as the source when it's not the first token-wise iteration
             s0 = s;
         }
index cb7e9330c8c67ec756a20501516c3bb0725b3f4c..598f3228c401ace65fb8db4772bbeff66f7678e7 100644 (file)
@@ -5189,11 +5189,17 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
                    (op->src[1]->type == GGML_TYPE_F32 || op->src[1]->type == GGML_TYPE_F16) &&
                    (op->type         == GGML_TYPE_F32 || op->type         == GGML_TYPE_F16);
         case GGML_OP_SSM_SCAN: {
+            const int32_t K = ggml_get_op_params_i32(op, 0);
+
             if (op->src[3]->ne[0] == 1) {
                 // Mamba2
                 // (kernel only supports (d_state == 128 || d_state == 256) && d_head % 16 == 0)
                 return (op->src[0]->ne[0] == 128 || op->src[0]->ne[0] == 256) && op->src[0]->ne[1] % 16 == 0;
             } else {
+                if (K > 1) {
+                    return false;
+                }
+
                 // Mamba
                 // (kernel only supports d_state == 16, d_head == 1, n_head % 128 == 0, n_group == 1)
                 return op->src[0]->ne[0] == 16 && op->src[0]->ne[1] == 1 && op->src[0]->ne[2] % 128 == 0 && op->src[4]->ne[1] == 1;
index f3418c2af83d9e07fafff767a90a7b5a90858139..ef342f01f1dd85629c728e2d3ff8da86c80490f3 100644 (file)
@@ -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 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) {
     const float   * GGML_CUDA_RESTRICT src0 = src0_ptr;
     const float   * GGML_CUDA_RESTRICT src1 = src1_ptr;
     const float   * GGML_CUDA_RESTRICT src2 = src2_ptr;
@@ -217,6 +217,16 @@ __global__ void __launch_bounds__(d_state, 1)
         if (lane == 0) {
             y_warp[i * stride_y] = state_sum;
         }
+
+        // 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);
+#pragma unroll
+            for (int j = 0; j < c_factor; j++) {
+                s_snapshot_warp[WARP_SIZE * j + lane] = state[j];
+            }
+        }
     }
 
     // write back the state
@@ -232,7 +242,7 @@ 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,
-                              cudaStream_t stream) {
+                              const int64_t K, cudaStream_t stream) {
     // NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
     if (src3_nb1 == sizeof(float)) {
         // Mamba-2
@@ -245,7 +255,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);
+                    src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, 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;
@@ -255,12 +265,13 @@ 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);
+                    src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
         } else {
             GGML_ABORT("doesn't support d_state!=(128 or 256).");
         }
     } else {
         // Mamba-1
+        GGML_ASSERT(K == 1);
         constexpr int threads = 128;
         GGML_ASSERT(n_head % threads == 0);
         GGML_ASSERT(head_dim == 1);
@@ -769,10 +780,12 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     const int64_t ng  = src4->ne[1];  // n_group
     const int64_t n_t = src1->ne[2];  // number of tokens per sequence
     const int64_t n_s = src1->ne[3];  // number of sequences in the batch
+    const int32_t K_param = ggml_get_op_params_i32(dst, 0);
+    const int64_t K = K_param > 0 ? K_param : 1;
 
     const int64_t s_off = ggml_nelements(src1) * sizeof(float);
 
-    GGML_ASSERT(ggml_nelements(src1) + nc*nr*nh*n_s == ggml_nelements(dst));
+    GGML_ASSERT(ggml_nelements(src1) + K*nc*nr*nh*n_s == ggml_nelements(dst));
     GGML_ASSERT(src0->nb[0] == sizeof(float));
     GGML_ASSERT(src1->nb[0] == sizeof(float));
     GGML_ASSERT(src2->nb[0] == sizeof(float));
@@ -780,6 +793,7 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     GGML_ASSERT(src4->nb[0] == sizeof(float));
     GGML_ASSERT(src5->nb[0] == sizeof(float));
     GGML_ASSERT(src6->nb[0] == sizeof(int32_t));
+    GGML_ASSERT(src3->ne[0] == 1 || K == 1);
 
     const float * src0_d = (const float *) src0->data;
     const float * src1_d = (const float *) src1->data;
@@ -814,6 +828,7 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     const bool is_mamba2 = (src3->nb[1] == sizeof(float));
     const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
     const bool use_ssd = is_mamba2 && n_t > SSM_SSD_MIN_TOKENS
+                      && K == 1
                       && n_t <= SSM_SSD_MAX_TOKENS
                       && GGML_CUDA_CC_IS_NVIDIA(cc)
                       && cc >= GGML_CUDA_CC_TURING
@@ -841,5 +856,5 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     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, stream);
+                      s_off, nc, nr, nh, ng, n_t, n_s, K, stream);
 }
index c114e9981d2a66537f09aab7f99547ef4a627e1c..82ac4309cf11e0702940ac98142ed1d2d0a4500c 100644 (file)
@@ -12,7 +12,8 @@ struct ggml_et_ssm_scan_params {
     struct ggml_tensor src4;  // B:   [d_state, n_group, n_seq_tokens, n_seqs]
     struct ggml_tensor src5;  // C:   [d_state, n_group, n_seq_tokens, n_seqs]
     struct ggml_tensor src6;  // ids: [n_seqs] i32
-    struct ggml_tensor dst;   // packed [y, final_state]
+    struct ggml_tensor dst;   // packed [y, states]
+    int32_t            K;
 };
 
 static inline float softplus_f32(float x) {
@@ -72,6 +73,7 @@ int entry_point(struct ggml_et_ssm_scan_params * params, void * env) {
     const int64_t n_seq_tokens = src1->ne[2];
     const int64_t n_seqs       = src1->ne[3];
     const int64_t y_elems      = src1->ne[0] * src1->ne[1] * src1->ne[2] * src1->ne[3];
+    const int64_t K            = params->K;
 
     if (src0->nb[0] != sizeof(float) || src1->nb[0] != sizeof(float) || src2->nb[0] != sizeof(float) ||
         src3->nb[0] != sizeof(float) || src4->nb[0] != sizeof(float) || src5->nb[0] != sizeof(float) ||
@@ -79,7 +81,7 @@ int entry_point(struct ggml_et_ssm_scan_params * params, void * env) {
         return -1;
     }
 
-    if (n_group <= 0 || n_head % n_group != 0) {
+    if (K < 1 || n_group <= 0 || n_head % n_group != 0) {
         return -1;
     }
 
@@ -260,6 +262,15 @@ int entry_point(struct ggml_et_ssm_scan_params * params, void * env) {
                         sumf += st * C_row[state_idx];
                     }
 
+                    const int64_t slot = n_seq_tokens - 1 - token_idx;
+                    if (slot > 0 && slot < K) {
+                        float * state_snapshot =
+                            (float *) ((char *) state_dst + (size_t) slot * n_seqs * src0->nb[3]);
+                        for (int64_t i = 0; i < d_state; ++i) {
+                            state_snapshot[i] = state_dst[i];
+                        }
+                    }
+
                     dst_data[seq_idx * (n_seq_tokens * n_head * head_dim) + token_idx * (n_head * head_dim) +
                              head_idx * head_dim + dim_idx] = sumf;
                 }
index 6c80fe8acde3a9a4064d6dce7c9c5142dd775921..7871d524081875d4ea08cbc8eab3da4669ea9780 100644 (file)
@@ -2064,6 +2064,7 @@ bool ggml_et_op_ssm_scan(ggml_backend_et_device_context * dev_ctx, const ggml_te
     params.src5 = *node->src[5];
     params.src6 = *node->src[6];
     params.dst  = *node;
+    params.K    = ggml_get_op_params_i32(node, 0);
 
     bool kernel_result = ggml_et_launch_kernel(dev_ctx, "ssm_scan_f32", &params, sizeof(params), 0xFFFFFFFF);
 
index 2c7ca7ece2055887f7cddcfeeac39cdba3b7dba0..032f7a263913f569cfaded052723a014e253fb49 100644 (file)
@@ -218,7 +218,8 @@ struct ggml_et_ssm_scan_params {
     ggml_tensor src4;  // B:   [d_state, n_group, n_seq_tokens, n_seqs]
     ggml_tensor src5;  // C:   [d_state, n_group, n_seq_tokens, n_seqs]
     ggml_tensor src6;  // ids: [n_seqs] i32
-    ggml_tensor dst;   // [y, final_state] packed output from ggml_ssm_scan()
+    ggml_tensor dst;   // [y, states] packed output from ggml_ssm_scan()
+    int32_t     K;
 };
 
 struct ggml_et_rwkv_wkv6_params {
index b70816c32d153a8a91f710e71fab7eeb60c4b4ba..312b00dc48aa3a380214d0fdeb11b6709ac0f154 100644 (file)
@@ -1376,9 +1376,10 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
                 ggml_is_contiguous_rows(op->src[1]) &&
                 ggml_is_contiguous_rows(op->src[2]) &&
                 ggml_is_contiguous_rows(op->src[3]);
-        case GGML_OP_SSM_CONV:
         case GGML_OP_SSM_SCAN:
             return has_simdgroup_reduction;
+        case GGML_OP_SSM_CONV:
+            return has_simdgroup_reduction;
         case GGML_OP_RWKV_WKV6:
         case GGML_OP_RWKV_WKV7:
             return true;
index cf32c5c5b246969570d3be605e16acbc74319859..1f6e8c48bcb8a2fc89f9483bbc07f209b27a7178 100644 (file)
@@ -880,6 +880,7 @@ typedef struct {
     int64_t  n_group;
     int64_t  n_seq_tokens;
     int64_t  n_seqs;
+    int64_t  K;
     uint64_t s_off;
     uint64_t nb00;
     uint64_t nb01;
index 6d324056dd2c5944bfd36b7737b9701612506cff..b7f9b2d0d9c2702457a74a41a60106b8ebe55a8a 100644 (file)
@@ -1710,6 +1710,10 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) {
     const int64_t n_group      = ne41;
     const int64_t n_seq_tokens = ne12;
     const int64_t n_seqs       = ne13;
+    const int64_t K            = ggml_get_op_params_i32(op, 0);
+
+    GGML_ASSERT(K >= 1);
+    GGML_ASSERT(ggml_nelements(op->src[1]) + K*d_state*d_inner*n_head*n_seqs == ggml_nelements(op));
 
     ggml_metal_kargs_ssm_scan args = {
         /*.d_state      =*/ d_state,
@@ -1718,6 +1722,7 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) {
         /*.n_group      =*/ n_group,
         /*.n_seq_tokens =*/ n_seq_tokens,
         /*.n_seqs       =*/ n_seqs,
+        /*.K            =*/ K,
         /*.s_off        =*/ ggml_nelements(op->src[1]) * sizeof(float),
         /*.nb00         =*/ nb00,
         /*.nb01         =*/ nb01,
index b38b23edc95dbf1d89ef2aa029de215e99da4e78..243c997fc44a9d2a9820241d795390873bab187f 100644 (file)
@@ -2429,6 +2429,8 @@ kernel void kernel_ssm_scan_f32(
     const int32_t nh  = args.n_head;
     const int32_t ng  = args.n_group;
     const int32_t n_t = args.n_seq_tokens;
+    const int32_t n_s = args.n_seqs;
+    const int32_t K   = args.K;
 
     const int32_t s_off = args.s_off;
 
@@ -2487,6 +2489,12 @@ kernel void kernel_ssm_scan_f32(
             // recurse
             s0 = s;
 
+            const int32_t slot = n_t - 1 - (i2 + t);
+            if (slot > 0 && slot < K) {
+                device float * s_snapshot = (device float *) ((device char *) s_buff + (int64_t) slot*n_s*args.nb03);
+                s_snapshot[i] = s;
+            }
+
             B  += args.ns42;
             C  += args.ns52;
         }
index ae65298138409f50a995d2e74c7cabdcce00bbdf..7fceb85d254151c42a31ec98e5f714324e3b4c62 100644 (file)
@@ -10,6 +10,7 @@ static void ssm_scan_f32_group(
         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,
         const sycl::nd_item<2> & item) {
 
     const int lane     = item.get_local_id(1) % WARP_SIZE;
@@ -64,6 +65,15 @@ static void ssm_scan_f32_group(
         if (lane == 0) {
             y_warp[i * stride_y] = state_sum;
         }
+
+        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 * item.get_group_range(0) + seq_idx) * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
+#pragma unroll
+            for (int j = 0; j < c_factor; j++) {
+                s_snapshot_warp[WARP_SIZE * j + lane] = state[j];
+            }
+        }
     }
 
 #pragma unroll
@@ -79,6 +89,7 @@ static void ssm_scan_f32_sycl(
         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,
         dpct::queue_ptr stream) {
 
     // NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
@@ -94,7 +105,7 @@ static void ssm_scan_f32_sycl(
                 ssm_scan_f32_group<128 / WARP_SIZE, 128>(
                     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, item);
+                    src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K, item);
             });
     } else if (d_state == 256) {
         constexpr int threads   = 256;
@@ -107,7 +118,7 @@ static void ssm_scan_f32_sycl(
                 ssm_scan_f32_group<256 / WARP_SIZE, 256>(
                     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, item);
+                    src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K, item);
             });
     } else {
         GGML_ABORT("ssm_scan: unsupported d_state (must be 128 or 256)");
@@ -133,9 +144,12 @@ inline void ggml_sycl_op_ssm_scan(ggml_backend_sycl_context & ctx, ggml_tensor *
     const int64_t ng  = src4->ne[1];
     const int64_t n_t = src1->ne[2];
     const int64_t n_s = src1->ne[3];
+    const int64_t K   = ggml_get_op_params_i32(dst, 0);
     const int64_t s_off = ggml_nelements(src1) * sizeof(float);
 
-    GGML_ASSERT(ggml_nelements(src1) + nc * nr * nh * n_s == ggml_nelements(dst));
+    GGML_ASSERT(K >= 1);
+    GGML_ASSERT(ggml_nelements(src1) + K * nc * nr * nh * n_s == ggml_nelements(dst));
+    GGML_ASSERT(src3->ne[0] == 1 || K == 1);
 
     dpct::queue_ptr stream = ctx.stream();
     SYCL_CHECK(ggml_sycl_set_device(ctx.device));
@@ -147,7 +161,7 @@ inline void ggml_sycl_op_ssm_scan(ggml_backend_sycl_context & ctx, ggml_tensor *
         static_cast<const int32_t *>(src6->data), static_cast<float *>(dst->data),
         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, stream);
+        s_off, nc, nr, nh, ng, n_t, n_s, K, stream);
 }
 
 void ggml_sycl_ssm_scan(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
index c815d4ff99b55af54ab3684a821866d719ab1df9..ff4a3390448b9716c016180cf5cb04658513b045 100644 (file)
@@ -1861,6 +1861,7 @@ struct vk_op_ssm_scan_push_constants {
     uint32_t nb42, nb43, nb52, nb53;
     uint32_t s_off;
     uint32_t n_head, d_head, n_group, n_tok;
+    uint32_t n_seq, K;
 };
 struct vk_op_ssm_conv_push_constants {
     uint32_t nb01, nb02;
@@ -12731,7 +12732,8 @@ static void ggml_vk_ssm_scan(ggml_backend_vk_context * ctx, vk_context& subctx,
         (uint32_t)src4->nb[2], (uint32_t)src4->nb[3],
         (uint32_t)src5->nb[2], (uint32_t)src5->nb[3],
         (uint32_t)s_off,
-        n_head, head_dim, n_group, n_tok
+        n_head, head_dim, n_group, n_tok,
+        n_seq, (uint32_t) ggml_get_op_params_i32(dst, 0)
     };
 
     vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst);
@@ -19417,8 +19419,9 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph *
         } else if (tensor->op == GGML_OP_ADD_ID) {
             tensor_clone = ggml_add_id(ggml_ctx, src_clone[0], src_clone[1], src_clone[2]);
         } else if (tensor->op == GGML_OP_SSM_SCAN) {
+            const int32_t K = ggml_get_op_params_i32(tensor, 0);
             tensor_clone = ggml_ssm_scan(ggml_ctx, src_clone[0], src_clone[1], src_clone[2],
-                                         src_clone[3], src_clone[4], src_clone[5], src_clone[6]);
+                                         src_clone[3], src_clone[4], src_clone[5], src_clone[6], K);
         } else if (tensor->op == GGML_OP_SSM_CONV) {
             tensor_clone = ggml_ssm_conv(ggml_ctx, src_clone[0], src_clone[1]);
         } else if (tensor->op == GGML_OP_ROLL) {
index c7416206dbdaed3ea3bb185d02d9a398a8a3ce70..4fecb3aa5ace14dd64eaa9644d5d2abfd31ce944 100644 (file)
@@ -33,6 +33,8 @@ layout(push_constant) uniform PushConstants {
     uint d_head;
     uint n_group;
     uint n_tok;
+    uint n_seq;
+    uint K;
 };
 
 float softplus(float x) {
@@ -114,6 +116,14 @@ void main() {
         if (lane == 0) {
             d[y_base_idx + i * stride_y] = state_sum;
         }
+
+        const uint slot = n_tok - 1u - i;
+        if (slot > 0u && slot < K) {
+            const uint snapshot_base_idx = s_base_idx + slot * n_seq * (nb03 / 4u);
+            [[unroll]] for (uint j = 0; j < c_factor; j++) {
+                d[snapshot_base_idx + SUBGROUP_SIZE * j + lane] = state[j];
+            }
+        }
     }
 
     // write back the state
index 6741752b36112e0d5bc24fbde9c23ae708cccea1..394aeeda2744bb40b80ccfd3c70b00c6f155cef9 100644 (file)
@@ -1327,6 +1327,7 @@ static webgpu_encoded_op ggml_webgpu_ssm_scan(webgpu_context & ctx,
         (uint32_t) src4->ne[1],
         (uint32_t) src1->ne[2],
         (uint32_t) ggml_nelements(src1),
+        (uint32_t) ggml_get_op_params_i32(dst, 0),
     };
 
     std::vector<wgpu::BindGroupEntry> entries = {
index 2d4c4e5a0b9186dfc8e0a06b46a19637efabb6b7..57f012ad0f87c317c9d8a053d4e04330301351a4 100644 (file)
@@ -41,6 +41,7 @@ struct Params {
     n_seq_tokens: u32,
 
     y_elems: u32,
+    K: u32,
 };
 
 @group(0) @binding(0) var<storage, read_write> s_in: array<f32>;
@@ -123,6 +124,7 @@ fn main(
     let head_seq = wg_linear / params.d_inner;
     let ir = head_seq % params.n_head;
     let i3 = head_seq / params.n_head;
+    let n_seqs = params.y_elems / (params.n_seq_tokens * params.n_head * params.d_inner);
 
     let state_slot = read_state_slot(i3);
     let g = ir / (params.n_head / params.n_group);
@@ -179,6 +181,15 @@ fn main(
 #endif
             s_prev = s;
 
+            let slot = params.n_seq_tokens - 1u - token;
+            if (slot > 0u && slot < params.K) {
+                let snapshot_idx =
+                    params.offset_dst + params.y_elems + tid + i1 * params.d_state +
+                    ir * (params.d_state * params.d_inner) +
+                    (slot * n_seqs + i3) * (params.d_state * params.d_inner * params.n_head);
+                dst[snapshot_idx] = s;
+            }
+
 #ifdef USE_SUBGROUP_REDUCTION
 #ifdef XBC_OVERLAP
             let subgroup_partial = subgroupAdd(s * read_merged_f32(c_idx));
index da7f3a5f2e30540db794cf87464dcbcd27f700bb..d0d369c41737ff94e63325e0b8050d1d05b91546 100644 (file)
@@ -5588,7 +5588,10 @@ struct ggml_tensor * ggml_ssm_scan(
         struct ggml_tensor  * A,
         struct ggml_tensor  * B,
         struct ggml_tensor  * C,
-        struct ggml_tensor  * ids) {
+        struct ggml_tensor  * ids,
+        int64_t               K) {
+    GGML_ASSERT(K >= 1);
+    GGML_ASSERT(K <= INT32_MAX);
     GGML_ASSERT(ggml_is_contiguous(s));
     GGML_ASSERT(ggml_is_contiguous(dt));
     GGML_ASSERT(ggml_is_contiguous(A));
@@ -5625,11 +5628,12 @@ struct ggml_tensor * ggml_ssm_scan(
         if (A->ne[0] != 1) {
             // Mamba-1 has more granular decay factors
             GGML_ASSERT(A->ne[0] == d_state);
+            GGML_ASSERT(K == 1);
         }
     }
 
     // concatenated y + ssm_states
-    struct ggml_tensor * result = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, ggml_nelements(x) + s->ne[0]*s->ne[1]*s->ne[2]*ids->ne[0]);
+    struct ggml_tensor * result = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, ggml_nelements(x) + K*s->ne[0]*s->ne[1]*s->ne[2]*ids->ne[0]);
 
     result->op   = GGML_OP_SSM_SCAN;
     result->src[0] = s;
@@ -5640,6 +5644,8 @@ struct ggml_tensor * ggml_ssm_scan(
     result->src[5] = C;
     result->src[6] = ids;
 
+    ggml_set_op_params_i32(result, 0, (int32_t) K);
+
     return result;
 }
 
index 8ed9391d7c74b4342737ec1e7012832da0952bda..292ab26103d862bd167f5ea20be16a54473040e0 100644 (file)
@@ -1001,6 +1001,8 @@ bool llm_arch_supports_rs_rollback(const llm_arch & arch) {
         case LLM_ARCH_QWEN35:
         case LLM_ARCH_QWEN35MOE:
         case LLM_ARCH_DEEPSEEK4:
+        case LLM_ARCH_NEMOTRON_H:
+        case LLM_ARCH_NEMOTRON_H_MOE:
             return true;
         default:
             return false;
index aa9fb2c3b481431faed9505007f5f883600028d4..cd013cdb14476b0195b1c4fd265617d3638ebf00 100644 (file)
@@ -103,7 +103,7 @@ llama_context::llama_context(
 
     cparams.n_rs_seq = params.n_rs_seq;
     if (cparams.n_rs_seq > 0 && !llm_arch_supports_rs_rollback(model.arch)) {
-        LLAMA_LOG_DEBUG("%s: n_rs_seq=%u requested but model arch does not support recurrent partial rollback; clamping to 0\n",
+        LLAMA_LOG_DEBUG("%s: n_rs_seq=%u requested but model does not support recurrent partial rollback; clamping to 0\n",
                         __func__, cparams.n_rs_seq);
         cparams.n_rs_seq = 0;
     }
index 51ba0543968212fb91da5c3bf5b6207f9b58f6a5..5c5e97fbc9eb92dc8d44095450a8076017880eb2 100644 (file)
@@ -1002,7 +1002,7 @@ static bool weight_buft_supported(const llama_hparams & hparams, ggml_tensor * w
                 ggml_tensor * B   = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, d_state, n_group, n_seq_tokens, n_seqs);
                 ggml_tensor * C   = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, d_state, n_group, n_seq_tokens, n_seqs);
                 ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, n_seqs);
-                op_tensor = ggml_ssm_scan(ctx, s, x, dt, w, B, C, ids);
+                op_tensor = ggml_ssm_scan(ctx, s, x, dt, w, B, C, ids, /*K=*/1);
             } break;
         case GGML_OP_RWKV_WKV6:
             {
index fd3fe3f032309db44f5d2fb84de02dc4b28e2a1d..1f994ae0af34b17ddd929fb7fdcf3d9ce5506339 100644 (file)
@@ -2,6 +2,8 @@
 
 #include "llama-memory-recurrent.h"
 
+#include <algorithm>
+
 llm_build_mamba_base::llm_build_mamba_base(const llm_graph_params & params) : llm_graph_context(params) {}
 
 ggml_tensor * llm_build_mamba_base::build_mamba_layer(llm_graph_input_rs * inp,
@@ -118,7 +120,7 @@ ggml_tensor * llm_build_mamba_base::build_mamba_layer(llm_graph_input_rs * inp,
             // Custom operator to optimize the parallel associative scan
             // as described in the Annex D of the Mamba paper.
             // => {d_inner, n_seq_tokens, n_seqs} and {d_state, d_inner, n_seqs}
-            return ggml_ssm_scan(ctx, ssm, x, dt, A, B, C, ids);
+            return ggml_ssm_scan(ctx, ssm, x, dt, A, B, C, ids, /*K=*/1);
         };
 
         ggml_tensor * y_ssm = build_rs(inp, ssm_states_all, hparams.n_embd_s(), ubatch.n_seqs, get_ssm_rows);
@@ -153,7 +155,8 @@ ggml_tensor * llm_build_mamba_base::build_mamba2_layer(llm_graph_input_rs * inp,
                                                           int                  il) const {
     const auto * mctx_cur = inp->mctx;
 
-    const auto kv_head = mctx_cur->get_head();
+    const auto kv_head  = mctx_cur->get_head();
+    const auto mem_size = mctx_cur->get_size();
 
     const int64_t d_conv   = hparams.ssm_d_conv;
     const int64_t d_inner  = hparams.ssm_d_inner;
@@ -164,6 +167,7 @@ ggml_tensor * llm_build_mamba_base::build_mamba2_layer(llm_graph_input_rs * inp,
     const int64_t n_seqs   = ubatch.n_seqs;
 
     const int64_t n_seq_tokens = ubatch.n_seq_tokens;
+    const int64_t K            = cparams.n_rs_seq > 0 ? (int64_t) cparams.n_rs_seq + 1 : 1;
 
     GGML_ASSERT(n_seqs != 0);
     GGML_ASSERT(ubatch.equal_seqs());
@@ -173,6 +177,7 @@ ggml_tensor * llm_build_mamba_base::build_mamba2_layer(llm_graph_input_rs * inp,
 
     ggml_tensor * conv_states_all = mctx_cur->get_r_l(il);
     ggml_tensor * ssm_states_all  = mctx_cur->get_s_l(il);
+    const int64_t state_slots     = ssm_states_all->ne[1];
 
     ggml_tensor * conv = build_rs(inp, conv_states_all, hparams.n_embd_r(), n_seqs);
     conv               = ggml_reshape_3d(ctx0, conv, d_conv - 1, d_inner + 2 * n_group * d_state, n_seqs);
@@ -198,15 +203,19 @@ ggml_tensor * llm_build_mamba_base::build_mamba2_layer(llm_graph_input_rs * inp,
         // => {d_conv - 1 + n_seq_tokens, d_inner + 2*n_group*d_state, n_seqs}
         ggml_tensor * conv_x = ggml_concat(ctx0, conv, ggml_transpose(ctx0, xBC), 0);
 
-        // copy last (d_conv - 1) columns back into the state cache
-        ggml_tensor * last_conv = ggml_view_3d(ctx0, conv_x, d_conv - 1, d_inner + 2 * n_group * d_state, n_seqs,
-                                               conv_x->nb[1], conv_x->nb[2], n_seq_tokens * (conv_x->nb[0]));
+        const int64_t row_count = (d_conv - 1) * (d_inner + 2 * n_group * d_state);
+        const size_t  row_size  = ggml_row_size(conv_states_all->type, row_count);
+        const int64_t n_written = std::min<int64_t>(n_seq_tokens, K);
 
-        ggml_build_forward_expand(gf, ggml_cpy(ctx0, last_conv,
-                                               ggml_view_1d(ctx0, conv_states_all,
-                                                            (d_conv - 1) * (d_inner + 2 * n_group * d_state) * (n_seqs),
-                                                            kv_head * (d_conv - 1) * (d_inner + 2 * n_group * d_state) *
-                                                                ggml_element_size(conv_states_all))));
+        for (int64_t slot = 0; slot < n_written; ++slot) {
+            ggml_tensor * last_conv = ggml_view_3d(ctx0, conv_x, d_conv - 1, d_inner + 2 * n_group * d_state, n_seqs,
+                                                   conv_x->nb[1], conv_x->nb[2], (n_seq_tokens - slot) * conv_x->nb[0]);
+
+            ggml_build_forward_expand(gf, ggml_cpy(ctx0, last_conv,
+                                                   ggml_view_2d(ctx0, conv_states_all, row_count, n_seqs,
+                                                                conv_states_all->nb[1],
+                                                                ((size_t) slot * mem_size + kv_head) * row_size)));
+        }
 
         // 1D convolution
         // The equivalent is to make a self-overlapping view of conv_x
@@ -244,20 +253,27 @@ ggml_tensor * llm_build_mamba_base::build_mamba2_layer(llm_graph_input_rs * inp,
         // (this is necessary in order to properly use the states before they are overwritten,
         //  while avoiding to make unnecessary copies of the states)
         auto get_ssm_rows = [&](ggml_context * ctx, ggml_tensor * states, ggml_tensor * ids) {
-            ggml_tensor * ssm = ggml_reshape_4d(ctx, states, d_state, head_dim, n_head, mctx_cur->get_size());
+            ggml_tensor * ssm = ggml_reshape_4d(ctx, states, d_state, head_dim, n_head, state_slots);
 
             // TODO: use semistructured matrices to implement state-space duality
             // => {d_inner, n_seq_tokens, n_seqs} and {d_state, d_inner, n_seqs}
-            return ggml_ssm_scan(ctx, ssm, x, dt, A, B, C, ids);
+            // K > 1 asks the backend to return rollback snapshots in addition to the final state.
+            return ggml_ssm_scan(ctx, ssm, x, dt, A, B, C, ids, K);
         };
 
         ggml_tensor * y_ssm = build_rs(inp, ssm_states_all, hparams.n_embd_s(), ubatch.n_seqs, get_ssm_rows);
+        const int64_t D            = d_state * d_inner;
+        const int64_t n_written    = std::min<int64_t>(n_seq_tokens, K);
+        const size_t  row_size     = ggml_row_size(ssm_states_all->type, D);
+        const size_t  y_row_size   = ggml_row_size(y_ssm->type, D);
+        const size_t  state_offset = ggml_nelements(x) * ggml_element_size(x);
 
-        // store last states
         ggml_build_forward_expand(
-            gf, ggml_cpy(ctx0, ggml_view_1d(ctx0, y_ssm, d_state * d_inner * n_seqs, ggml_nelements(x) * x->nb[0]),
-                         ggml_view_1d(ctx0, ssm_states_all, d_state * d_inner * n_seqs,
-                                      kv_head * d_state * d_inner * ggml_element_size(ssm_states_all))));
+            gf, ggml_cpy(ctx0,
+                         ggml_view_3d(ctx0, y_ssm, D, n_seqs, n_written,
+                                      y_row_size, y_row_size * n_seqs, state_offset),
+                         ggml_view_3d(ctx0, ssm_states_all, D, n_seqs, n_written,
+                                      ssm_states_all->nb[1], (size_t) mem_size * row_size, kv_head * row_size)));
 
         ggml_tensor * y = ggml_view_4d(ctx0, y_ssm, head_dim, n_head, n_seq_tokens, n_seqs, x->nb[1], n_head * x->nb[1],
                                        n_seq_tokens * n_head * x->nb[1], 0);
index 0b81513c3687bad54040e5de785b0e5e57fac837..d946b3cff6da23f4d48c99b2fd1fd0bdf77803cb 100644 (file)
@@ -382,7 +382,7 @@ ggml_tensor * llama_model_plamo2::graph::build_plamo2_mamba_layer(llm_graph_inpu
             // Custom operator to optimize the parallel associative scan
             // as described in the Annex D of the Mamba paper.
             // => {d_inner, n_seq_tokens, n_seqs} and {d_state, d_inner, n_seqs}
-            return ggml_ssm_scan(ctx, ssm, x, dt, A, B, C, ids);
+            return ggml_ssm_scan(ctx, ssm, x, dt, A, B, C, ids, /*K=*/1);
         };
 
         ggml_tensor * y_ssm = build_rs(inp, ssm_states_all, hparams.n_embd_s(), ubatch.n_seqs, get_ssm_rows);
index 419e1eba4c2cdb465d20453004eeeca5af28037f..08c6f5a479630b05f38c0cd30b5c596868ceddca 100644 (file)
@@ -217,6 +217,16 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
     set_tests_properties(test-recurrent-state-rollback PROPERTIES
         FIXTURES_REQUIRED generate-models
     )
+
+    llama_test(
+        test-recurrent-state-rollback
+        NAME test-recurrent-state-rollback-nemotron-h
+        LABEL main
+        ARGS -m "${MODEL_DIR}/nemotron_h-dense.gguf"
+    )
+    set_tests_properties(test-recurrent-state-rollback-nemotron-h PROPERTIES
+        FIXTURES_REQUIRED generate-models
+    )
 endif()
 
 llama_build_and_test(test-chat-peg-parser.cpp peg-parser/simple-tokenize.cpp)
index 08c29eec63f10c00ca25d365e8668ca89fbd3a46..3349a64b174da332b66f7634a2bdeef329214268 100644 (file)
@@ -4111,9 +4111,10 @@ struct test_ssm_scan : public test_case {
     const int64_t n_seq_tokens;
     const int64_t n_seqs;
     const bool    xbc_overlap;
+    const int64_t K;
 
     std::string vars() override {
-        return VARS_TO_STR8(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, xbc_overlap);
+        return VARS_TO_STR9(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, xbc_overlap, K);
     }
 
     test_ssm_scan(ggml_type type = GGML_TYPE_F32,
@@ -4123,8 +4124,9 @@ struct test_ssm_scan : public test_case {
             int64_t n_group = 1,
             int64_t n_seq_tokens = 32,
             int64_t n_seqs = 32,
-            bool xbc_overlap = false)
-        : 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), xbc_overlap(xbc_overlap) {}
+            bool xbc_overlap = false,
+            int64_t K = 1)
+        : 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), xbc_overlap(xbc_overlap), K(K) {}
 
     double max_nmse_err() override {
         // SSD path (head_dim > 1) uses FP16 intermediates (M matrix, X_dt); Mamba-1 is pure FP32.
@@ -4153,7 +4155,7 @@ struct test_ssm_scan : public test_case {
             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_tensor * out = ggml_ssm_scan(ctx, s, x, dt, A, B, C, ids);
+        ggml_tensor * out = ggml_ssm_scan(ctx, s, x, dt, A, B, C, ids, K);
         return out;
     }
 
@@ -4185,6 +4187,114 @@ struct test_ssm_scan : public test_case {
     }
 };
 
+struct test_ssm_scan_rollback : 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;
+
+    std::string vars() override {
+        return VARS_TO_STR8(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, K);
+    }
+
+    std::string op_desc(ggml_tensor * t) override {
+        GGML_UNUSED(t);
+        return "SSM_SCAN_ROLLBACK";
+    }
+
+    bool run_whole_graph() override {
+        return true;
+    }
+
+    double max_err() override {
+        return 1e-6;
+    }
+
+    double err(const float * a, const float * b, size_t n) override {
+        double result = 0.0;
+        for (size_t i = 0; i < n; ++i) {
+            result = std::max(result, (double) fabsf(a[i]));
+            result = std::max(result, (double) fabsf(b[i]));
+        }
+        return result;
+    }
+
+    test_ssm_scan_rollback(ggml_type type = GGML_TYPE_F32,
+            int64_t d_state = 32,
+            int64_t head_dim = 64,
+            int64_t n_head  = 16,
+            int64_t n_group = 2,
+            int64_t n_seq_tokens = 8,
+            int64_t n_seqs = 2,
+            int64_t K = 3)
+        : 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 {
+        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_tensor * full = ggml_ssm_scan(ctx, s, x, dt, A, B, C, ids, K);
+
+        const int64_t y_elems     = head_dim * n_head * n_seq_tokens * n_seqs;
+        const int64_t state_elems = d_state  * head_dim * n_head      * n_seqs;
+
+        ggml_tensor * out = nullptr;
+        for (int64_t slot = 0; slot < K; ++slot) {
+            const int64_t prefix_tokens = n_seq_tokens - slot;
+
+            ggml_tensor * x_prefix  = ggml_cont(ctx, ggml_view_4d(ctx, x,  head_dim, n_head,  prefix_tokens, n_seqs, x->nb[1],  x->nb[2],  x->nb[3],  0));
+            ggml_tensor * dt_prefix = ggml_cont(ctx, ggml_view_3d(ctx, dt, n_head,   prefix_tokens, n_seqs, dt->nb[1], dt->nb[2], 0));
+            ggml_tensor * B_prefix  = ggml_cont(ctx, ggml_view_4d(ctx, B,  d_state,  n_group, prefix_tokens, n_seqs, B->nb[1],  B->nb[2],  B->nb[3],  0));
+            ggml_tensor * C_prefix  = ggml_cont(ctx, ggml_view_4d(ctx, C,  d_state,  n_group, prefix_tokens, n_seqs, C->nb[1],  C->nb[2],  C->nb[3],  0));
+
+            ggml_tensor * prefix = ggml_ssm_scan(ctx, s, x_prefix, dt_prefix, A, B_prefix, C_prefix, ids, /*K=*/1);
+
+            ggml_tensor * full_state   = ggml_view_1d(ctx, full,   state_elems, (y_elems + slot*state_elems)*ggml_element_size(full));
+            ggml_tensor * prefix_state = ggml_view_1d(ctx, prefix, state_elems, (head_dim*n_head*prefix_tokens*n_seqs)*ggml_element_size(prefix));
+            ggml_tensor * diff         = ggml_sum(ctx, ggml_sqr(ctx, ggml_sub(ctx, full_state, prefix_state)));
+
+            out = out == nullptr ? diff : ggml_add(ctx, out, diff);
+        }
+
+        return out;
+    }
+
+    void initialize_tensors(ggml_context * ctx) override {
+        std::random_device rd;
+        std::default_random_engine rng(rd());
+        for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) {
+            if (t->type == GGML_TYPE_I32) {
+                if (ggml_is_view_op(t->op)) { continue; }
+                for (int64_t r = 0; r < ggml_nrows(t); r++) {
+                    std::vector<int32_t> data(t->ne[0]);
+                    for (int i = 0; i < t->ne[0]; i++) {
+                        data[i] = i;
+                    }
+                    std::shuffle(data.begin(), data.end(), rng);
+                    ggml_backend_tensor_set(t, data.data(), r * t->nb[1], t->ne[0] * sizeof(int32_t));
+                }
+            } else if (ggml_is_view_op(t->op)) {
+                continue;
+            } else if (t->ne[1] == n_head && t->ne[2] == 1) {
+                init_tensor_uniform(t, -1.0f, -0.5f);
+            } else {
+                init_tensor_uniform(t);
+            }
+        }
+    }
+};
+
 // GGML_OP_RWKV_WKV6
 struct test_rwkv_wkv6 : public test_case {
     const ggml_type type;
@@ -8952,6 +9062,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 80, 128, 1, 256, 1)); // Nemotron-9B SSD path
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 80, 128, 1, 512, 1)); // Nemotron-9B SSD multi-chunk (2 aligned chunks)
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 80, 8, 300, 2)); // Mamba-2 SSD multi-chunk (partial 2nd chunk, 2 seqs)
+    test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 4, 2, false, /*K=*/4)); // Mamba-2 rollback snapshots
+    test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 8, 2, false, /*K=*/3)); // Mamba-2 rollback overflow
+    test_cases.emplace_back(new test_ssm_scan_rollback(GGML_TYPE_F32, 128, 64, 16, 2, 8, 2, /*K=*/3)); // rollback snapshots match prefix states
 
     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));