From f1c2570995ee58fd82d4bca67a01249ef324bb5e Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Wed, 29 Apr 2026 17:41:23 +0300 Subject: [PATCH] spec : fix draft model checkpoints --- common/speculative.cpp | 63 +++++++++++++++++------------------------- 1 file changed, 26 insertions(+), 37 deletions(-) diff --git a/common/speculative.cpp b/common/speculative.cpp index bda9993b15..e538217cc2 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -167,8 +167,6 @@ struct common_speculative_checkpoint { size_t size() const { return data.size(); } - - size_t ckpt_size = 0; }; struct common_speculative_state_draft : public common_speculative_state { @@ -176,7 +174,7 @@ struct common_speculative_state_draft : public common_speculative_state { llama_context * ctx_dft; bool use_ckpt = false; - struct common_speculative_checkpoint ckpt; + common_speculative_checkpoint ckpt; common_sampler * smpl; @@ -249,26 +247,16 @@ struct common_speculative_state_draft : public common_speculative_state { llama_batch_free(batch); } - void begin(const llama_tokens & prompt) override { - if (use_ckpt && ckpt.size() > 0) { - // delete checkpoint - LOG_DBG("%s: delete checkpoint, prompt.size=%zu, pos_min=%d, pos_max=%d, n_tokens=%" PRId64 ", size=%.3f MiB\n", - __func__, prompt.size(), ckpt.pos_min, ckpt.pos_max, ckpt.n_tokens, (float) ckpt.data.size() / 1024 / 1024); - ckpt.pos_min = 0; - ckpt.pos_max = 0; - ckpt.n_tokens = 0; - ckpt.ckpt_size = 0; - ckpt.data.clear(); - } + void begin(const llama_tokens & /*prompt*/) override { } - size_t draft_create_checkpoint(int n_tokens_prompt, int n_tokens_batch) { + size_t create_checkpoint(int n_tokens_prompt) { int slot_id = 0; const size_t checkpoint_size = llama_state_seq_get_size_ext(ctx_dft, slot_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); ckpt.pos_min = llama_memory_seq_pos_min(llama_get_memory(ctx_dft), slot_id); ckpt.pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), slot_id); - ckpt.n_tokens = n_tokens_prompt - n_tokens_batch; + ckpt.n_tokens = n_tokens_prompt; ckpt.data.resize(checkpoint_size); const size_t n = llama_state_seq_get_data_ext(ctx_dft, ckpt.data.data(), checkpoint_size, slot_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); @@ -281,13 +269,13 @@ struct common_speculative_state_draft : public common_speculative_state { return n; } - size_t draft_restore_checkpoint(size_t ckpt_size_part_expected) { + size_t restore_checkpoint() { int slot_id = 0; LOG_DBG("%s: pos_min = %d, pos_max = %d\n", __func__, ckpt.pos_min, ckpt.pos_max); const size_t n = llama_state_seq_set_data_ext(ctx_dft, ckpt.data.data(), ckpt.size(), slot_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); - if (n != ckpt_size_part_expected) { - GGML_ABORT("%s: failed to restore context checkpoint (pos_min=%d, pos_max=%d, size=%zu, get_data_ext->%zu, set_data_ext->%zu", - __func__, ckpt.pos_min, ckpt.pos_max, ckpt.size(), ckpt_size_part_expected, n); + if (n != ckpt.size()) { + GGML_ABORT("%s: failed to restore context checkpoint (pos_min=%d, pos_max=%d, size=%zu", + __func__, ckpt.pos_min, ckpt.pos_max, ckpt.size()); } llama_memory_seq_rm(llama_get_memory(ctx_dft), slot_id, ckpt.pos_max + 1, -1); @@ -351,8 +339,8 @@ struct common_speculative_state_draft : public common_speculative_state { for (int i = 0; i < (int) prompt_dft.size(); ++i) { int cur = 0; while (i_start + cur < (int) prompt_cur.size() && - i + cur < (int) prompt_dft.size() && - prompt_cur[i_start + cur] == prompt_dft[i + cur]) { + i + cur < (int) prompt_dft.size() && + prompt_cur[i_start + cur] == prompt_dft[i + cur]) { cur++; } @@ -364,17 +352,18 @@ struct common_speculative_state_draft : public common_speculative_state { LOG_DBG("%s: reuse_i = %d, reuse_n = %d, #prompt_dft = %zu, #prompt_cur = %zu\n", __func__, reuse_i, reuse_n, prompt_dft.size(), prompt_cur.size()); - if (use_ckpt && ckpt.ckpt_size == 0 && reuse_n > 0) { - LOG_DBG("%s: no checkpoint available, no reuse, (reuse_i=%d, reuse_n=%d) -> (0, 0)\n", - __func__, reuse_i, reuse_n); + if (use_ckpt && ckpt.n_tokens > reuse_n) { + LOG_DBG("%s: checkpoint is outdated -> delete it (reuse_i=%d, reuse_n=%d) -> (0, 0), ckpt.n_tokens = %lld\n", + __func__, reuse_i, reuse_n, ckpt.n_tokens); reuse_i = 0; reuse_n = 0; + + ckpt = {}; } result.clear(); result.reserve(sparams.n_max); - bool needs_ckpt = use_ckpt && prompt_dft.size() > 0; if (reuse_n == 0 || (use_ckpt && reuse_i > 0)) { llama_memory_clear(mem_dft, false); prompt_dft.clear(); @@ -414,14 +403,13 @@ struct common_speculative_state_draft : public common_speculative_state { if (reuse_n < (int) prompt_dft.size() || do_restore) { if (use_ckpt) { - if (ckpt.n_tokens > (int64_t) prompt_dft.size()) { - LOG_INF("%s: checkpoint is too large, prompt_tgt.size=%zu, ckpt.n_tokens=%" PRId64 ", reuse_n=%d, prompt_dft.size=%zu\n", - __func__, prompt_tgt.size(), ckpt.n_tokens, reuse_n, prompt_dft.size()); + if (ckpt.n_tokens > 0) { + LOG_DBG("%s: restoring checkpoint, reuse_n=%d, prompt_dft.size=%zu\n", + __func__, reuse_n, prompt_dft.size()); + restore_checkpoint(); + reuse_n = ckpt.n_tokens; + prompt_dft.resize(reuse_n); } - draft_restore_checkpoint(ckpt.ckpt_size); - reuse_n = ckpt.n_tokens; - prompt_dft.resize(reuse_n); - needs_ckpt = false; } else { bool is_removed = llama_memory_seq_rm (mem_dft, 0, reuse_n, -1); if (!is_removed) { @@ -433,10 +421,6 @@ struct common_speculative_state_draft : public common_speculative_state { } } - if (needs_ckpt) { - ckpt.ckpt_size = draft_create_checkpoint(prompt_dft.size(), batch.n_tokens); - } - // prepare a batch to evaluate any new tokens in the prompt common_batch_clear(batch); @@ -450,12 +434,17 @@ struct common_speculative_state_draft : public common_speculative_state { // we should rarely end-up here during normal decoding if (batch.n_tokens > 0) { //LOG_DBG("%s: draft prompt batch: %s\n", __func__, string_from(ctx, batch).c_str()); + LOG_DBG("%s: draft prompt batch: %d tokens\n", __func__, batch.n_tokens); int ret = llama_decode(ctx_dft, batch); if (ret != 0 && ret != 1) { LOG_WRN("%s: llama_decode returned %d, prompt_cur.size=%zu\n", __func__, ret, prompt_cur.size()); } + + if (use_ckpt) { + create_checkpoint(prompt_dft.size()); + } } const llama_pos n_past = prompt_dft.size();