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"