Commit 436f6f89e for llama.cpp
commit 436f6f89e1e581249900b37a5b8a12a36a6d0912
Author: Pascal <admin@serveurperso.com>
Date: Sat Oct 3 14:02:05 2026 +0200
graph: gather the recurrent states once so the reserve covers every split (#29856)
build_rs gathered the extra states (n_rs - n_seqs rows) with their own
get_rows. The worst-case reserve has n_rs == n_seqs, so that node was
sized at zero rows, and any ubatch whose cells are not contiguous forced
a graph reallocation at an unchanged node count, which aborts under
GGML_SCHED_NO_REALLOC.
A single get_rows now gathers the n_rs states: the ubatch states and the
extra states are views of it, and its size only depends on n_rs, which
the reserve already sets to the maximum. A custom getter (mamba ssm_scan)
gathers from the second state, so a single sequence ubatch copies no
state. The views are built once per graph in the input to keep the host
overhead of the graph unchanged.
diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp
index ae477c08d..16a3ed2ab 100644
--- a/src/llama-graph.cpp
+++ b/src/llama-graph.cpp
@@ -351,8 +351,7 @@ bool llm_graph_input_rs::can_reuse(const llm_graph_params & params) {
res &= s_copy->ne[0] == mctx->get_n_rs();
- res &= s_copy_main->ne[0] == params.ubatch.n_seqs;
- res &= s_copy_extra->ne[0] == mctx->get_n_rs() - params.ubatch.n_seqs;
+ res &= s_copy_main->ne[0] == params.ubatch.n_seqs;
res &= head == mctx->get_head();
res &= rs_z == mctx->get_rs_z();
@@ -1132,8 +1131,7 @@ bool llm_graph_input_mem_hybrid::can_reuse(const llm_graph_params & params) {
res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();
- res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;
- res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;
+ res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;
res &= inp_rs->head == mctx->get_recr()->get_head();
res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();
@@ -1175,8 +1173,7 @@ bool llm_graph_input_mem_hybrid_k::can_reuse(const llm_graph_params & params) {
res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();
- res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;
- res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;
+ res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;
res &= inp_rs->head == mctx->get_recr()->get_head();
res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();
@@ -1263,8 +1260,7 @@ bool llm_graph_input_mem_hybrid_iswa::can_reuse(const llm_graph_params & params)
res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();
- res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;
- res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;
+ res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;
res &= inp_rs->head == mctx->get_recr()->get_head();
res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();
@@ -3519,8 +3515,8 @@ llm_graph_input_dsv4 * llm_graph_context::build_inp_dsv4() const {
ggml_tensor * llm_graph_context::build_rs(
ggml_tensor * s,
+ ggml_tensor * state_copy,
ggml_tensor * state_copy_main,
- ggml_tensor * state_copy_extra,
int32_t state_size,
int32_t n_seqs,
uint32_t n_rs,
@@ -3539,12 +3535,19 @@ 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);
+ // one gather of the states i0..n_rs (ubatch states then extra states), sized by n_rs so the reserve covers every split
+ // {state_size, rs_size} -> {state_size, n_rs - i0}
+ const int64_t i0 = n_rs - state_copy->ne[0];
+
+ ggml_tensor * states_all = ggml_get_rows(ctx0, states, state_copy);
+
+ ggml_tensor * output_states = get_state_rows ?
+ get_state_rows(ctx0, states, state_copy_main) :
+ ggml_view_2d(ctx0, states_all, state_size, n_seqs, states_all->nb[1], 0);
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);
+ ggml_tensor * states_extra = ggml_view_2d(ctx0, states_all, state_size, n_rs - n_seqs, states_all->nb[1], (n_seqs - i0)*states_all->nb[1]);
ggml_build_forward_expand(gf,
ggml_cpy(ctx0,
states_extra,
@@ -3567,8 +3570,8 @@ static std::unique_ptr<llm_graph_input_rs> build_rs_inp_impl(
ggml_set_input(inp->s_copy);
ggml_set_name(inp->s_copy, "rs_s_copy");
- inp->s_copy_main = ggml_view_1d(ctx0, inp->s_copy, n_seqs, 0);
- inp->s_copy_extra = ggml_view_1d(ctx0, inp->s_copy, n_rs - n_seqs, n_seqs * inp->s_copy->nb[0]);
+ inp->s_copy_main = ggml_view_1d(ctx0, inp->s_copy, n_seqs, 0);
+ inp->s_copy_tail = ggml_view_1d(ctx0, inp->s_copy, n_rs - 1, inp->s_copy->nb[0]);
inp->head = mctx_cur->get_head();
inp->rs_z = mctx_cur->get_rs_z();
@@ -3592,7 +3595,11 @@ ggml_tensor * llm_graph_context::build_rs(
const llm_graph_get_rows_fn & get_state_rows) const {
const auto * kv_state = inp->mctx;
- return build_rs(s, inp->s_copy_main, inp->s_copy_extra, state_size, n_seqs,
+ // a custom getter reads the states of the ubatch straight from the cache, so the gather skips the first
+ // state: it still holds the n_rs - n_seqs extra states and copies no state of a single sequence ubatch
+ ggml_tensor * state_copy = get_state_rows ? inp->s_copy_tail : inp->s_copy;
+
+ return build_rs(s, state_copy, inp->s_copy_main, 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);
}
diff --git a/src/llama-graph.h b/src/llama-graph.h
index 3daa425bc..366eccbcd 100644
--- a/src/llama-graph.h
+++ b/src/llama-graph.h
@@ -273,8 +273,8 @@ public:
// views of s_copy, computed once per graph
// and shared across layers which use build_rs
- ggml_tensor * s_copy_main; // I32 [n_seqs]
- ggml_tensor * s_copy_extra; // I32 [n_rs - n_seqs]
+ ggml_tensor * s_copy_main; // I32 [n_seqs]
+ ggml_tensor * s_copy_tail; // I32 [n_rs - 1]
const llama_memory_recurrent_context * mctx;
@@ -1327,15 +1327,15 @@ struct llm_graph_context {
// `llama_memory_recurrent`
ggml_tensor * build_rs(
ggml_tensor * s,
+ ggml_tensor * state_copy,
ggml_tensor * state_copy_main,
- ggml_tensor * state_copy_extra,
int32_t state_size,
int32_t n_seqs,
uint32_t n_rs,
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 = nullptr) const;
llm_graph_input_rs * build_rs_inp() const;
@@ -1344,7 +1344,7 @@ struct llm_graph_context {
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 = nullptr) const;
ggml_tensor * build_rwkv_token_shift_load(
llm_graph_input_rs * inp,