Commit 81e39ad34 for llama.cpp
commit 81e39ad34368329b5db77620721510fd91107e80
Author: Georgi Gerganov <ggerganov@gmail.com>
Date: Thu Oct 1 19:55:34 2026 +0300
llama : clamp kpool re-pool bound to existing pools (#29805)
* tests : simplify function signature
* llama : clamp kpool re-pool bound to existing pools
The n_tokens/kpool + n_seqs_unq bound on n_new_g overshoots when a batch
fills the whole cache: n_ctx tokens complete exactly n_ctx/kpool pools, so
the +1 pads new_pool_idxs/new_pool_rep one entry past n_pool_real. Graph
reserve only covers n_pool_real entries, so the first full-context decode
builds bigger tensors than reserved and ggml-alloc demands a graph
reallocation (abort under GGML_SCHED_DEBUG_REALLOC=1).
Clamp the bound to n_pool_real: a ubatch can never mark more pools than
the cache holds, and reserve's n_pool_max already covers that.
Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-MOPD
* cont : cap to n_pool_max
diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp
index de64a3700..3b43df92d 100644
--- a/src/llama-memory-hybrid-idx.cpp
+++ b/src/llama-memory-hybrid-idx.cpp
@@ -762,10 +762,12 @@ void llama_memory_hybrid_idx_context::kpool_build_state(const llama_ubatch & uba
}
}
- // a ubatch touches at most t_s/kpool + 1 pools of a sequence with t_s tokens: pad to that bound so the
- // graph keeps its shape as the count moves, e.g. between 0 and n_seq while several sequences decode
+ // a ubatch touches at most t_s/kpool + 1 pools per sequence, pad to that bound so the graph keeps its shape
+ // as the count moves; reserve sizes the list for every pool the cache can hold, so never pad past n_pool_max
+ const auto * idx = mem->get_mem_idx();
+ const uint32_t n_pool_max = idx->get_size() / kpool * idx->get_n_seq_max();
const uint32_t bound = ubatch.n_tokens/kpool + ubatch.n_seqs_unq;
- st.n_new_g = std::max({st.n_new, 1u, std::min(bound, kpool_pad(st.n_pool_real) - 1)});
+ st.n_new_g = std::max({st.n_new, 1u, std::min({bound, kpool_pad(st.n_pool_real) - 1, n_pool_max})});
}
const llama_memory_hybrid_idx_context::kpool_state & llama_memory_hybrid_idx_context::kpool_cur() const {
diff --git a/tests/test-recurrent-state-rollback.cpp b/tests/test-recurrent-state-rollback.cpp
index f8eda55c8..cac432e3e 100644
--- a/tests/test-recurrent-state-rollback.cpp
+++ b/tests/test-recurrent-state-rollback.cpp
@@ -34,10 +34,10 @@ static const char * test_status_str(test_status status) {
return "";
}
-static bool decode_tokens(llama_context * ctx, const std::vector<llama_token> & tokens, uint32_t count) {
+static bool decode_tokens(llama_context * ctx, const std::vector<llama_token> & tokens) {
common_batch batch(ctx);
- for (uint32_t pos = 0; pos < count; ++pos) {
- batch.add(tokens[pos], pos, 0, pos + 1 == count);
+ for (uint32_t pos = 0; pos < tokens.size(); ++pos) {
+ batch.add(tokens[pos], pos, 0, pos + 1 == tokens.size());
}
return llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get()) == 0;
}
@@ -74,7 +74,7 @@ static llama_context_ptr init_ctx(llama_model * model, llama_context_params cpar
// Use a full ubatch so buffer discovery preserves prefill allocation sizes.
const uint32_t n_tokens = llama_n_ubatch(ctx.get());
- if (!decode_tokens(ctx.get(), std::vector<llama_token>(n_tokens, 0), n_tokens)) {
+ if (!decode_tokens(ctx.get(), std::vector<llama_token>(n_tokens, 0))) {
return nullptr;
}
llama_synchronize(ctx.get());
@@ -347,7 +347,7 @@ static test_status test_rollback(const common_params & params, llama_model * mod
// Decode the full prompt on the source, then roll back three positions.
// Replaying them crosses DSV4's ratio-4 compressor boundary.
// Rollback leaves the recurrent memory in a snapshot state (rs_idx != 0).
- if (!decode_tokens(ctx_src.get(), tokens, n_tokens)) {
+ if (!decode_tokens(ctx_src.get(), tokens)) {
LOG_ERR("%s: failed to decode prompt\n", __func__);
return test_status::FAIL;
}
@@ -432,7 +432,7 @@ static test_status test_rollback(const common_params & params, llama_model * mod
t = 0;
}
}
- if (!decode_tokens(ctx_dirty.get(), noise, n_tokens)) {
+ if (!decode_tokens(ctx_dirty.get(), noise)) {
LOG_ERR("%s: dirty prompt decode failed\n", __func__);
return test_status::FAIL;
}