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:
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;
}
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));
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;
}
}
}
+ 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;
}
(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;
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;
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
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
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;
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);
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));
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;
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
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);
}
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) {
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) ||
return -1;
}
- if (n_group <= 0 || n_head % n_group != 0) {
+ if (K < 1 || n_group <= 0 || n_head % n_group != 0) {
return -1;
}
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;
}
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);
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 {
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;
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;
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,
/*.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,
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;
// 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;
}
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;
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
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!
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;
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)");
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));
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) {
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;
(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);
} 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) {
uint d_head;
uint n_group;
uint n_tok;
+ uint n_seq;
+ uint K;
};
float softplus(float x) {
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
(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 = {
n_seq_tokens: u32,
y_elems: u32,
+ K: u32,
};
@group(0) @binding(0) var<storage, read_write> s_in: array<f32>;
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);
#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));
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));
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;
result->src[5] = C;
result->src[6] = ids;
+ ggml_set_op_params_i32(result, 0, (int32_t) K);
+
return result;
}
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;
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;
}
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:
{
#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,
// 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);
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;
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());
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);
// => {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
// (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);
// 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);
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)
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,
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.
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;
}
}
};
+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;
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));