From 1692f9e50bb20fd96b963af38a282daf78feea64 Mon Sep 17 00:00:00 2001 From: lnigam Date: Fri, 14 Aug 2026 19:50:40 +0530 Subject: [PATCH] ggml : recurrent state rollback for ggml_ssm_scan (#26623) * 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 * 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 Co-authored-by: Gaurav Garg --- ggml/include/ggml.h | 3 +- ggml/src/ggml-cpu/ggml-cpu.cpp | 2 + ggml/src/ggml-cpu/ops.cpp | 12 +- ggml/src/ggml-cuda/ggml-cuda.cu | 6 + ggml/src/ggml-cuda/ssm-scan.cu | 27 +++- .../src/ggml-et/et-kernels/src/ssm_scan_f32.c | 15 ++- ggml/src/ggml-et/ggml-et-ops.cpp | 1 + ggml/src/ggml-et/ggml-et-ops.h | 3 +- ggml/src/ggml-metal/ggml-metal-device.m | 3 +- ggml/src/ggml-metal/ggml-metal-impl.h | 1 + ggml/src/ggml-metal/ggml-metal-ops.cpp | 5 + ggml/src/ggml-metal/ggml-metal.metal | 8 ++ ggml/src/ggml-sycl/ssm_scan.cpp | 22 +++- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 7 +- .../ggml-vulkan/vulkan-shaders/ssm_scan.comp | 10 ++ ggml/src/ggml-webgpu/ggml-webgpu.cpp | 1 + .../ggml-webgpu/wgsl-shaders/ssm_scan.wgsl | 11 ++ ggml/src/ggml.c | 10 +- src/llama-arch.cpp | 2 + src/llama-context.cpp | 2 +- src/llama-model-loader.cpp | 2 +- src/models/mamba-base.cpp | 48 ++++--- src/models/plamo2.cpp | 2 +- tests/CMakeLists.txt | 10 ++ tests/test-backend-ops.cpp | 121 +++++++++++++++++- 25 files changed, 291 insertions(+), 43 deletions(-) diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 5cb49d0ee..c2ccd9725 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -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: diff --git a/ggml/src/ggml-cpu/ggml-cpu.cpp b/ggml/src/ggml-cpu/ggml-cpu.cpp index c0c9aa3cf..8cece71f1 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.cpp +++ b/ggml/src/ggml-cpu/ggml-cpu.cpp @@ -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; } diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 25bb74383..001e1ae85 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -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; } diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index cb7e9330c..598f3228c 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -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; diff --git a/ggml/src/ggml-cuda/ssm-scan.cu b/ggml/src/ggml-cuda/ssm-scan.cu index f3418c2af..ef342f01f 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 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); } diff --git a/ggml/src/ggml-et/et-kernels/src/ssm_scan_f32.c b/ggml/src/ggml-et/et-kernels/src/ssm_scan_f32.c index c114e9981..82ac4309c 100644 --- a/ggml/src/ggml-et/et-kernels/src/ssm_scan_f32.c +++ b/ggml/src/ggml-et/et-kernels/src/ssm_scan_f32.c @@ -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; } diff --git a/ggml/src/ggml-et/ggml-et-ops.cpp b/ggml/src/ggml-et/ggml-et-ops.cpp index 6c80fe8ac..7871d5240 100644 --- a/ggml/src/ggml-et/ggml-et-ops.cpp +++ b/ggml/src/ggml-et/ggml-et-ops.cpp @@ -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", ¶ms, sizeof(params), 0xFFFFFFFF); diff --git a/ggml/src/ggml-et/ggml-et-ops.h b/ggml/src/ggml-et/ggml-et-ops.h index 2c7ca7ece..032f7a263 100644 --- a/ggml/src/ggml-et/ggml-et-ops.h +++ b/ggml/src/ggml-et/ggml-et-ops.h @@ -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 { diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index b70816c32..312b00dc4 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -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; diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index cf32c5c5b..1f6e8c48b 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -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; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 6d324056d..b7f9b2d0d 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -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, diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index b38b23edc..243c997fc 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -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; } diff --git a/ggml/src/ggml-sycl/ssm_scan.cpp b/ggml/src/ggml-sycl/ssm_scan.cpp index ae6529813..7fceb85d2 100644 --- a/ggml/src/ggml-sycl/ssm_scan.cpp +++ b/ggml/src/ggml-sycl/ssm_scan.cpp @@ -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(src6->data), static_cast(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) { diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index c815d4ff9..ff4a33904 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/ssm_scan.comp b/ggml/src/ggml-vulkan/vulkan-shaders/ssm_scan.comp index c7416206d..4fecb3aa5 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/ssm_scan.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/ssm_scan.comp @@ -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 diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 6741752b3..394aeeda2 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -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 entries = { diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl index 2d4c4e5a0..57f012ad0 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl @@ -41,6 +41,7 @@ struct Params { n_seq_tokens: u32, y_elems: u32, + K: u32, }; @group(0) @binding(0) var s_in: array; @@ -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)); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index da7f3a5f2..d0d369c41 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -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; } diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 8ed9391d7..292ab2610 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -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; diff --git a/src/llama-context.cpp b/src/llama-context.cpp index aa9fb2c3b..cd013cdb1 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -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; } diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index 51ba05439..5c5e97fbc 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -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: { diff --git a/src/models/mamba-base.cpp b/src/models/mamba-base.cpp index fd3fe3f03..1f994ae0a 100644 --- a/src/models/mamba-base.cpp +++ b/src/models/mamba-base.cpp @@ -2,6 +2,8 @@ #include "llama-memory-recurrent.h" +#include + 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(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(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); diff --git a/src/models/plamo2.cpp b/src/models/plamo2.cpp index 0b81513c3..d946b3cff 100644 --- a/src/models/plamo2.cpp +++ b/src/models/plamo2.cpp @@ -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); diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 419e1eba4..08c6f5a47 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -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) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 08c29eec6..3349a64b1 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -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 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> 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));