Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 24 additions & 5 deletions src/llama-graph.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -360,6 +360,8 @@ bool llm_graph_input_rs::can_reuse(const llm_graph_params & params) {
res &= head == mctx->get_head();
res &= rs_z == mctx->get_rs_z();

res &= s_copy_main_identity == mctx->is_s_copy_main_identity(params.ubatch.n_seqs);

return res;
}

Expand Down Expand Up @@ -1156,6 +1158,8 @@ bool llm_graph_input_mem_hybrid::can_reuse(const llm_graph_params & params) {
res &= inp_rs->head == mctx->get_recr()->get_head();
res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();

res &= inp_rs->s_copy_main_identity == mctx->get_recr()->is_s_copy_main_identity(params.ubatch.n_seqs);

return res;
}

Expand Down Expand Up @@ -1204,6 +1208,8 @@ bool llm_graph_input_mem_hybrid_k::can_reuse(const llm_graph_params & params) {
res &= inp_rs->head == mctx->get_recr()->get_head();
res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();

res &= inp_rs->s_copy_main_identity == mctx->get_recr()->is_s_copy_main_identity(params.ubatch.n_seqs);

return res;
}

Expand Down Expand Up @@ -1294,6 +1300,8 @@ bool llm_graph_input_mem_hybrid_iswa::can_reuse(const llm_graph_params & params)
res &= inp_rs->head == mctx->get_recr()->get_head();
res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();

res &= inp_rs->s_copy_main_identity == mctx->get_recr()->is_s_copy_main_identity(params.ubatch.n_seqs);

return res;
}

