From 20994d2d21649e8fdf735f6a14fae351e81a0e36 Mon Sep 17 00:00:00 2001 From: bri-prism <288398250+bri-prism@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:08:31 -0700 Subject: [PATCH] speculative: add opt-in prompt lookup for DFlash2 --- common/ngram-cache.h | 62 +++++++++++++++++++++++++++++++++++++++++- common/speculative.cpp | 20 ++++++++++++++ 2 files changed, 81 insertions(+), 1 deletion(-) diff --git a/common/ngram-cache.h b/common/ngram-cache.h index 6e7cfea966df..0ae6736af33e 100644 --- a/common/ngram-cache.h +++ b/common/ngram-cache.h @@ -2,8 +2,9 @@ #include "llama.h" -#include +#include #include +#include #include #define LLAMA_NGRAM_MIN 1 @@ -99,3 +100,62 @@ common_ngram_cache common_ngram_cache_load(const std::string & filename); // ngram_cache_target: the ngram cache to which to add the information from ngram_cache_add. // ngram_cache_add: the ngram cache to add to ngram_cache_target. void common_ngram_cache_merge(common_ngram_cache & ngram_cache_target, common_ngram_cache & ngram_cache_add); + +// Prompt-copy proposals only; the target must verify every returned token. +inline std::vector common_prompt_lookup_draft(const std::vector & prompt, + const std::vector & history, + llama_token anchor, + int depth) { + if (depth <= 0 || (size_t) depth >= prompt.size() || history.size() < prompt.size() || + !std::equal(prompt.begin(), prompt.end(), history.begin())) { + return {}; + } + const size_t n = (size_t) depth; + size_t chosen = prompt.size(); + if (history.size() == prompt.size()) { + for (size_t c = 1; c < prompt.size() - n; ++c) { + if (prompt[c - 1] == prompt.back() && prompt[c] == anchor) { + if (chosen != prompt.size()) { + return {}; + } + chosen = c + 1; + } + } + } else { + auto committed = history; + committed.push_back(anchor); + const size_t longest = std::min({ size_t(64), committed.size() - n, prompt.size() - n }); + if (longest < 16) { + return {}; + } + size_t best = 15; + bool ambiguous = false; + for (size_t end = 16; end <= prompt.size() - n; ++end) { + if (prompt[end - 1] != anchor) { + continue; + } + size_t length = 1; + while (length < std::min(longest, end) && + prompt[end - length - 1] == committed[committed.size() - length - 1]) { + ++length; + } + if (length < 16 || length < best) { + continue; + } + if (length > best) { + best = length; + chosen = end; + ambiguous = false; + } else if (!std::equal(prompt.begin() + chosen, prompt.begin() + chosen + n, prompt.begin() + end)) { + ambiguous = true; + } + } + if (ambiguous) { + return {}; + } + } + if (chosen == prompt.size()) { + return {}; + } + return { prompt.begin() + chosen, prompt.begin() + chosen + n }; +} diff --git a/common/speculative.cpp b/common/speculative.cpp index bca517f5ad43..7b53a51257dd 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -1658,6 +1658,8 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { int32_t block_size = 0; llama_token mask_token_id = 0; + bool prompt_lookup = false; + std::vector lookup_prompts; bool is_dflash2 = false; bool is_mrope = false; int32_t selector_top_k = 0; @@ -1724,6 +1726,11 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { selector_top_k = llama_model_dflash_selector_top_k(model_dft); is_dflash2 = selector_top_k > 0; + const char * lookup_env = std::getenv("LLAMA_DFLASH2_LOOKUP"); + prompt_lookup = is_dflash2 && lookup_env && std::strcmp(lookup_env, "1") == 0; + if (prompt_lookup) { + lookup_prompts.resize(n_seq); + } if (is_dflash2) { if (const char * value = getenv("LLAMA_DFLASH2_BEAM_WIDTH")) { const int width = std::atoi(value); @@ -1822,6 +1829,9 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { return; } + if (prompt_lookup) { + lookup_prompts[seq_id] = prompt; + } const int32_t N = (int32_t) prompt.size(); if (N <= 0) { return; @@ -1974,6 +1984,16 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { continue; } + if (prompt_lookup && dp.prompt) { + const int limit = dp.n_max > 0 ? std::min(params.n_max, dp.n_max) : params.n_max; + auto found = common_prompt_lookup_draft(lookup_prompts[seq_id], *dp.prompt, dp.id_last, limit); + if (!found.empty() && found.size() >= (size_t) params.n_min) { + *dp.result = std::move(found); + LOG_DBG("DFlash2 prompt lookup: seq=%d tokens=%zu, neural draft skipped\n", seq_id, + dp.result->size()); + continue; + } + } common_sampler_reset(smpls[seq_id].get()); const int32_t n = (int32_t) dp.n_past;