]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
ggml-cuda: add chunked SSD matmul for Mamba-2 prefill acceleration (#22675)
authorBhavik Sharda <redacted>
Tue, 28 Jul 2026 12:03:42 +0000 (17:33 +0530)
committerGitHub <redacted>
Tue, 28 Jul 2026 12:03:42 +0000 (17:33 +0530)
* ggml-cuda: add chunked SSD matmul for Mamba-2 prefill acceleration

* cuda: added SSD CICD fixes for CUDA / HIP / MUSA / MSVC.

* ggml-cuda: review comments fixed.

* ggml-cuda: Fuse M matrix materialization into pre_matmul kernel and enabled test.

* ggml-cuda: test updates and fixes

* ggml-cuda: test updates to remove hardcoding of tensor initialise data limits.

* ggml-cuda: ssd minor review comment fixed.

* ggml-cuda: ssd minor CICD fixed.

* CUDA SSD: Fixes correctness by promoting s0_stride_seq to int64_t, improves memory coalescing in ssm_ssd_prepare_dt_kernel, and boosts efficiency by merging B_weighted and C_scaled; also addresses prior review comments.

* cuda: fix sdata read-write race in prepare_dt fallback scan loop

ggml/src/ggml-cuda/ssm-scan.cu
tests/test-backend-ops.cpp

index 3022249c77d5b0f8f81ede471633216b65b089e4..f3418c2af83d9e07fafff767a90a7b5a90858139 100644 (file)
@@ -9,6 +9,21 @@ using namespace cub;
 
 #include "ssm-scan.cuh"
 
+
+// Minimum number of tokens to use SSD (State Space Duality) matmul path instead of scan path.
+// For n_tok <= this threshold, the scan kernel is used (lower overhead for short sequences).
+#define SSM_SSD_MIN_TOKENS 128
+
+// prepare_dt kernel dimensions: one block per (head, seq), each block handles DT_MAX_ITEMS items.
+#define SSM_SSD_DT_BLOCK     256
+#define SSM_SSD_DT_MAX_ITEMS  32
+
+// Maximum tokens the SSD path supports, derived from the prepare_dt kernel block capacity.
+#define SSM_SSD_MAX_TOKENS (SSM_SSD_DT_BLOCK * SSM_SSD_DT_MAX_ITEMS)
+
+// Chunk size for chunked SSD. Caps matmul cost at O(chunk^2) per chunk.
+#define SSM_SSD_CHUNK_SIZE 256
+
 // We would like to keep pragma unroll for cases where L_template is not 0,
 // so we suppress the clang transformation warning.
 #ifdef __clang__
@@ -316,6 +331,429 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
     }
 }
 
