diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index b79a5270b52a..61a7f48c1e19 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -2307,11 +2307,14 @@ struct server_context_impl { llama_pos pos_next = slot.prompt.tokens.pos_next(n_past); + const bool is_recurrent = llama_model_is_recurrent(model); + // note: when n_swa == 0, the model does not use SWA const auto n_swa = std::max(0, llama_model_n_swa(model)); - // the largest pos_min required for a checkpoint to be useful - const auto pos_min_thold = std::max(0, pos_next - n_swa); + // For hybrid/recurrent: SWA threshold not meaningful, set to 0. + // For pure transformer/SWA: preserve existing behavior. + const auto pos_min_thold = is_recurrent ? (llama_pos) 0 : (llama_pos) std::max(0, pos_next - n_swa); if (n_past > 0 && n_past < slot.prompt.n_tokens()) { const auto pos_min = llama_memory_seq_pos_min(llama_get_memory(ctx), slot.id); @@ -2371,6 +2374,10 @@ struct server_context_impl { slot.prompt.checkpoints.rbegin(), slot.prompt.checkpoints.rend(), [&, func_name = __func__](const auto & cur) { + if (is_recurrent) { + // For hybrid/recurrent: use position-matching semantics + return cur.pos_max <= n_past && cur.pos_max < pos_next; + } // guarantee that a checkpoint will result in at least one token being processed [TAG_PROMPT_LOGITS] LOG_INF("slot %12.*s: id %2d | task %d | Checking checkpoint with [%d, %d] against %d...\n", 12, func_name, (slot).id, ((slot).task ? (slot).task->id : -1), cur.pos_min, cur.pos_max, pos_min_thold); @@ -2397,10 +2404,15 @@ struct server_context_impl { } if (do_reset) { - SLT_WRN(slot, "forcing full prompt re-processing due to lack of cache data (likely due to SWA or hybrid/recurrent memory, see %s)\n", - "https://github.com/ggml-org/llama.cpp/pull/13194#issuecomment-2868343055"); - pos_next = 0; - n_past = 0; + if (is_recurrent) { + // For hybrid/recurrent: preserve current state, do not zero n_past + SLT_WRN(slot, "no matching recurrent checkpoint; preserving prompt state (n_past = %d)\n", n_past); + } else { + SLT_WRN(slot, "forcing full prompt re-processing due to lack of cache data (likely due to SWA or hybrid/recurrent memory, see %s)\n", + "https://github.com/ggml-org/llama.cpp/pull/13194#issuecomment-2868343055"); + pos_next = 0; + n_past = 0; + } } } } @@ -2410,7 +2422,7 @@ struct server_context_impl { for (auto it = slot.prompt.checkpoints.begin(); it != slot.prompt.checkpoints.end();) { const auto & cur = *it; if (cur.pos_max > pos_next) { - SLT_WRN(slot, "erased invalidated context checkpoint (pos_min = %d, pos_max = %d, n_tokens = %" PRId64 ", n_swa = %d, pos_next = %d, size = %.3f MiB)\n", cur.pos_min, cur.pos_max, cur.n_tokens, n_swa, pos_next, (float) cur.data.size() / 1024 / 1024); + SLT_WRN(slot, "erased invalidated context checkpoint (pos_min = %d, pos_max = %d, n_tokens = %" PRId64 ", pos_next = %d, size = %.3f MiB)\n", cur.pos_min, cur.pos_max, cur.n_tokens, pos_next, (float) cur.data.size() / 1024 / 1024); it = slot.prompt.checkpoints.erase(it); } else { ++it;