Commit c173a53bd for llama.cpp
commit c173a53bdfca1047c710018dc934a6d67a8b010f
Author: Georgi Gerganov <ggerganov@gmail.com>
Date: Mon Oct 5 11:36:25 2026 +0300
llama : fix unexpected graph reallocation in the k-pool models (#29958)
* llama : fix unexpected graph reallocation in the k-pool models
Both k-pool models built a graph shape that depends on state the
full-context reserve cannot know:
- qwen4exp branched on inp->cache_safe, which turns false as soon as
llama_memory_seq_cp shares cells (e.g. batched-bench -pps): the QSA
layers swapped scatter+gather for fill+concat and dropped the
new_pool_rep leaf, so the decode graph had 12 fewer nodes than the
reserved one
- glm5-next branched on gather = n_tokens <= 16 && n_kv > n_sel, so the
TG decode built the gather shape (7564 nodes) while the last reserve,
the PP one, had the dense shape (7762 nodes)
Either mismatch forces a decode-time re-reserve that drops the
worst-case sizing and bakes in the current state, so the next state
growth (n_pool, n_kv, n_new) needs more room at an unchanged graph size
and aborts under GGML_SCHED_DEBUG_REALLOC=1. Reproduce with, e.g.:
GGML_SCHED_DEBUG_REALLOC=1 ./bin/llama-batched-bench \
-hf ggml-org/GLM-5.3-Flash-GGUF:Q2_K -npp 2500 -ntg 32 -npl 1,2 \
-c 32768 -pps -kvu
Always scatter+gather the pooled keys, and pick gather from context
constants only: n_ubatch bounds every ubatch, top_k + kpool - 1 bounds
n_sel. Every graph of a context then shares one shape, which the
reserve covers, and the dense path measured faster than the gather path
at 2.5k and 16k context.
Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-MOPD
* llama : drop the unused k-pool cache_safe graph API
The k-pool graphs no longer branch on cache_safe, so nothing reads
get_kpool_cache_safe() or the conditional new_pool_rep any more: both
models always pass the scatter target, which set_input_kpool now
requires instead of merely preferring.
Also drop the cache_safe copy in kpool_build_sizes(), a sizes-only
helper. The layout and state flag itself stays, it still decides which
pools a layout with shared cells must re-pool.
Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-MOPD
* tests : add a shared-seq graph reserve regression test
Decode a prompt into seq 0, share its cells with seq 1 via
llama_memory_seq_cp (what llama-batched-bench does for -pps), then keep
decoding both sequences. For the k-pool models sharing clears
cache_safe, which changes the graph topology while the pools keep
growing, so a scheduler that re-reserves with the current state
instead of the worst-case one aborts under GGML_SCHED_DEBUG_REALLOC=1.
The test registration sets that flag, and the test aborts on both
k-pool models before 2220411ec1.
kimi-linear and minimax-01 are skipped: they reserve the final pp graph
with n_seqs = 1 (see [TAG_RESERVE_DIAG_DECAY] in llama-context.cpp), so
every multi-seq graph has a different layout and re-reserves by design.
Assisted-by: pi:llama.cpp/MiMo-V2.6-Flash-MOPD
* cont : add TODOs
* cont : fix comment
* cuda: match the moe weighted reduction on empty ubatches
ggml_cuda_match_moe_weighted_reduction rejected tensors with zero
rows. A ubatch without outputs shrinks the last layer to zero rows
through inp_out_ids, so graph_optimize dropped its alloc dep there and
the scheduler graph lost one node compared to the reserved one. The
scheduler then re-reserved at the size of that ubatch, and the next
ubatch with the same node count but larger tensors aborted under
GGML_SCHED_DEBUG_REALLOC=1.
The compute loop already skips empty nodes before trying any fusion,
so the guard only made the alloc deps depend on the row count.
* tests: build the rollback test only where internal symbols link
The shared-seq case calls llm_arch_from_string, which libllama does
not export through LLAMA_API, so linking test-recurrent-state-rollback
fails on Windows with shared libraries. Its build now sits in the
NOT WIN32 OR NOT BUILD_SHARED_LIBS block, next to test-llama-archs and
the test registration it already lives under.
* tests: skip archs by name in the shared-seq reserve test
The skip of kimi-linear and minimax-01 went through llm_arch_from_string,
which libllama does not export through LLAMA_API, so the test could not
link on Windows with shared libraries. It now compares the
general.architecture string directly, and the test builds on every
platform again.
---------
Co-authored-by: Pascal <admin@serveurperso.com>
diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu
index 303d67e4c..c72419613 100644
--- a/ggml/src/ggml-cuda/ggml-cuda.cu
+++ b/ggml/src/ggml-cuda/ggml-cuda.cu
@@ -3164,7 +3164,7 @@ static bool ggml_cuda_match_moe_weighted_reduction(
const int n_expert_used = (int) weighted->ne[1];
const int64_t n_tokens = weighted->ne[2] * weighted->ne[3];
- if (n_expert_used < 2 || n_expert_used > MOE_WEIGHTED_REDUCTION_MAX_EXPERTS || n_tokens <= 0) {
+ if (n_expert_used < 2 || n_expert_used > MOE_WEIGHTED_REDUCTION_MAX_EXPERTS) {
return false;
}
diff --git a/src/llama-memory-hybrid-idx.cpp b/src/llama-memory-hybrid-idx.cpp
index 0bfba3ce0..65ce4f6fb 100644
--- a/src/llama-memory-hybrid-idx.cpp
+++ b/src/llama-memory-hybrid-idx.cpp
@@ -664,7 +664,6 @@ llama_memory_hybrid_idx_context::kpool_state llama_memory_hybrid_idx_context::kp
kpool_state st;
st.n_pool_real = lay.n_pool_real;
- st.cache_safe = lay.cache_safe;
return st;
}
@@ -784,10 +783,6 @@ uint32_t llama_memory_hybrid_idx_context::get_n_kpool_new() const {
return kpool_cur().n_new_g;
}
-bool llama_memory_hybrid_idx_context::get_kpool_cache_safe() const {
- return kpool_cur().cache_safe;
-}
-
void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells, ggml_tensor * pool_idxs, ggml_tensor * pool_mask, ggml_tensor * tail_idxs,
ggml_tensor * gather_mask, bool gather, ggml_tensor * new_pool_idxs, ggml_tensor * new_pool_rep,
const llama_ubatch * ubatch, ggml_tensor * new_pool_pos) const {
@@ -816,13 +811,11 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
GGML_ASSERT(pool_mask->ne[0] == (int64_t) n_pool && pool_mask->ne[1] == (int64_t) n_tokens);
GGML_ASSERT(tail_idxs->ne[0] == (int64_t) kpool - 1 && tail_idxs->ne[1] == (int64_t) n_tokens);
GGML_ASSERT(pool_idxs->ne[0] == (int64_t) kpool && pool_idxs->ne[1] == (int64_t) n_pool);
- GGML_ASSERT(st.cache_safe == (new_pool_rep != nullptr));
GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_idxs->buffer));
GGML_ASSERT(new_pool_idxs->ne[0] == (int64_t) kpool && new_pool_idxs->ne[1] == (int64_t) n_new_g);
- if (new_pool_rep != nullptr) {
- GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_rep->buffer));
- GGML_ASSERT(new_pool_rep->ne[0] == (int64_t) n_new_g);
- }
+ // the graph always scatters the fresh pooled keys back into the cache, see build_qsa_sel
+ GGML_ASSERT(new_pool_rep != nullptr && ggml_backend_buffer_is_host(new_pool_rep->buffer));
+ GGML_ASSERT(new_pool_rep->ne[0] == (int64_t) n_new_g);
if (new_pool_pos != nullptr) {
GGML_ASSERT(ggml_backend_buffer_is_host(new_pool_pos->buffer));
GGML_ASSERT(new_pool_pos->ne[0] == 4*(int64_t) n_new_g);
@@ -891,7 +884,7 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
int32_t * pcell = (int32_t *) pool_cells->data;
int32_t * pidx = (int32_t *) pool_idxs->data;
int32_t * nidx = (int32_t *) new_pool_idxs->data;
- int64_t * nrep = new_pool_rep != nullptr ? (int64_t *) new_pool_rep->data : nullptr;
+ int64_t * nrep = (int64_t *) new_pool_rep->data;
int32_t * npos = new_pool_pos != nullptr ? (int32_t *) new_pool_pos->data : nullptr;
if (npos != nullptr) {
@@ -924,9 +917,7 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
for (uint32_t k = 0; k < kpool; ++k) {
nidx[(size_t) i_new*kpool + k] = (int32_t) gcell(sq, sq.cells[j + k].second);
}
- if (nrep != nullptr) {
- nrep[i_new] = gcell(sq, rep);
- }
+ nrep[i_new] = gcell(sq, rep);
if (npos != nullptr) {
// a pooled key is rotated to the M-RoPE position of its first member
const uint32_t c = sq.cells[j].second;
@@ -959,9 +950,7 @@ void llama_memory_hybrid_idx_context::set_input_kpool(ggml_tensor * pool_cells,
for (uint32_t k = 0; k < kpool; ++k) {
nidx[(size_t) i*kpool + k] = (int32_t) pad_cell;
}
- if (nrep != nullptr) {
- nrep[i] = pad_cell;
- }
+ nrep[i] = pad_cell;
}
}
diff --git a/src/llama-memory-hybrid-idx.h b/src/llama-memory-hybrid-idx.h
index 1e63a4099..18e3413b1 100644
--- a/src/llama-memory-hybrid-idx.h
+++ b/src/llama-memory-hybrid-idx.h
@@ -194,7 +194,6 @@ public:
// glm5-next and qwen4exp, complete pools of kpool cells per sequence, scored as whole pools.
uint32_t get_n_kpool () const; // Padded pool count, where the last pool is always unused.
uint32_t get_n_kpool_new() const; // Pools to re-pool this ubatch, padded to a stable bound, never below 1.
- bool get_kpool_cache_safe() const;
kpool_access get_kpool_access(ggml_context * ctx, int32_t il, int64_t n_embd) const;
ggml_tensor * gather_mla_rows(ggml_context * ctx, ggml_tensor * idxs, int64_t n_rows, int64_t n_embd, int32_t il) const;
// new_pool_pos (I32 [4*n_new]): M-RoPE position of each new pool's first member, for pooled keys rotated at pooling time
diff --git a/src/models/glm5-next.cpp b/src/models/glm5-next.cpp
index 09382fdac..839f223d2 100644
--- a/src/models/glm5-next.cpp
+++ b/src/models/glm5-next.cpp
@@ -263,7 +263,6 @@ public:
// The scatter mask shape follows n_kv.
res &= n_kv == idx->get_n_kv();
res &= n_new == std::max(mctx->get_n_kpool_new(), 1u);
- res &= cache_safe == mctx->get_kpool_cache_safe();
return res;
}
@@ -282,7 +281,6 @@ public:
const uint32_t kpool;
uint32_t n_new = 0;
uint32_t n_sel = 0;
- bool cache_safe = true;
bool gather = false;
uint32_t n_kv = 0;
};
@@ -296,8 +294,7 @@ llama_model_glm5_next::llm_graph_input_kpool * llama_model_glm5_next::graph::bui
const uint32_t n_kv = mctx_idx->get_n_kv();
// a ubatch that completes no pool still builds one dummy entry, so the graph does not
// change shape every kpool tokens
- const uint32_t n_new = std::max(mctx_hyb->get_n_kpool_new(), 1u);
- const bool cache_safe = mctx_hyb->get_kpool_cache_safe();
+ const uint32_t n_new = std::max(mctx_hyb->get_n_kpool_new(), 1u);
// the fused lightning indexer wants an f16 mask
const auto type_mask = cparams.fused_lid ? GGML_TYPE_F16 : GGML_TYPE_F32;
@@ -321,14 +318,17 @@ llama_model_glm5_next::llm_graph_input_kpool * llama_model_glm5_next::graph::bui
inp->n_kv = n_kv;
- // Gather selected latents for small decode batches when n_kv exceeds n_sel.
+ // Gather selected latents for small batches when the context exceeds the selection width.
{
constexpr int64_t max_ub = 16;
const int64_t n_top_pool = std::min<int64_t>(n_pool, hparams.indexer_top_k / kpool);
const int64_t n_sel = kpool*n_top_pool + (hparams.indexer_kpool_select_tail ? kpool - 1 : 0);
inp->n_sel = (uint32_t) n_sel;
- inp->gather = (int64_t) n_tokens <= max_ub && (int64_t) n_kv > n_sel;
+ // both terms are context constants: n_ubatch bounds every ubatch and top_k + kpool - 1 bounds
+ // n_sel, so the graph shape follows neither n_tokens nor n_kv, which the reserve cannot predict
+ // TODO: remove "gather" logic and everything related. the backends now support sparse attension so this is obsolete
+ inp->gather = (int64_t) cparams.n_ubatch <= max_ub && (int64_t) cparams.n_ctx > hparams.indexer_top_k + kpool - 1;
// Both paths read the slot mask: gather adds it to the scores, scatter maps its dead slots to dump rows.
inp->gather_mask = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_sel, 1, 1, n_tokens);
@@ -338,14 +338,13 @@ llama_model_glm5_next::llm_graph_input_kpool * llama_model_glm5_next::graph::bui
}
inp->n_new = n_new;
- inp->cache_safe = cache_safe;
inp->new_pool_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, n_new);
ggml_set_input(inp->new_pool_idxs);
- if (cache_safe) {
- inp->new_pool_rep = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, n_new);
- ggml_set_input(inp->new_pool_rep);
- }
+ // the scatter target is part of the graph shape: llama_context reserves the full-context graph,
+ // so this must not depend on cache_safe, which only the decode-time graph can know
+ inp->new_pool_rep = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, n_new);
+ ggml_set_input(inp->new_pool_rep);
return (llm_graph_input_kpool *) res->add_input(std::move(inp));
}
@@ -814,20 +813,11 @@ ggml_tensor * llama_model_glm5_next::graph::build_kpool_select(
pooled_new = ggml_reshape_2d(ctx0, pooled_new, n_embd_indexer, n_new);
cb(pooled_new, "indexer_pool_k_new", il);
- if (inp_kpool->cache_safe) {
- // Write before the pool gather.
- ggml_build_forward_expand(gf, kpool_cache.scatter_pooled(pooled_new, inp_kpool->new_pool_rep));
- }
-
- ggml_tensor * pooled = nullptr;
- if (inp_kpool->cache_safe) {
- pooled = kpool_cache.gather_pooled(inp_kpool->pool_cells);
- } else {
- GGML_ASSERT(n_new <= n_pool);
- ggml_tensor * pad = ggml_fill(ctx0,
- ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_indexer, n_pool - n_new), 0.0f);
- pooled = ggml_concat(ctx0, pooled_new, pad, 1);
- }
+ // scatter the fresh pooled keys, then gather all n_pool of them by cell, in both cache modes:
+ // the reserved graph cannot branch on cache_safe, and without sharing every pool is re-pooled
+ // anyway (n_new == n_pool_real, layout order), so the gather returns exactly pooled_new
+ ggml_build_forward_expand(gf, kpool_cache.scatter_pooled(pooled_new, inp_kpool->new_pool_rep));
+ ggml_tensor * pooled = kpool_cache.gather_pooled(inp_kpool->pool_cells);
pooled = ggml_reshape_3d(ctx0, pooled, n_embd_indexer, 1, n_pool);
cb(pooled, "indexer_pool_k", il);
diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp
index 416c9263b..250eb74da 100644
--- a/src/models/qwen4exp.cpp
+++ b/src/models/qwen4exp.cpp
@@ -674,7 +674,6 @@ public:
// the scatter mask shape follows n_kv
res &= n_kv == idx->get_n_kv();
res &= n_new == mctx->get_n_kpool_new();
- res &= cache_safe == mctx->get_kpool_cache_safe();
return res;
}
@@ -693,7 +692,6 @@ public:
uint32_t n_new = 0; // padded to a stable bound, never below 1
uint32_t n_sel = 0;
uint32_t n_kv = 0;
- bool cache_safe = true;
};
llama_model_qwen4exp::llm_graph_input_kpool * llama_model_qwen4exp::graph::build_inp_kpool(const llama_memory_hybrid_idx_context * mctx_hyb) {
@@ -721,18 +719,17 @@ llama_model_qwen4exp::llm_graph_input_kpool * llama_model_qwen4exp::graph::build
ggml_build_forward_expand(gf, inp->pool_mask);
ggml_build_forward_expand(gf, inp->tail_idxs);
- inp->n_kv = mctx_idx->get_n_kv();
- inp->n_new = mctx_hyb->get_n_kpool_new();
- inp->cache_safe = mctx_hyb->get_kpool_cache_safe();
+ inp->n_kv = mctx_idx->get_n_kv();
+ inp->n_new = mctx_hyb->get_n_kpool_new();
// the top blocks plus the tail
- inp->n_sel = kpool*std::min<uint32_t>(n_pool, hparams.indexer_top_k / kpool) + kpool - 1;
+ inp->n_sel = kpool*std::min<uint32_t>(n_pool, hparams.indexer_top_k / kpool) + kpool - 1;
inp->new_pool_idxs = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, inp->n_new);
ggml_set_input(inp->new_pool_idxs);
- if (inp->cache_safe) {
- inp->new_pool_rep = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, inp->n_new);
- ggml_set_input(inp->new_pool_rep);
- }
+ // the scatter target is part of the graph shape: llama_context reserves the full-context graph,
+ // so this must not depend on cache_safe, which only the decode-time graph can know
+ inp->new_pool_rep = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, inp->n_new);
+ ggml_set_input(inp->new_pool_rep);
inp->new_pool_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 4*inp->n_new);
ggml_set_input(inp->new_pool_pos);
@@ -790,17 +787,11 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
pooled_new = ggml_reshape_2d(ctx0, pooled_new, idx_dim, n_new);
cb(pooled_new, "indexer_pool_k_new", il);
- ggml_tensor * pooled = nullptr;
- if (inp_kpool->cache_safe) {
- // write before the pool gather
- ggml_build_forward_expand(gf, kpool_cache.scatter_pooled(pooled_new, inp_kpool->new_pool_rep));
- pooled = kpool_cache.gather_pooled(inp_kpool->pool_cells);
- } else {
- // shared cells re-pool every pool, in layout order
- GGML_ASSERT(n_new < n_pool);
- ggml_tensor * pad = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, idx_dim, n_pool - n_new), 0.0f);
- pooled = ggml_concat(ctx0, pooled_new, pad, 1);
- }
+ // scatter the fresh pooled keys, then gather all n_pool of them by cell, in both cache modes:
+ // the reserved graph cannot branch on cache_safe, and without sharing every pool is re-pooled
+ // anyway (n_new == n_pool_real, layout order), so the gather returns exactly pooled_new
+ ggml_build_forward_expand(gf, kpool_cache.scatter_pooled(pooled_new, inp_kpool->new_pool_rep));
+ ggml_tensor * pooled = kpool_cache.gather_pooled(inp_kpool->pool_cells);
pooled = ggml_reshape_3d(ctx0, pooled, idx_dim, 1, n_pool);
cb(pooled, "indexer_k", il);
diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt
index b5ef02c02..c44189ceb 100644
--- a/tests/CMakeLists.txt
+++ b/tests/CMakeLists.txt
@@ -213,7 +213,11 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS)
LABEL main
ARGS --models "${MODEL_DIR}"
)
- set_tests_properties(test-recurrent-state-rollback PROPERTIES FIXTURES_REQUIRED generate-models)
+ # abort on graph reallocations that the initial reserve should have covered
+ set_tests_properties(test-recurrent-state-rollback PROPERTIES
+ FIXTURES_REQUIRED generate-models
+ ENVIRONMENT "GGML_SCHED_DEBUG_REALLOC=1"
+ )
# Test state save/load functionality across all architectures, using the generated dummy models
llama_test(
diff --git a/tests/test-recurrent-state-rollback.cpp b/tests/test-recurrent-state-rollback.cpp
index cac432e3e..4043a229c 100644
--- a/tests/test-recurrent-state-rollback.cpp
+++ b/tests/test-recurrent-state-rollback.cpp
@@ -1,3 +1,6 @@
+// TODO: merge with test-save-load-state.cpp
+// TODO: merge with test-state-restore-fragmented.cpp
+
#include "arg.h"
#include "common.h"
#include "ggml-backend.h"
@@ -474,6 +477,87 @@ static test_status test_rollback(const common_params & params, llama_model * mod
return test_status::PASS;
}
+// Decode a prompt into seq 0, share its cells with a second sequence, then keep
+// decoding both.
+static test_status test_shared_seq_reserve(const common_params & params, llama_model * model, uint8_t fill) {
+ const int n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model));
+
+ // these archs reserve the final pp graph with n_seqs = 1, so every multi-seq
+ // graph has a different layout and re-reserves by design
+ // see [TAG_RESERVE_DIAG_DECAY] in llama-context.cpp
+ char arch_str[64] = {};
+ llama_model_meta_val_str(model, "general.architecture", arch_str, sizeof(arch_str));
+ if (strcmp(arch_str, "kimi-linear") == 0 || strcmp(arch_str, "minimax-01") == 0) {
+ LOG_INF("%s: skipping %s, its reserve uses n_seqs = 1\n", __func__, arch_str);
+ return test_status::SKIP;
+ }
+
+ constexpr uint32_t n_seqs = 2;
+ constexpr uint32_t n_prompt = 128;
+ constexpr uint32_t n_continue = 32;
+
+ auto cparams = common_context_params_to_llama(params);
+ cparams.n_seq_max = n_seqs;
+ cparams.n_ctx = 512;
+ cparams.n_batch = 256;
+ cparams.n_ubatch = 64;
+ cparams.kv_unified = true; // only a unified cache shares cells on seq_cp
+
+ llama_context_ptr ctx = init_ctx(model, cparams, fill);
+ if (!ctx) {
+ LOG_ERR("%s: failed to init context\n", __func__);
+ return test_status::FAIL;
+ }
+
+ const auto tok = [&](uint32_t seq, llama_pos pos) {
+ return (llama_token) ((7*(uint32_t) pos + 31*seq + 1) % (uint32_t) n_vocab);
+ };
+
+ {
+ common_batch batch(ctx.get());
+ for (llama_pos pos = 0; pos < (llama_pos) n_prompt; ++pos) {
+ batch.add(tok(0, pos), pos, 0, false);
+ }
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
+ LOG_ERR("%s: prompt decode failed\n", __func__);
+ return test_status::FAIL;
+ }
+ }
+
+ // this is what llama-batched-bench does for -pps
+ llama_memory_seq_cp(llama_get_memory(ctx.get()), 0, 1, -1, -1);
+
+ for (uint32_t i = 0; i < n_continue; ++i) {
+ const llama_pos pos = (llama_pos) (n_prompt + i);
+
+ common_batch batch(ctx.get());
+ for (uint32_t s = 0; s < n_seqs; ++s) {
+ batch.add(tok(s, pos), pos, (llama_seq_id) s, true);
+ }
+ if (llama_process(ctx.get(), LLAMA_PROCESS_TYPE_DECODE, batch.get()) != 0) {
+ LOG_ERR("%s: shared-seq decode failed at step %u\n", __func__, i);
+ return test_status::FAIL;
+ }
+
+ for (uint32_t s = 0; s < n_seqs; ++s) {
+ const float * logits = llama_get_logits_ith(ctx.get(), (int) s);
+ if (logits == nullptr) {
+ LOG_ERR("%s: missing shared-seq logits at index %u\n", __func__, s);
+ return test_status::FAIL;
+ }
+ for (int t = 0; t < n_vocab; ++t) {
+ if (!std::isfinite(logits[t])) {
+ LOG_ERR("%s: non-finite shared-seq logit at step %u, seq %u, index %d\n", __func__, i, s, t);
+ return test_status::FAIL;
+ }
+ }
+ }
+ }
+
+ LOG_INF("%s: shared-seq decode succeeded (%u tokens after seq_cp)\n", __func__, n_continue*n_seqs);
+ return test_status::PASS;
+}
+
static test_status merge_status(test_status a, test_status b) {
if (a == test_status::FAIL || b == test_status::FAIL) {
return test_status::FAIL;
@@ -487,6 +571,7 @@ static test_status merge_status(test_status a, test_status b) {
struct test_results {
test_status rollback = test_status::SKIP;
test_status replay = test_status::SKIP;
+ test_status shared = test_status::SKIP;
};
// Run every test for an initialized model over both cache fills.
@@ -496,9 +581,11 @@ static test_results run_tests(const common_params & params, llama_model * model)
LOG_INF("%s: testing with cache fill 0x%02x\n", __func__, fill);
const test_status rb = test_rollback(params, model, fill);
const test_status rp = test_multi_seq_split_replay(params, model, fill);
+ const test_status ss = test_shared_seq_reserve(params, model, fill);
res.rollback = merge_status(res.rollback, rb);
res.replay = merge_status(res.replay, rp);
- if (rb == test_status::FAIL || rp == test_status::FAIL) {
+ res.shared = merge_status(res.shared, ss);
+ if (rb == test_status::FAIL || rp == test_status::FAIL || ss == test_status::FAIL) {
break;
}
}
@@ -517,7 +604,7 @@ static test_results run_tests_for_model(const std::string & model_path, const st
if (model == nullptr) {
LOG_ERR("%s: failed to init model '%s'\n", __func__, model_path.c_str());
// a model that cannot be loaded is a failure, not a skip
- return { test_status::FAIL, test_status::FAIL };
+ return { test_status::FAIL, test_status::FAIL, test_status::FAIL };
}
if (!llama_model_is_recurrent(model) && !llama_model_is_hybrid(model)) {
@@ -603,12 +690,12 @@ int main(int argc, char ** argv) {
// silence everything but the table itself (LOG has verbosity LOG_LEVEL_OUTPUT = 0)
common_log_set_verbosity_thold(0);
- LOG("%-*s %-8s %s\n", (int) name_width, "Model", "rollback", "split replay");
+ LOG("%-*s %-8s %-11s %s\n", (int) name_width, "Model", "rollback", "split replay", "shared seq");
common_log_flush(common_log_main());
- size_t n_pass[2] = { 0, 0 };
- size_t n_skip[2] = { 0, 0 };
- size_t n_fail[2] = { 0, 0 };
+ size_t n_pass[3] = { 0, 0, 0 };
+ size_t n_skip[3] = { 0, 0, 0 };
+ size_t n_fail[3] = { 0, 0, 0 };
for (const auto & model_path : models) {
const auto name = std::filesystem::path(model_path).filename().string();
@@ -618,12 +705,12 @@ int main(int argc, char ** argv) {
// all status strings have the same raw length, so the columns line up;
// pad the first status to the width of the "rollback" header + separator
- LOG(" %s %s", test_status_str(res.rollback), test_status_str(res.replay));
+ LOG(" %s %s %s", test_status_str(res.rollback), test_status_str(res.replay), test_status_str(res.shared));
LOG("\n");
common_log_flush(common_log_main());
- const test_status all[2] = { res.rollback, res.replay };
- for (int t = 0; t < 2; ++t) {
+ const test_status all[3] = { res.rollback, res.replay, res.shared };
+ for (int t = 0; t < 3; ++t) {
switch (all[t]) {
case test_status::PASS: n_pass[t]++; break;
case test_status::FAIL: n_fail[t]++; break;
@@ -639,12 +726,14 @@ int main(int argc, char ** argv) {
__func__, n_pass[0], n_skip[0], n_fail[0], models.size());
LOG_INF("%s: split replay: %zu passed, %zu skipped, %zu failed (of %zu)\n",
__func__, n_pass[1], n_skip[1], n_fail[1], models.size());
+ LOG_INF("%s: shared seq: %zu passed, %zu skipped, %zu failed (of %zu)\n",
+ __func__, n_pass[2], n_skip[2], n_fail[2], models.size());
- return (n_fail[0] + n_fail[1]) == 0 ? 0 : 1;
+ return (n_fail[0] + n_fail[1] + n_fail[2]) == 0 ? 0 : 1;
}
// single-model mode
const test_results res = run_tests_for_model(params.model.path, params);
- return (res.rollback == test_status::FAIL || res.replay == test_status::FAIL) ? 1 : 0;
+ return (res.rollback == test_status::FAIL || res.replay == test_status::FAIL || res.shared == test_status::FAIL) ? 1 : 0;
}
diff --git a/tests/test-save-load-state.cpp b/tests/test-save-load-state.cpp
index b9d2fd8e6..43a513ae5 100644
--- a/tests/test-save-load-state.cpp
+++ b/tests/test-save-load-state.cpp
@@ -1,3 +1,6 @@
+// TODO: merge with test-recurrent-state-rollback.cpp
+// TODO: merge with test-state-restore-fragmented.cpp
+
#include "arg.h"
#include "common.h"
#include "log.h"
diff --git a/tests/test-state-restore-fragmented.cpp b/tests/test-state-restore-fragmented.cpp
index ea3006949..d20c231d7 100644
--- a/tests/test-state-restore-fragmented.cpp
+++ b/tests/test-state-restore-fragmented.cpp
@@ -6,6 +6,9 @@
// The fix changes find_slot(ubatch, true) to find_slot(ubatch, false)
// in state_read_meta(), allowing non-contiguous slot allocation.
+// TODO: merge with test-save-load-state.cpp
+// TODO: merge with test-recurrent-state-rollback.cpp
+
#include "arg.h"
#include "common.h"
#include "llama.h"