+#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+// ============================================================================
+// SSD (State Space Duality) kernels for Mamba-2 prefill (n_tok > SSM_SSD_MIN_TOKENS)
+//
+// Instead of a sequential scan, SSD reformulates the output as:
+//   Y = (L (.) (C @ B^T)) @ (X * dt)  +  decay * C @ s_init
+// where L is a causal decay mask derived from A and dt.
+//
+// This converts the O(T*N) sequential scan into parallel matmuls.
+// ============================================================================
+// Softplus(dt) and inclusive prefix sum per head using CUB BlockScan.
+// Grid: (n_head, n_seqs)
+template <int BLOCK_SIZE, int MAX_ITEMS>
+__global__ void ssm_ssd_prepare_dt_kernel(
+        const float * __restrict__ dt_raw,
+        float * __restrict__ dt_sp_out,
+        float * __restrict__ cs_out,
+        const int n_head, const int n_tok,
+        const int dt_stride_tok,   // elements between tokens in dt
+        const int dt_stride_seq) { // elements between sequences in dt
+
+    const int h = blockIdx.x;
+    const int s = blockIdx.y;
+
+    const float * dt_seq = dt_raw + s * dt_stride_seq;
+
+    float * dt_sp_seq = dt_sp_out + s * n_tok * n_head;
+    float * cs_seq    = cs_out    + s * n_tok * n_head;
+
+    const int items_per_thread = (n_tok + BLOCK_SIZE - 1) / BLOCK_SIZE;
+
+    // Phase 1: softplus with interleaved distribution (t = i*BLOCK_SIZE + threadIdx.x).
+    // Each warp reads BLOCK_SIZE consecutive tokens, giving coalesced dt_raw loads
+    // (stride n_head between threads vs. items_per_thread*n_head in blocked layout).
+    float local_vals[MAX_ITEMS];
+    for (int i = 0; i < items_per_thread; i++) {
+        const int t = i * BLOCK_SIZE + threadIdx.x;
+        if (t < n_tok) {
+            float val = dt_seq[h + t * dt_stride_tok];
+            float sp = (val <= 20.0f) ? log1pf(expf(val)) : val;
+            local_vals[i] = sp;
+            dt_sp_seq[t * n_head + h] = sp;
+        } else {
+            local_vals[i] = 0.0f;
+        }
+    }
+
+    // Phase 2+3: per-step inclusive scan to build cs[] in token order.
+    // With interleaved distribution the per-thread total scan would not give token-order
+    // prefix sums, so we scan one BLOCK_SIZE slab at a time and carry a running total.
+#ifdef USE_CUB
+    using BlockScan = cub::BlockScan<float, BLOCK_SIZE>;
+    __shared__ typename BlockScan::TempStorage scan_temp;
+    __shared__ float step_total;
+
+    float running = 0.0f;
+    for (int i = 0; i < items_per_thread; i++) {
+        float inclusive;
+        BlockScan(scan_temp).InclusiveSum(local_vals[i], inclusive);
+        const int t = i * BLOCK_SIZE + threadIdx.x;
+        if (t < n_tok) {
+            cs_seq[t * n_head + h] = running + inclusive;
+        }
+        if (threadIdx.x == BLOCK_SIZE - 1) {
+            step_total = inclusive;
+        }
+        __syncthreads();
+        running += step_total;
+    }
+#else
+    // Fallback: sequential prefix scan in shared memory, one slab at a time.
+    __shared__ float sdata[BLOCK_SIZE];
+    float running = 0.0f;
+    for (int i = 0; i < items_per_thread; i++) {
+        const int t = i * BLOCK_SIZE + threadIdx.x;
+        sdata[threadIdx.x] = local_vals[i];
+        __syncthreads();
+        if (threadIdx.x == 0) {
+            for (int j = 1; j < BLOCK_SIZE; j++) {
+                sdata[j] += sdata[j - 1];
+            }
+        }
+        __syncthreads();
+        if (t < n_tok) {
+            cs_seq[t * n_head + h] = running + sdata[threadIdx.x];
+        }
+        running += sdata[BLOCK_SIZE - 1];
+        __syncthreads();
+    }
+#endif
+}
+
+// Prepare SSD matmul inputs for one chunk: X_dt, B_weighted, C_scaled.
+// T_matmul controls precision for X_dt, B_weighted (float or half).
+// C_scaled is always float (pairs with float s_cur in step 3c).
+// Computation is always FP32; only the final store converts to T_matmul.
+// Also materializes the causal M matrix = exp(A*(cs_out - cs_in)) * CB (fused with prep to save a launch).
+// Grid: (ceil(max(C*head_dim, d_state*C, chunk_len^2) / BLOCK), n_head, n_seqs)
+template <int BLOCK_SIZE, typename T_matmul>
+__global__ void ssm_ssd_pre_matmul_kernel(
+        const float * __restrict__ cs,         // {n_tok, n_head} cumulative dt sums
+        const float * __restrict__ dt_sp,      // {n_tok, n_head} softplus(dt)
+        const float * __restrict__ A,          // {1, n_head}
+        const float * __restrict__ x,          // {head_dim, n_head, n_tok, n_seqs}
+        const float * __restrict__ B,          // {d_state, n_group, n_tok, n_seqs}
+        const float * __restrict__ C_src,      // {d_state, n_group, n_tok, n_seqs}
+        T_matmul * __restrict__ X_dt,          // {head_dim, C, n_head} x * dt, d-fastest
+        T_matmul * __restrict__ B_weighted,    // {d_state, C, n_head} B * decay_from_end
+        float * __restrict__ C_scaled,         // {d_state, C, n_head} C * decay_to_pos (always float)
+        const float * __restrict__ CB,         // {chunk_len, chunk_len, n_group, n_seqs}
+        half * __restrict__ M_out,             // {chunk_len, chunk_len, n_head, n_seqs}
+        const int chunk_len, const int head_dim, const int n_head, const int n_group,
+        const int d_state, const int A_stride,
+        const int x_stride_tok, const int x_stride_seq,
+        const int B_stride_tok, const int B_stride_seq,
+        const int C_stride_tok, const int C_stride_seq,
+        const int chunk_offset,
+        const int n_tok_total) {
+
+    const int h = blockIdx.y;
+    const int s = blockIdx.z;
+    const int g = h / (n_head / n_group);
+
+    const float A_h = A[h * A_stride];
+    const int idx = blockIdx.x * BLOCK_SIZE + threadIdx.x;
+
+    const int cs_seq_off = s * n_tok_total * n_head;
+    const float cs_base = (chunk_offset > 0) ? cs[cs_seq_off + (chunk_offset - 1) * n_head + h] : 0.0f;
+    const float cs_last = cs[cs_seq_off + (chunk_offset + chunk_len - 1) * n_head + h] - cs_base;
+
+    // Prepare X_dt = x * dt, stored d-fastest for coalesced reads and writes.
+    const int n_xdt = chunk_len * head_dim;
+    if (idx < n_xdt) {
+        const int d = idx % head_dim;
+        const int t = idx / head_dim;
+
+        const float x_val = x[s * x_stride_seq + (chunk_offset + t) * x_stride_tok + d + h * head_dim];
+        const float dt_val = dt_sp[cs_seq_off + (chunk_offset + t) * n_head + h];
+
+        X_dt[d + t * head_dim + h * n_xdt + s * n_xdt * n_head] = (T_matmul)(x_val * dt_val);
+    }
+
+    // Prepare B_weighted and C_scaled together: both share the same index space (d_state * chunk_len)
+    // and the same cs_t load, so merging halves the cs[] global memory traffic.
+    const int n_bw = d_state * chunk_len;
+    if (idx < n_bw) {
+        const int n = idx % d_state;
+        const int t = idx / d_state;
+
+        const float cs_t = cs[cs_seq_off + (chunk_offset + t) * n_head + h] - cs_base;
+
+        const float B_val = B[s * B_stride_seq + (chunk_offset + t) * B_stride_tok + g * d_state + n];
+        B_weighted[n + t * d_state + h * n_bw + s * n_bw * n_head] = (T_matmul)(B_val * __expf(A_h * (cs_last - cs_t)));
+
+        const float C_val = C_src[s * C_stride_seq + (chunk_offset + t) * C_stride_tok + g * d_state + n];
+        C_scaled[n + t * d_state + h * n_bw + s * n_bw * n_head] = C_val * __expf(A_h * cs_t);
+    }
+
+    // Materialize M = exp(A*(cs_out - cs_in)) * CB with causal mask.
+    const int n_M = chunk_len * chunk_len;
+    if (idx < n_M) {
+        const int t_out = idx % chunk_len;
+        const int t_in  = idx / chunk_len;
+
+        half val;
+        if (t_in <= t_out) {
+            const float cs_out = cs[cs_seq_off + (chunk_offset + t_out) * n_head + h] - cs_base;
+            const float cs_in  = cs[cs_seq_off + (chunk_offset + t_in)  * n_head + h] - cs_base;
+            const float decay  = __expf(A_h * (cs_out - cs_in));
+            const float * CB_g = CB + (int64_t)s * chunk_len * chunk_len * n_group
+                                   + (int64_t)g * chunk_len * chunk_len;
+            const float cb_val = CB_g[t_out + t_in * chunk_len];
+            val = __float2half(decay * cb_val);
+        } else {
+            val = __float2half(0.0f);
+        }
+
+        M_out[(int64_t)s * n_M * n_head + (int64_t)h * n_M + t_in * chunk_len + t_out] = val;
+    }
+}
+
+// Scale running state in-place: s_cur *= decay_total(chunk).
+// Called BEFORE cuBLAS state update (beta=1) to fuse inter-chunk decay.
+// Eliminates the s_old buffer and D2D memcpy vs the old approach of:
+//   memcpy(s_old, s_cur) -> cuBLAS(beta=0) -> s_cur += decay * s_old
+// Grid: (ceil(d_state * head_dim / BLOCK), n_head, n_seqs)
+template <int BLOCK_SIZE>
+__global__ void ssm_ssd_scale_state_kernel(
+        float * __restrict__ s_cur,            // {d_state, head_dim, n_head, n_seqs}
+        const float * __restrict__ cs,         // {n_tok, n_head} cumulative dt sums
+        const float * __restrict__ A,          // {1, n_head}
+        const int d_state, const int head_dim, const int n_head,
+        const int chunk_offset, const int chunk_len,
+        const int n_tok_total, const int A_stride) {
+
+    const int h = blockIdx.y;
+    const int s = blockIdx.z;
+    const int idx = blockIdx.x * BLOCK_SIZE + threadIdx.x;
+    const int state_per_head = d_state * head_dim;
+    if (idx >= state_per_head) return;
+
+    const float A_h = A[h * A_stride];
+    const int cs_seq_off = s * n_tok_total * n_head;
+    const float cs_base = (chunk_offset > 0) ? cs[cs_seq_off + (chunk_offset - 1) * n_head + h] : 0.0f;
+    const float cs_last = cs[cs_seq_off + (chunk_offset + chunk_len - 1) * n_head + h] - cs_base;
+    const float decay_total = __expf(A_h * cs_last);
+
+    const int off = s * state_per_head * n_head + h * state_per_head + idx;
+    s_cur[off] *= decay_total;
+}
+
+// Copy initial state from src0[ids[s]] into s_cur for each sequence.
+// 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 int32_t * __restrict__ ids,      // {n_seqs}
+        float * __restrict__ 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;
+    const int idx = blockIdx.x * BLOCK_SIZE + threadIdx.x;
+    if (idx >= state_size) return;
+
+    const float * s_src = src0 + (int64_t)ids[s] * s0_stride_seq;
+    s_cur[s * state_size + idx] = s_src[idx];
+}
+
+// SSD (State Space Duality) dispatch for Mamba-2 prefill.
+// Chunked matmuls: CB, materialize M + cuBLAS Y, S@C, B@X_dt.
+// All strides are in elements (floats), not bytes.
+static void ssm_scan_ssd_f32_cuda(
+        ggml_backend_cuda_context & ctx,
+        const float * src0_d, const float * src1_d, const float * src2_d, const float * src3_d,
+        const float * src4_d, const float * src5_d, const int32_t * src6_d, float * dst_d,
+        const int64_t s0_stride_seq,                                   // state (src0) stride between seqs
+        const int x_stride_tok,  const int x_stride_seq,               // x (src1) strides
+        const int dt_stride_tok, const int dt_stride_seq,              // dt (src2) strides
+        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,
+        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();
+    const int64_t d_inner = head_dim * n_head;
+
+    const int64_t chunk_size = SSM_SSD_CHUNK_SIZE;
+    const int64_t n_chunks = (n_tok + chunk_size - 1) / chunk_size;
+
+    const int64_t state_per_head = d_state * head_dim;
+
+    using matmul_t = half;
+    static constexpr cudaDataType_t matmul_dtype = CUDA_R_16F;
+
+    ggml_cuda_pool_alloc<float>    dt_sp_buf(ctx.pool(), n_tok * n_head * n_seq);
+    ggml_cuda_pool_alloc<float>    cs_buf(ctx.pool(), n_tok * n_head * n_seq);
+    ggml_cuda_pool_alloc<float>    CB_buf(ctx.pool(), chunk_size * chunk_size * n_group * n_seq);
+    ggml_cuda_pool_alloc<matmul_t> X_dt_buf(ctx.pool(), chunk_size * head_dim * n_head * n_seq);
+    ggml_cuda_pool_alloc<matmul_t> B_w_buf(ctx.pool(), d_state * chunk_size * n_head * n_seq);
+    ggml_cuda_pool_alloc<float>    C_s_buf(ctx.pool(), d_state * chunk_size * n_head * n_seq);
+    float    * dt_sp      = dt_sp_buf.get();
+    float    * cs         = cs_buf.get();
+    float    * CB         = CB_buf.get();
+    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
+    {
+        dim3 grid(n_head, n_seq);
+        ssm_ssd_prepare_dt_kernel<SSM_SSD_DT_BLOCK, SSM_SSD_DT_MAX_ITEMS><<<grid, SSM_SSD_DT_BLOCK, 0, stream>>>(
+            src2_d, dt_sp, cs, n_head, n_tok, dt_stride_tok, dt_stride_seq);
+        CUDA_CHECK(cudaGetLastError());
+    }
+
+    // Step 2: initialize running state from src0[ids[s]]
+    {
+        constexpr int BLOCK = 256;
+        const int64_t state_size = d_state * head_dim * n_head;
+        dim3 grid((state_size + BLOCK - 1) / BLOCK, n_seq);
+        ssm_ssd_init_state_kernel<BLOCK><<<grid, BLOCK, 0, stream>>>(
+            src0_d, src6_d, s_cur, state_size, s0_stride_seq);
+        CUDA_CHECK(cudaGetLastError());
+    }
+
+    // Step 3: chunked SSD loop
+    // Per chunk: pre_matmul (incl. M) + 4 cuBLAS (CB, Y, S@C, state update) + scale_state
+    cublasHandle_t handle = ctx.cublas_handle();
+    CUBLAS_CHECK(cublasSetStream(handle, stream));
+    const float alpha_one  = 1.0f;
+    const float beta_zero  = 0.0f;
+    const float beta_one   = 1.0f;
+    const int lda_C_src = C_stride_tok;  // leading dim for C in CB = C^T @ B
+    const int ldb_B_src = B_stride_tok;  // leading dim for B in CB = C^T @ B
+
+    // Scratch buffer for causal M matrix, reused across chunks (max size at chunk_size)
+    const int64_t n_M_max = chunk_size * chunk_size;
+    ggml_cuda_pool_alloc<half> M_buf(ctx.pool(), n_M_max * n_head * n_seq);
+    half * M_mat = M_buf.get();
+
+    for (int64_t k = 0; k < n_chunks; k++) {
+        const int64_t chunk_offset = k * chunk_size;
+        const int64_t chunk_len = (chunk_offset + chunk_size <= n_tok) ? chunk_size : (n_tok - chunk_offset);
+
+        // 3a: CB = C^T @ B per group
+        for (int64_t s = 0; s < n_seq; s++) {
+            const float * C_s = src5_d + s * C_stride_seq + chunk_offset * C_stride_tok;
+            const float * B_s = src4_d + s * B_stride_seq + chunk_offset * B_stride_tok;
+            float *      CB_s = CB + s * chunk_len * chunk_len * n_group;
+
+            if (n_group == 1) {
+                CUBLAS_CHECK(cublasSgemm(handle, CUBLAS_OP_T, CUBLAS_OP_N,
+                    chunk_len, chunk_len, d_state,
+                    &alpha_one, C_s, lda_C_src, B_s, ldb_B_src,
+                    &beta_zero, CB_s, (int)chunk_len));
+            } else {
+                CUBLAS_CHECK(cublasGemmStridedBatchedEx(handle, CUBLAS_OP_T, CUBLAS_OP_N,
+                    chunk_len, chunk_len, d_state,
+                    &alpha_one,
+                    C_s, CUDA_R_32F, lda_C_src, d_state,
+                    B_s, CUDA_R_32F, ldb_B_src, d_state,
+                    &beta_zero,
+                    CB_s, CUDA_R_32F, (int)chunk_len, (long long)(chunk_len * chunk_len),
+                    n_group,
+                    CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
+            }
+        }
+
+        // 3b: prepare X_dt, B_weighted, C_scaled + materialize causal M matrix
+        const int64_t n_M = chunk_len * chunk_len;
+        {
+            constexpr int BLOCK = 256;
+            const int64_t n_xdt   = chunk_len * head_dim;
+            const int64_t n_bw    = d_state * chunk_len;
+            int64_t max_work = n_xdt;
+            if (n_bw  > max_work) max_work = n_bw;
+            if (n_M   > max_work) max_work = n_M;
+            dim3 grid((max_work + BLOCK - 1) / BLOCK, n_head, n_seq);
+            ssm_ssd_pre_matmul_kernel<BLOCK, matmul_t><<<grid, BLOCK, 0, stream>>>(
+                cs, dt_sp, src3_d, src1_d, src4_d, src5_d,
+                X_dt, B_weighted, C_scaled,
+                CB, M_mat,
+                chunk_len, head_dim, n_head, n_group, d_state, A_stride,
+                x_stride_tok, x_stride_seq, B_stride_tok, B_stride_seq, C_stride_tok, C_stride_seq,
+                chunk_offset, n_tok);
+            CUDA_CHECK(cudaGetLastError());
+        }
+
+        // 3c: dst = S_cur^T @ C_scaled (state contribution)
+        {
+            const int64_t stride_S  = state_per_head;
+            const int64_t stride_Cs = d_state * chunk_len;
+
+            for (int64_t s = 0; s < n_seq; s++) {
+                float * dst_chunk = dst_d + s * d_inner * n_tok + chunk_offset * d_inner;
+
+                CUBLAS_CHECK(cublasGemmStridedBatchedEx(handle, CUBLAS_OP_T, CUBLAS_OP_N,
+                    head_dim, chunk_len, d_state,
+                    &alpha_one,
+                    s_cur    + s * stride_S  * n_head, CUDA_R_32F, d_state, stride_S,
+                    C_scaled + s * stride_Cs * n_head, CUDA_R_32F, d_state, stride_Cs,
+                    &beta_zero,
+                    dst_chunk, CUDA_R_32F, d_inner, head_dim,
+                    n_head,
+                    CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
+            }
+        }
+
+        // 3d: dst += X_dt @ M^T (intra-chunk contribution, adds to 3c result)
+        // M is stored as M[t_out, t_in] (lower-triangular), transpose needed for Y = X @ M^T.
+        {
+            const int64_t stride_M = n_M;
+            const int64_t stride_X_h = (int64_t)chunk_len * head_dim;
+
+            for (int64_t s = 0; s < n_seq; s++) {
+                float * dst_chunk = dst_d + s * d_inner * n_tok + chunk_offset * d_inner;
+                CUBLAS_CHECK(cublasGemmStridedBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_T,
+                    head_dim, chunk_len, chunk_len,
+                    &alpha_one,
+                    X_dt       + s * stride_X_h * n_head, matmul_dtype, head_dim, stride_X_h,
+                    M_mat      + s * stride_M   * n_head, matmul_dtype, chunk_len, stride_M,
+                    &beta_one,
+                    dst_chunk, CUDA_R_32F, d_inner, head_dim,
+                    n_head,
+                    CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
+            }
+        }
+
+        // 3e: s_cur = B_weighted @ X_dt^T + decay_total * s_cur_old (state update)
+        {
+            // Scale s_cur in-place by per-head decay_total BEFORE cuBLAS overwrites it
+            constexpr int BLOCK = 256;
+            dim3 grid((state_per_head + BLOCK - 1) / BLOCK, n_head, n_seq);
+            ssm_ssd_scale_state_kernel<BLOCK><<<grid, BLOCK, 0, stream>>>(
+                s_cur, cs, src3_d,
+                d_state, head_dim, n_head,
+                chunk_offset, chunk_len, n_tok, A_stride);
+            CUDA_CHECK(cudaGetLastError());
+
+            // cuBLAS with beta=1: s_cur = B_weighted @ X_dt^T + 1.0 * s_cur (already scaled)
+            const int64_t stride_Bw = d_state * chunk_len;
+            const int64_t stride_X  = chunk_len * head_dim;
+            const int64_t stride_S  = state_per_head;
+
+            for (int64_t s = 0; s < n_seq; s++) {
+                // X_dt is d-fastest {hd, C}, read as OP_T to get {C, hd}
+                CUBLAS_CHECK(cublasGemmStridedBatchedEx(handle, CUBLAS_OP_N, CUBLAS_OP_T,
+                    d_state, head_dim, chunk_len,
+                    &alpha_one,
+                    B_weighted + s * stride_Bw * n_head, matmul_dtype, d_state, stride_Bw,
+                    X_dt       + s * stride_X  * n_head, matmul_dtype, head_dim, stride_X,
+                    &beta_one,
+                    s_cur      + s * stride_S  * n_head, CUDA_R_32F, d_state, stride_S,
+                    n_head,
+                    CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT));
+            }
+        }
+    }
+}
+#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+
 void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     const struct ggml_tensor * src0 = dst->src[0];  // s
     const struct ggml_tensor * src1 = dst->src[1];  // x
@@ -357,6 +795,49 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     GGML_ASSERT(src6->type == GGML_TYPE_I32);
     GGML_ASSERT(dst->type == GGML_TYPE_F32);
 
+    // Byte strides are narrowed to int for both scan and SSD paths.
+    GGML_ASSERT(src0->nb[2] <= (size_t)INT_MAX);
+    GGML_ASSERT(src0->nb[3] <= (size_t)INT_MAX);
+    GGML_ASSERT(src1->nb[2] <= (size_t)INT_MAX);
+    GGML_ASSERT(src1->nb[3] <= (size_t)INT_MAX);
+    GGML_ASSERT(src2->nb[1] <= (size_t)INT_MAX);
+    GGML_ASSERT(src2->nb[2] <= (size_t)INT_MAX);
+    GGML_ASSERT(src3->nb[1] <= (size_t)INT_MAX);
+    GGML_ASSERT(src4->nb[2] <= (size_t)INT_MAX);
+    GGML_ASSERT(src4->nb[3] <= (size_t)INT_MAX);
+    GGML_ASSERT(src5->nb[2] <= (size_t)INT_MAX);
+    GGML_ASSERT(src5->nb[3] <= (size_t)INT_MAX);
+
+#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+    // Mamba-2 with scalar A per head: use SSD matmul path for long sequences.
+    // Requires NVIDIA Turing+ otherwise fallback to scan.
+    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
+                      && n_t <= SSM_SSD_MAX_TOKENS
+                      && GGML_CUDA_CC_IS_NVIDIA(cc)
+                      && cc >= GGML_CUDA_CC_TURING
+                      && nr % 8 == 0;  // cuBLAS requires 8-element (16-byte) alignment
+
+    if (use_ssd) {
+        // ssm_ssd_init_state_kernel uses flat linear indexing within each sequence,
+        // so src0 must be fully contiguous across all inner dimensions.
+        // The scan path handles non-contiguous nb[2] via src0_nb2 but does not handle nb[1].
+        GGML_ASSERT(src0->nb[1] == nc         * sizeof(float));
+        GGML_ASSERT(src0->nb[2] == nc * nr    * sizeof(float));
+
+        ssm_scan_ssd_f32_cuda(ctx,
+            src0_d, src1_d, src2_d, src3_d, src4_d, src5_d, src6_d, dst_d,
+            (int64_t)(src0->nb[3] / sizeof(float)),
+            (int)(src1->nb[2] / sizeof(float)), (int)(src1->nb[3] / sizeof(float)),
+            (int)(src2->nb[1] / sizeof(float)), (int)(src2->nb[2] / sizeof(float)),
+            (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);
+        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],
index f6b60e52a4003f43d79795edbeb433b01e38306c..a5b660f47a0a60c0caf1159dfe8d878f3989633f 100644 (file)
@@ -4000,7 +4000,7 @@ struct test_ssm_scan : public test_case {
 
     test_ssm_scan(ggml_type type = GGML_TYPE_F32,
             int64_t d_state = 32,
-            int64_t head_dim = 1, // non-zero for Mamba-2
+            int64_t head_dim = 1, // 1 = Mamba-1; > 1 = Mamba-2 (scalar A per head)
             int64_t n_head  = 32,
             int64_t n_group = 1,
             int64_t n_seq_tokens = 32,
@@ -4008,6 +4008,11 @@ struct test_ssm_scan : public test_case {
             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) {}
 
+    double max_nmse_err() override {
+        // SSD path (head_dim > 1) uses FP16 intermediates (M matrix, X_dt); Mamba-1 is pure FP32.
+        return (head_dim > 1) ? 2e-7 : 1e-7;
+    }
+
     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 * dt  = ggml_new_tensor_3d(ctx, type, n_head,   n_seq_tokens, n_seqs);
@@ -4034,14 +4039,14 @@ struct test_ssm_scan : public test_case {
         return out;
     }
 
-    // similar to test_mul_mat_id
+
     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; }
-                // ids
+                // ids: permutation of [0..n_seqs)
                 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++) {
@@ -4050,6 +4055,11 @@ struct test_ssm_scan : public test_case {
                     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) {
+                // A {1 or d_state, n_head}: negative decay (2-D tensor, ne[2]==1 distinguishes from 3-D/4-D tensors)
+                init_tensor_uniform(t, -1.0f, -0.5f);
             } else {
                 init_tensor_uniform(t);
             }
@@ -8770,6 +8780,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, 64, 16, 2, 32, 4)); // Mamba-2
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 256, 64,  8, 2, 32, 4)); // Falcon-H1
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 128, 4, 4, 16, 2, true)); // x/B/C overlap
+    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_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 1, 1));
     test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 32, 1));
@@ -9990,6 +10003,8 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
     test_cases.emplace_back(new test_ssm_conv_bias_silu(GGML_TYPE_F32, {4,   3328, 1, 1}, {4, 3328, 1, 1}, true));  // generate
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 48, 1, 512, 1)); // prefill
     test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 48, 1, 1,   1)); // generate
+    test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 80, 128, 1, 512, 1)); // Nemotron-9B prefill
+    test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 80, 128, 1, 1,   1)); // Nemotron-9B generate
 
     // acc
     test_cases.emplace_back(new test_acc(GGML_TYPE_F32, {256, 17, 1, 1}, {256, 16, 1, 1}, -1));