#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__
}
}
+#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
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],