server: introduce self-speculative decoding

This commit is contained in:
Sascha Rogmann
2025-12-29 20:46:32 +01:00
parent 8f91ca54ec
commit 1fb2658b0d
6 changed files with 105 additions and 11 deletions
+57
View File
@@ -359,3 +359,60 @@ llama_tokens common_speculative_gen_draft(
}
return result;
}
llama_tokens common_speculative_gen_self_draft(const llama_tokens & tokens, llama_token sampled,
size_t n_draft_min, size_t n_draft_max) {
const size_t cur_len = tokens.size();
// vector for tokens we want to verify.
// return empty vector if there is no match.
llama_tokens draft_tokens;
if (cur_len <= static_cast<size_t>(n_draft_min + n_draft_max + 1)) {
return draft_tokens;
}
// pattern search
llama_tokens pattern;
pattern.reserve(n_draft_min);
for (size_t j = cur_len - n_draft_min + 1; j < cur_len; ++j) {
pattern.push_back(tokens[j]);
}
pattern.push_back(sampled); // add the last token to the pattern
size_t match_pos = 0; // we ignore position 0, position 0 == no match
// search backwards, but skip the current match (we are currently there)
for (size_t j = cur_len - n_draft_min - 1; j > 0; --j) {
bool match = true;
for (size_t k = 0; k < pattern.size(); ++k) {
if (tokens[j + k] != pattern[k]) {
match = false;
break;
}
}
if (match) {
match_pos = j;
break;
}
}
if (match_pos == 0) {
return draft_tokens;
}
const size_t copy_max = std::min(
n_draft_max,
cur_len - (match_pos + n_draft_min)
);
if (copy_max < n_draft_min) {
return draft_tokens;
}
LOG_DBG("%s: #tokens = %ld: found matching pattern at pos %ld, length %ld, draft length %ld\n",
__func__, (int64_t) cur_len,
(int64_t) match_pos, (int64_t) pattern.size(), copy_max);
draft_tokens.reserve(copy_max);
for (size_t j = 0; j < copy_max; ++j) {
draft_tokens.push_back(tokens[match_pos + n_draft_min + j]);
}
return draft_tokens;
}