Expand Down Expand Up @@ -3519,7 +3527,8 @@ ggml_tensor * llm_graph_context::build_rs(
uint32_t rs_head,
uint32_t rs_size,
int32_t rs_zero,
const llm_graph_get_rows_fn & get_state_rows) const {
const llm_graph_get_rows_fn & get_state_rows,
bool main_inplace) const {

GGML_UNUSED(rs_size);
ggml_tensor * states = ggml_reshape_2d(ctx0, s, state_size, s->ne[1]);
Expand All @@ -3532,8 +3541,14 @@ ggml_tensor * llm_graph_context::build_rs(
// copy states
// NOTE: assuming the copy destinations are ALL contained between rs_head and rs_head + n_rs
// {state_size, rs_size} -> {state_size, n_seqs}
ggml_tensor * output_states = get_state_rows(ctx0, states, state_copy_main);
ggml_build_forward_expand(gf, output_states);
ggml_tensor * output_states;
if (main_inplace) {
// rows are already in place; not expanded here, so the view stays next to its consumer
output_states = ggml_view_2d(ctx0, states, state_size, n_seqs, states->nb[1], rs_head*states->nb[1]);
} else {
output_states = get_state_rows(ctx0, states, state_copy_main);
ggml_build_forward_expand(gf, output_states);
}

// copy extra states which won't be changed further (between n_seqs and n_rs)
ggml_tensor * states_extra = ggml_get_rows(ctx0, states, state_copy_extra);
Expand Down Expand Up @@ -3564,6 +3579,7 @@ static std::unique_ptr<llm_graph_input_rs> build_rs_inp_impl(

inp->head = mctx_cur->get_head();
inp->rs_z = mctx_cur->get_rs_z();
inp->s_copy_main_identity = mctx_cur->is_s_copy_main_identity(n_seqs);

return inp;
}
Expand All @@ -3581,12 +3597,15 @@ ggml_tensor * llm_graph_context::build_rs(
ggml_tensor * s,
int32_t state_size,
int32_t n_seqs,
const llm_graph_get_rows_fn & get_state_rows) const {
const llm_graph_get_rows_fn & get_state_rows,
bool allow_inplace) const {
const auto * kv_state = inp->mctx;

const bool main_inplace = allow_inplace && inp->s_copy_main_identity;

return build_rs(s, inp->s_copy_main, inp->s_copy_extra, state_size, n_seqs,
kv_state->get_n_rs(), kv_state->get_head(), kv_state->get_size(), kv_state->get_rs_z(),
get_state_rows);
get_state_rows, main_inplace);
}

ggml_tensor * llm_graph_context::build_rwkv_token_shift_load(
Expand Down
11 changes: 9 additions & 2 deletions src/llama-graph.h
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,9 @@ class llm_graph_input_rs : public llm_graph_input_i {
// used in view offsets, need to match for valid graph reuse
uint32_t head;
int32_t rs_z;

// part of the reuse key: a seq_cp between two ubatches can change it
bool s_copy_main_identity = false;
};

class llm_graph_input_cross_embd : public llm_graph_input_i {
Expand Down Expand Up @@ -1339,16 +1342,20 @@ struct llm_graph_context {
uint32_t rs_head,
uint32_t rs_size,
int32_t rs_zero,
const llm_graph_get_rows_fn & get_state_rows = ggml_get_rows) const;
const llm_graph_get_rows_fn & get_state_rows = ggml_get_rows,
bool main_inplace = false) const;

llm_graph_input_rs * build_rs_inp() const;

// allow_inplace: when no state is copied, return a view of the cache rows instead of a get_rows copy
// the consumer must read the whole state before the write-back cpy
ggml_tensor * build_rs(
llm_graph_input_rs * inp,
ggml_tensor * s,
int32_t state_size,
int32_t n_seqs,
const llm_graph_get_rows_fn & get_state_rows = ggml_get_rows) const;
const llm_graph_get_rows_fn & get_state_rows = ggml_get_rows,
bool allow_inplace = false) const;

ggml_tensor * build_rwkv_token_shift_load(
llm_graph_input_rs * inp,
Expand Down
14 changes: 14 additions & 0 deletions src/llama-memory-recurrent.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1313,6 +1313,20 @@ ggml_tensor * llama_memory_recurrent_context::get_p_l(int32_t il) const {
return mem->p_l[il];
}

bool llama_memory_recurrent_context::is_s_copy_main_identity(uint32_t n_seqs) const {
// with n_rs_seq == 0, s_copy(i) is cells[head + i].src0
if (is_full || mem->n_rs_seq > 0 || n_seqs > get_n_rs()) {
return false;
}
for (uint32_t i = 0; i < n_seqs; ++i) {
const uint32_t cell = mem->head + i;
if (mem->cells[cell].src0 != (int32_t) cell) {
return false;
}
}
return true;
}

int32_t llama_memory_recurrent_context::s_copy(int i) const {
const uint32_t cell_idx = i + mem->head;
const int32_t src0 = mem->cells[cell_idx].src0;
Expand Down
3 changes: 3 additions & 0 deletions src/llama-memory-recurrent.h
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,9 @@ class llama_memory_recurrent_context : public llama_memory_context_i {

int32_t s_copy(int i) const;

// s_copy(i) == head + i for all i < n_seqs, without side effects; always false with n_rs_seq > 0
bool is_s_copy_main_identity(uint32_t n_seqs) const;

private:
const llama_memory_status status;

Expand Down
6 changes: 3 additions & 3 deletions src/models/qwen35.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -397,7 +397,7 @@ ggml_tensor * llama_model_qwen35::graph::build_layer_attn_linear(

ggml_tensor * conv_input = build_conv_state(inp, conv_states_all, qkv_mixed, conv_kernel_size, conv_channels, il);

ggml_tensor * state = build_rs(inp, ssm_states_all, hparams.n_embd_s(), n_seqs);
ggml_tensor * state = build_rs(inp, ssm_states_all, hparams.n_embd_s(), n_seqs, ggml_get_rows, /*allow_inplace=*/true);
state = ggml_reshape_4d(ctx0, state, head_v_dim, head_v_dim, num_v_heads, n_seqs);
cb(state, "state_predelta", il);

Expand Down Expand Up @@ -466,8 +466,8 @@ ggml_tensor * llama_model_qwen35::graph::build_layer_attn_linear(
// Apply gated normalization: self.norm(core_attn_out, z)
ggml_tensor * attn_out_norm = build_norm_gated(output, model.layers[il].ssm_norm, z_2d, il);

// Final reshape: [head_dim, n_heads, n_tokens, n_seqs] -> [n_tokens, n_seqs, n_heads * head_dim]
ggml_tensor * final_output = ggml_reshape_3d(ctx0, attn_out_norm, head_v_dim * num_v_heads, n_seq_tokens, n_seqs);
// 2-D: with a [d, n_seq_tokens, n_seqs] src1 the projection is n_seqs matmuls that each read the whole ssm_out weight
ggml_tensor * final_output = ggml_reshape_2d(ctx0, attn_out_norm, head_v_dim * num_v_heads, n_seq_tokens * n_seqs);
cb(final_output, "final_output", il);

// Output projection
Expand Down
6 changes: 3 additions & 3 deletions src/models/qwen35moe.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -414,7 +414,7 @@ ggml_tensor * llama_model_qwen35moe::graph::build_layer_attn_linear(

ggml_tensor * conv_input = build_conv_state(inp, conv_states_all, qkv_mixed, conv_kernel_size, conv_channels, il);

ggml_tensor * state = build_rs(inp, ssm_states_all, hparams.n_embd_s(), n_seqs);
ggml_tensor * state = build_rs(inp, ssm_states_all, hparams.n_embd_s(), n_seqs, ggml_get_rows, /*allow_inplace=*/true);
state = ggml_reshape_4d(ctx0, state, head_v_dim, head_v_dim, num_v_heads, n_seqs);
cb(state, "state_predelta", il);

Expand Down Expand Up @@ -483,8 +483,8 @@ ggml_tensor * llama_model_qwen35moe::graph::build_layer_attn_linear(
// Apply gated normalization: self.norm(core_attn_out, z)
ggml_tensor * attn_out_norm = build_norm_gated(output, model.layers[il].ssm_norm, z_2d, il);

// Final reshape: [head_dim, n_heads, n_tokens, n_seqs] -> [n_tokens, n_seqs, n_heads * head_dim]
ggml_tensor * final_output = ggml_reshape_3d(ctx0, attn_out_norm, head_v_dim * num_v_heads, n_seq_tokens, n_seqs);
// 2-D: with a [d, n_seq_tokens, n_seqs] src1 the projection is n_seqs matmuls that each read the whole ssm_out weight
ggml_tensor * final_output = ggml_reshape_2d(ctx0, attn_out_norm, head_v_dim * num_v_heads, n_seq_tokens * n_seqs);
cb(final_output, "final_output", il);

// Output projection
Expand Down
9 changes: 9 additions & 0 deletions tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -247,6 +247,15 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
set_tests_properties(test-recurrent-state-rollback-kimi-k3 PROPERTIES
FIXTURES_REQUIRED generate-models
)
llama_test(
test-recurrent-state-rollback
NAME test-recurrent-state-rollback-qwen35moe
LABEL main
ARGS -m "${MODEL_DIR}/qwen35moe-moe.gguf"
)
set_tests_properties(test-recurrent-state-rollback-qwen35moe PROPERTIES
FIXTURES_REQUIRED generate-models
)
# non-unified multi-seq decodes reach the QSA path with one stream per ubatch (#73)
llama_test(
test-recurrent-state-rollback
Expand Down
Loading