Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 61 additions & 1 deletion common/ngram-cache.h
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,9 @@

#include "llama.h"

#include <unordered_map>
#include <algorithm>
#include <string>
#include <unordered_map>
#include <vector>

#define LLAMA_NGRAM_MIN 1
Expand Down Expand Up @@ -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<llama_token> common_prompt_lookup_draft(const std::vector<llama_token> & prompt,
const std::vector<llama_token> & 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 };
}
20 changes: 20 additions & 0 deletions common/speculative.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<llama_tokens> lookup_prompts;
bool is_dflash2 = false;
bool is_mrope = false;
int32_t selector_top_k = 0;
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand Down
Loading