echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0479
1#include "speculative.h"2 3#include "common.h"4#include "ggml.h"5#include "llama.h"6#include "log.h"7#include "ngram-cache.h"8#include "ngram-map.h"9#include "ngram-mod.h"10#include "sampling.h"11 12#include <algorithm>13#include <cstring>14#include <iomanip>15#include <map>16#include <cinttypes>17 18#define SPEC_VOCAB_MAX_SIZE_DIFFERENCE 12819#define SPEC_VOCAB_CHECK_START_TOKEN_ID 520 21const std::vector<enum common_speculative_type> common_speculative_types = {22 COMMON_SPECULATIVE_TYPE_NONE,23 COMMON_SPECULATIVE_TYPE_DRAFT,24 COMMON_SPECULATIVE_TYPE_EAGLE3,25 COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE,26 COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K,27 COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V,28 COMMON_SPECULATIVE_TYPE_NGRAM_MOD,29 COMMON_SPECULATIVE_TYPE_NGRAM_CACHE30};31 32const std::map<std::string, enum common_speculative_type> common_speculative_type_from_name_map = {33 {"none", COMMON_SPECULATIVE_TYPE_NONE},34 {"draft", COMMON_SPECULATIVE_TYPE_DRAFT},35 {"eagle3", COMMON_SPECULATIVE_TYPE_EAGLE3},36 {"ngram_simple", COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE},37 {"ngram_map_k", COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K},38 {"ngram_map_k4v", COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V},39 {"ngram_mod", COMMON_SPECULATIVE_TYPE_NGRAM_MOD},40 {"ngram_cache", COMMON_SPECULATIVE_TYPE_NGRAM_CACHE}41};42 43struct common_speculative_config {44 common_speculative_type type;45 common_params_speculative params;46 47 common_speculative_config(common_speculative_type t,48 const common_params_speculative & p = common_params_speculative{}) : type(t), params(p) {}49};50 51static bool common_speculative_are_compatible(52 const llama_model * model_tgt,53 const llama_model * model_dft) {54 const llama_vocab * vocab_tgt = llama_model_get_vocab(model_tgt);55 const llama_vocab * vocab_dft = llama_model_get_vocab(model_dft);56 57 const bool vocab_type_tgt = llama_vocab_type(vocab_tgt);58 LOG_DBG("%s: vocab_type tgt: %d\n", __func__, vocab_type_tgt);59 60 const bool vocab_type_dft = llama_vocab_type(vocab_dft);61 LOG_DBG("%s: vocab_type dft: %d\n", __func__, vocab_type_dft);62 63 if (vocab_type_tgt != vocab_type_dft) {64 LOG_DBG("%s: draft model vocab type must match target model to use speculation but ", __func__);65 LOG_DBG("vocab_type_dft = %d while vocab_type_tgt = %d\n", vocab_type_dft, vocab_type_tgt);66 return false;67 }68 69 if (70 llama_vocab_get_add_bos(vocab_tgt) != llama_vocab_get_add_bos(vocab_dft) ||71 llama_vocab_get_add_eos(vocab_tgt) != llama_vocab_get_add_eos(vocab_dft) ||72 llama_vocab_bos(vocab_tgt) != llama_vocab_bos(vocab_dft) ||73 llama_vocab_eos(vocab_tgt) != llama_vocab_eos(vocab_dft)74 ) {75 LOG_DBG("%s: draft model special tokens must match target model to use speculation\n", __func__);76 return false;77 }78 79 {80 const int n_vocab_tgt = llama_vocab_n_tokens(vocab_tgt);81 const int n_vocab_dft = llama_vocab_n_tokens(vocab_dft);82 const int vocab_diff = n_vocab_tgt > n_vocab_dft83 ? n_vocab_tgt - n_vocab_dft84 : n_vocab_dft - n_vocab_tgt;85 86 if (vocab_diff > SPEC_VOCAB_MAX_SIZE_DIFFERENCE) {87 LOG_DBG("%s: draft model vocab must closely match target model to use speculation but ", __func__);88 LOG_DBG("target vocab size %d does not match draft vocab size %d - difference %d, max allowed %d\n",89 n_vocab_tgt, llama_vocab_n_tokens(vocab_dft), vocab_diff, SPEC_VOCAB_MAX_SIZE_DIFFERENCE);90 return false;91 }92 93 for (int i = SPEC_VOCAB_CHECK_START_TOKEN_ID; i < std::min(n_vocab_tgt, n_vocab_dft); ++i) {94 const char * token_text_tgt = llama_vocab_get_text(vocab_tgt, i);95 const char * token_text_dft = llama_vocab_get_text(vocab_dft, i);96 97 if (std::strcmp(token_text_tgt, token_text_dft) != 0) {98 LOG_DBG("%s: draft model vocab must match target model to use speculation but ", __func__);99 LOG_DBG("token %d content differs - target '%s', draft '%s'\n", i,100 common_token_to_piece(vocab_tgt, i).c_str(),101 common_token_to_piece(vocab_dft, i).c_str());102 return false;103 }104 }105 }106 107 return true;108}109 110// state of an implementation of speculative decoding111//112// each implementation has a unique type and a state that is implementation-specific113// in a subclass of common_speculative_state114struct common_speculative_state {115 const enum common_speculative_type type;116 117 size_t n_call_begin = 0; // number of times this implementation was called for refresh.118 size_t n_call_draft = 0; // number of times this implementation was called for generation.119 size_t n_call_accept = 0; // number of times this implementation was called for accumulation.120 121 size_t n_gen_drafts = 0; // number of times a draft or part was generated by this implementation.122 size_t n_acc_drafts = 0; // number of times a draft or part was accepted by the target model.123 size_t n_gen_tokens = 0; // number of tokens generated by this implementation.124 size_t n_acc_tokens = 0; // number of tokens accepted by the target model.125 126 // TODO: track performance of most recent calls127 const bool gen_perf = true; // whether to generate performance stats.128 129 int64_t t_begin_us = 0; // total time spent in refresh of this implementation in microseconds.130 int64_t t_draft_us = 0; // total time spent in generating drafts in this implementation in microseconds.131 int64_t t_accept_us = 0; // total time spent in accumulation of this implementation in microseconds.132 133 common_speculative_state(enum common_speculative_type type) : type(type) {}134 135 virtual ~common_speculative_state() = default;136 137 virtual void begin(const llama_tokens & prompt) = 0;138 139 virtual void draft(140 const common_params_speculative & params,141 const llama_tokens & prompt_tgt,142 llama_token id_last,143 llama_tokens & result) = 0;144 145 virtual void accept(uint16_t n_accepted) = 0;146};147 148struct common_speculative_checkpoint {149 llama_pos pos_min = 0;150 llama_pos pos_max = 0;151 152 int64_t n_tokens = 0;153 154 std::vector<uint8_t> data;155 156 size_t size() const {157 return data.size();158 }159 160 size_t ckpt_size = 0;161};162 163struct common_speculative_state_draft : public common_speculative_state {164 llama_context * ctx_tgt; // only used for retokenizing from ctx_dft165 llama_context * ctx_dft;166 167 bool use_ckpt = false;168 struct common_speculative_checkpoint ckpt;169 170 common_sampler * smpl;171 172 llama_batch batch;173 llama_tokens prompt_dft;174 175 bool vocab_cmpt = true; // whether retokenization is needed176 std::unordered_map<std::string, std::string> vocab_map;177 178 common_speculative_state_draft(179 enum common_speculative_type type,180 llama_context * ctx_tgt,181 llama_context * ctx_dft,182 const std::vector<std::pair<std::string, std::string>> & replacements,183 bool use_ckpt)184 : common_speculative_state(type)185 , ctx_tgt(ctx_tgt)186 , ctx_dft(ctx_dft)187 , use_ckpt(use_ckpt)188 {189 batch = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);190 smpl = nullptr;191 192 // TODO: optimize or pass from outside?193 // {194 // common_params_sampling params;195 // params.no_perf = false;196 //197 // params.top_k = 40;198 // params.top_p = 0.9;199 //200 // params.samplers = {201 // COMMON_SAMPLER_TYPE_TOP_K,202 // COMMON_SAMPLER_TYPE_TOP_P,203 // COMMON_SAMPLER_TYPE_INFILL,204 // };205 //206 // result->smpl = common_sampler_init(llama_get_model(ctx_dft), params);207 // }208 {209 common_params_sampling params;210 params.no_perf = false;211 params.top_k = 10;212 params.samplers = {213 COMMON_SAMPLER_TYPE_TOP_K,214 };215 216 smpl = common_sampler_init(llama_get_model(ctx_dft), params);217 }218 219 vocab_cmpt = common_speculative_are_compatible(llama_get_model(ctx_tgt), llama_get_model(ctx_dft));220 LOG_DBG("vocab_cmpt = %d\n", vocab_cmpt);221 222 if (!vocab_cmpt) {223 LOG_WRN("the target and draft vocabs are not compatible - tokens will be translated between the two\n");224 225 for (const auto & pair : replacements) {226 vocab_map[pair.first] = pair.second;227 }228 }229 }230 231 ~common_speculative_state_draft() override {232 llama_perf_context_print(ctx_dft);233 234 llama_free(ctx_dft);235 236 common_sampler_free(smpl);237 238 llama_batch_free(batch);239 }240 241 void begin(const llama_tokens & prompt) override {242 if (use_ckpt && ckpt.size() > 0) {243 // delete checkpoint244 LOG_DBG("%s: delete checkpoint, prompt.size=%zu, pos_min=%d, pos_max=%d, n_tokens=%" PRId64 ", size=%.3f MiB\n",245 __func__, prompt.size(), ckpt.pos_min, ckpt.pos_max, ckpt.n_tokens, (float) ckpt.data.size() / 1024 / 1024);246 ckpt.pos_min = 0;247 ckpt.pos_max = 0;248 ckpt.n_tokens = 0;249 ckpt.ckpt_size = 0;250 ckpt.data.clear();251 }252 }253 254 size_t draft_create_checkpoint(int n_tokens_prompt, int n_tokens_batch) {255 int slot_id = 0;256 const size_t checkpoint_size = llama_state_seq_get_size_ext(ctx_dft, slot_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);257 258 ckpt.pos_min = llama_memory_seq_pos_min(llama_get_memory(ctx_dft), slot_id);259 ckpt.pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), slot_id);260 ckpt.n_tokens = n_tokens_prompt - n_tokens_batch;261 ckpt.data.resize(checkpoint_size);262 263 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);264 if (n != checkpoint_size) {265 GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", checkpoint_size, n);266 }267 268 LOG_DBG("%s: pos_min = %d, pos_max = %d, size = %.3f MiB\n", __func__,269 ckpt.pos_min, ckpt.pos_max, (float) ckpt.data.size() / 1024 / 1024);270 return n;271 }272 273 size_t draft_restore_checkpoint(size_t ckpt_size_part_expected) {274 int slot_id = 0;275 LOG_DBG("%s: pos_min = %d, pos_max = %d\n", __func__, ckpt.pos_min, ckpt.pos_max);276 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);277 if (n != ckpt_size_part_expected) {278 GGML_ABORT("%s: failed to restore context checkpoint (pos_min=%d, pos_max=%d, size=%zu, get_data_ext->%zu, set_data_ext->%zu",279 __func__, ckpt.pos_min, ckpt.pos_max, ckpt.size(), ckpt_size_part_expected, n);280 }281 llama_memory_seq_rm(llama_get_memory(ctx_dft), slot_id, ckpt.pos_max + 1, -1);282 283 return n;284 }285 286 void draft(287 const common_params_speculative & params,288 const llama_tokens & prompt_tgt,289 llama_token id_last,290 llama_tokens & result) override {291 auto * spec = this;292 293 auto & batch = spec->batch;294 auto & ctx_tgt = spec->ctx_tgt;295 auto & ctx_dft = spec->ctx_dft;296 auto & smpl = spec->smpl;297 auto & prompt_dft = spec->prompt_dft;298 299 auto * mem_dft = llama_get_memory(ctx_dft);300 301 int reuse_i = 0; // index of part to be reused in prompt_dft302 int reuse_n = 0; // length of part to be reused in prompt_dft303 304 const int n_ctx = llama_n_ctx(ctx_dft) - params.n_max;305 306 llama_tokens prompt_cnv;307 if (!spec->vocab_cmpt) {308 std::string text;309 310 text = common_detokenize(ctx_tgt, prompt_tgt, true);311 text = replace_to_dft(text);312 313 LOG_DBG("%s: main->draft detokenized string: '%s'\n", __func__, text.c_str());314 315 prompt_cnv = common_tokenize(ctx_dft, text, false, true);316 317 // convert id_last to draft vocab. llama_detokenize is called directly to avoid an allocation318 const auto * model_tgt = llama_get_model(ctx_tgt);319 const auto * vocab_tgt = llama_model_get_vocab(model_tgt);320 321 int32_t n_chars = llama_detokenize(vocab_tgt, &id_last, 1, nullptr, 0, false, false);322 GGML_ASSERT(n_chars < 0 && "failed to detokenize id_last");323 324 text.resize(-n_chars);325 llama_detokenize(vocab_tgt, &id_last, 1, text.data(), text.size(), false, false);326 text = replace_to_dft(text);327 328 LOG_DBG("main->draft detokenized id_last(%d): '%s'\n", id_last, text.c_str());329 id_last = common_tokenize(ctx_dft, text, false, true)[0];330 }331 332 const llama_tokens & prompt_cur = spec->vocab_cmpt ? prompt_tgt : prompt_cnv;333 334 const int i_start = std::max<int>(0, (int) prompt_cur.size() - n_ctx);335 336 // reuse as much as possible from the old draft context337 // ideally, the draft context should be as big as the target context and we will always reuse the entire prompt338 for (int i = 0; i < (int) prompt_dft.size(); ++i) {339 int cur = 0;340 while (i_start + cur < (int) prompt_cur.size() &&341 i + cur < (int) prompt_dft.size() &&342 prompt_cur[i_start + cur] == prompt_dft[i + cur]) {343 cur++;344 }345 346 if ((cur >= 256 || n_ctx >= (int) prompt_cur.size()) && cur > reuse_n) {347 reuse_i = i;348 reuse_n = cur;349 }350 }351 352 LOG_DBG("%s: reuse_i = %d, reuse_n = %d, #prompt_dft = %zu, #prompt_cur = %zu\n",353 __func__, reuse_i, reuse_n, prompt_dft.size(), prompt_cur.size());354 if (use_ckpt && ckpt.ckpt_size == 0 && reuse_n > 0) {355 LOG_DBG("%s: no checkpoint available, no reuse, (reuse_i=%d, reuse_n=%d) -> (0, 0)\n",356 __func__, reuse_i, reuse_n);357 reuse_i = 0;358 reuse_n = 0;359 }360 361 result.clear();362 result.reserve(params.n_max);363 364 bool needs_ckpt = use_ckpt && prompt_dft.size() > 0;365 if (reuse_n == 0 || (use_ckpt && reuse_i > 0)) {366 llama_memory_clear(mem_dft, false);367 prompt_dft.clear();368 } else {369 // this happens when a previous draft has been discarded (for example, due to being too small), but the370 // target model agreed with it. in this case, we simply pass back the previous results to save compute371 if (reuse_i + reuse_n < (int64_t) prompt_dft.size() && prompt_dft[reuse_i + reuse_n] == id_last) {372 for (int i = reuse_i + reuse_n + 1; i < (int) prompt_dft.size(); ++i) {373 result.push_back(prompt_dft[i]);374 375 if (params.n_max <= (int) result.size()) {376 break;377 }378 }379 380 return;381 }382 383 bool do_restore = false;384 if (prompt_dft.size() > prompt_cur.size() && reuse_i + reuse_n < (int64_t) prompt_dft.size()) {385 // This can happen after a partial acceptance (speculative decoding with checkpoints)386 LOG_DBG("%s: #prompt_dft=%zu, #prompt_cur=%zu, shorten draft\n",387 __func__, prompt_dft.size(), prompt_cur.size());388 prompt_dft.resize(prompt_cur.size());389 do_restore = true;390 }391 392 if (reuse_i > 0) {393 bool is_removed = llama_memory_seq_rm (mem_dft, 0, 0, reuse_i);394 if (!is_removed) {395 LOG_ERR("%s: llama_memory_seq_rm failed, reuse_i=%d\n", __func__, reuse_i);396 }397 llama_memory_seq_add(mem_dft, 0, reuse_i, -1, -reuse_i);398 399 prompt_dft.erase(prompt_dft.begin(), prompt_dft.begin() + reuse_i);400 }401 402 if (reuse_n < (int) prompt_dft.size() || do_restore) {403 if (use_ckpt) {404 if (ckpt.n_tokens > (int64_t) prompt_dft.size()) {405 LOG_INF("%s: checkpoint is too large, prompt_tgt.size=%zu, ckpt.n_tokens=%" PRId64 ", reuse_n=%d, prompt_dft.size=%zu\n",406 __func__, prompt_tgt.size(), ckpt.n_tokens, reuse_n, prompt_dft.size());407 }408 draft_restore_checkpoint(ckpt.ckpt_size);409 reuse_n = ckpt.n_tokens;410 prompt_dft.resize(reuse_n);411 needs_ckpt = false;412 } else {413 bool is_removed = llama_memory_seq_rm (mem_dft, 0, reuse_n, -1);414 if (!is_removed) {415 LOG_ERR("%s: llama_memory_seq_rm failed, reuse_n=%d, prompt_dft.size=%zu\n",416 __func__, reuse_n, prompt_dft.size());417 }418 prompt_dft.erase(prompt_dft.begin() + reuse_n, prompt_dft.end());419 }420 }421 }422 423 if (needs_ckpt) {424 ckpt.ckpt_size = draft_create_checkpoint(prompt_dft.size(), batch.n_tokens);425 }426 427 // prepare a batch to evaluate any new tokens in the prompt428 common_batch_clear(batch);429 430 for (size_t i = i_start + reuse_n; i < prompt_cur.size(); ++i) {431 //LOG_DBG("i = %d, i_start = %d, reuse_n = %d, i - i_start = %d, id = %6d\n", i, i_start, reuse_n, i - i_start, prompt_cur[i]);432 common_batch_add(batch, prompt_cur[i], i - i_start, { 0 }, false);433 434 prompt_dft.push_back(prompt_cur[i]);435 }436 437 // we should rarely end-up here during normal decoding438 if (batch.n_tokens > 0) {439 //LOG_DBG("%s: draft prompt batch: %s\n", __func__, string_from(ctx, batch).c_str());440 441 int ret = llama_decode(ctx_dft, batch);442 if (ret != 0 && ret != 1) {443 LOG_WRN("%s: llama_decode returned %d, prompt_cur.size=%zu\n",444 __func__, ret, prompt_cur.size());445 }446 }447 448 const llama_pos n_past = prompt_dft.size();449 450 LOG_DBG("%s: n_past = %d\n", __func__, n_past);451 452 common_batch_clear(batch);453 common_batch_add (batch, id_last, n_past, { 0 }, true);454 455 prompt_dft.push_back(id_last);456 457 LOG_DBG("%s: draft prompt: %s\n", __func__, string_from(ctx_dft, prompt_dft).c_str());458 459 int ret = llama_decode(ctx_dft, batch);460 if (ret != 0 && ret != 1) {461 LOG_WRN("%s: llama_decode returned %d, prompt_cur.size=%zu, prompt_dft.size=%zu\n",462 __func__, ret, prompt_cur.size(), prompt_dft.size());463 }464 465 common_sampler_reset(smpl);466 467 // sample n_draft tokens from the draft model468 for (int i = 0; i < params.n_max; ++i) {469 common_batch_clear(batch);470 471 common_sampler_sample(smpl, ctx_dft, 0, true);472 473 const auto * cur_p = common_sampler_get_candidates(smpl, true);474 475 for (int k = 0; k < std::min(3, (int) cur_p->size); ++k) {476 LOG_DBG(" - draft candidate %3d, pos %3d: %6d (%8.3f) '%s'\n",477 k, i, cur_p->data[k].id, cur_p->data[k].p, common_token_to_piece(ctx_dft, cur_p->data[k].id).c_str());478 }479 480 // add drafted token for each sequence481 const llama_token id = cur_p->data[0].id;482 483 common_sampler_accept(smpl, id, true);484 485 result.push_back(id);486 487 if (params.n_max <= (int) result.size()) {488 break;489 }490 491 // only collect very high-confidence draft tokens492 if (cur_p->data[0].p < params.p_min) {493 break;494 }495 496 common_batch_add(batch, id, n_past + i + 1, { 0 }, true);497 498 // evaluate the drafted tokens on the draft model499 ret = llama_decode(ctx_dft, batch);500 if (ret != 0) {501 LOG_WRN("%s: llama_decode[%d] returned %d, prompt_cur.size=%zu, prompt_dft.size=%zu\n",502 __func__, i, ret, prompt_cur.size(), prompt_dft.size());503 }504 505 prompt_dft.push_back(id);506 }507 508 if (!spec->vocab_cmpt) {509 std::string detokenized = common_detokenize(ctx_dft, result, true);510 detokenized = replace_to_tgt(detokenized);511 LOG_DBG("draft->main detokenized string: '%s'\n", detokenized.c_str());512 result = common_tokenize(ctx_tgt, detokenized, false, true);513 if (result.size() > (size_t)params.n_max) {514 result.resize(params.n_max);515 }516 }517 }518 519 void accept(uint16_t n_accepted) override {520 // noop521 GGML_UNUSED(n_accepted);522 }523 524 std::string replace_to_dft(const std::string & input) const {525 std::string result = input;526 527 for (const auto & pair : this->vocab_map) {528 size_t pos = result.find(pair.first);529 while (pos != std::string::npos) {530 result.replace(pos, pair.first.length(), pair.second);531 pos = result.find(pair.first, pos + pair.second.length());532 }533 }534 535 return result;536 }537 538 std::string replace_to_tgt(const std::string & input) const {539 std::string result = input;540 541 for (const auto & pair : this->vocab_map) {542 size_t pos = result.find(pair.second);543 while (pos != std::string::npos) {544 result.replace(pos, pair.second.length(), pair.first);545 pos = result.find(pair.second, pos + pair.first.length());546 }547 }548 549 return result;550 }551};552 553struct common_speculative_state_eagle3 : public common_speculative_state {554 common_speculative_state_eagle3(enum common_speculative_type type) : common_speculative_state(type) {}555 556 void begin(const llama_tokens & prompt) override {557 GGML_UNUSED(prompt);558 }559 560 void draft(561 const common_params_speculative & params,562 const llama_tokens & prompt_tgt,563 llama_token id_last,564 llama_tokens & draft_tokens) override {565 // TODO: implement566 GGML_UNUSED(params);567 GGML_UNUSED(prompt_tgt);568 GGML_UNUSED(id_last);569 GGML_UNUSED(draft_tokens);570 }571 572 void accept(uint16_t n_accepted) override {573 // noop574 GGML_UNUSED(n_accepted);575 }576};577 578// state of self-speculation (simple implementation, not ngram-map)579struct common_speculative_state_ngram_simple : public common_speculative_state {580 common_ngram_simple_config config;581 582 common_speculative_state_ngram_simple(583 enum common_speculative_type type,584 common_ngram_simple_config config)585 : common_speculative_state(type), config(config) {}586 587 void begin(const llama_tokens & prompt) override {588 GGML_UNUSED(prompt);589 }590 591 void draft(592 const common_params_speculative & params,593 const llama_tokens & prompt_tgt,594 llama_token id_last,595 llama_tokens & result) override {596 597 result = common_ngram_simple_draft(config, prompt_tgt, id_last);598 GGML_UNUSED(params);599 }600 601 void accept(uint16_t n_accepted) override {602 // noop603 GGML_UNUSED(n_accepted);604 }605};606 607struct common_speculative_state_ngram_map_k : public common_speculative_state {608 // draft ngram map for speculative decoding without draft model609 common_ngram_map map;610 611 common_speculative_state_ngram_map_k(612 enum common_speculative_type type,613 common_ngram_map map)614 : common_speculative_state(type), map(std::move(map)) {}615 616 void begin(const llama_tokens & prompt) override {617 common_ngram_map_begin(map, prompt);618 }619 620 void draft(621 const common_params_speculative & params,622 const llama_tokens & prompt_tgt,623 llama_token id_last,624 llama_tokens & result) override {625 common_ngram_map_draft(map, prompt_tgt, id_last, result);626 GGML_UNUSED(params);627 }628 629 void accept(uint16_t n_accepted) override {630 common_ngram_map_accept(map, n_accepted);631 }632};633 634struct common_speculative_state_ngram_mod : public common_speculative_state {635 common_ngram_mod & mod;636 637 // the last position in the prompt that was added to the ngram container638 size_t i_last = 0;639 640 // length of the last drafted nโgram (number of tokens returned by draft)641 size_t n_draft_last = 0;642 643 // consecutive accept rounds with low acceptance fraction (< 0.5)644 int n_low = 0;645 646 // enable trace logging if LLAMA_TRACE is set647 const bool verbose;648 649 common_speculative_state_ngram_mod(enum common_speculative_type type, common_ngram_mod & mod)650 : common_speculative_state(type), mod(mod), verbose(std::getenv("LLAMA_TRACE") != nullptr) {651 static_assert(sizeof(llama_token) == sizeof(common_ngram_mod::entry_t));652 }653 654 void begin(const llama_tokens & prompt) override {655 i_last = 0;656 657 n_draft_last = 0;658 659 const size_t n = mod.get_n();660 661 if (prompt.size() < n) {662 return;663 }664 665 for (size_t i = 0; i < prompt.size() - n; ++i) {666 mod.add(prompt.data() + i);667 }668 669 i_last = prompt.size() - n;670 671 const double f = (double)mod.get_used() / (double)mod.size();672 LOG_INF("%s: ngram_mod occupancy = %zu/%zu (%.2f)\n", __func__, mod.get_used(), mod.size(), f);673 674 constexpr double f_thold = 0.25;675 if (f > f_thold) {676 LOG_WRN("%s: ngram_mod occupancy %.2f exceeds threshold (%.2f) - resetting\n", __func__, f, f_thold);677 678 mod.reset();679 }680 }681 682 void draft(683 const common_params_speculative & params,684 const llama_tokens & prompt_tgt,685 llama_token id_last,686 llama_tokens & result) override {687 GGML_UNUSED(params);688 689 n_draft_last = 0;690 691 const size_t cur_len = prompt_tgt.size();692 if (cur_len < mod.get_n()) {693 return;694 }695 696 const size_t n = mod.get_n();697 698 // add new ngrams in chunks699 if (i_last + 32 < cur_len) {700 for (size_t i = i_last; i < cur_len - n; ++i) {701 mod.add(prompt_tgt.data() + i);702 }703 704 i_last = cur_len - n;705 }706 707 result.resize(n + params.n_max);708 for (size_t i = 0; i < n - 1; ++i) {709 result[i] = prompt_tgt[cur_len - n + 1 + i];710 }711 result[n - 1] = id_last;712 713 for (int i = 0; i < params.n_max; ++i) {714 const llama_token token = mod.get(result.data() + i);715 if (token == common_ngram_mod::EMPTY) {716 if (i < params.n_min) {717 result.clear();718 return;719 }720 721 result.resize(n + i);722 break;723 }724 result[n + i] = token;725 }726 727 // only return the m tokens that were drafted728 for (size_t i = 0; n + i < result.size(); ++i) {729 result[i] = result[n + i];730 }731 result.resize(result.size() - n);732 733 // store length of drafted nโgram for later acceptance analysis734 n_draft_last = result.size();735 }736 737 void accept(uint16_t n_accepted) override {738 if (verbose) {739 LOG_INF("%s: accepted %d tokens from %zu drafted tokens\n", __func__, n_accepted, n_draft_last);740 }741 742 // compute acceptance fraction if we have a recorded draft length743 if (n_draft_last > 0) {744 const double f_acc = (double)n_accepted / (double)n_draft_last;745 if (f_acc < 0.5) {746 n_low++;747 if (n_low >= 3) {748 LOG_WRN("%s: low acceptance streak (%d) โ resetting ngram_mod\n", __func__, n_low);749 750 mod.reset();751 n_low = 0;752 }753 } else {754 n_low = 0;755 }756 }757 }758};759 760struct common_speculative_state_ngram_cache : public common_speculative_state {761 uint16_t n_draft;762 bool save_dynamic;763 bool save_static;764 765 common_ngram_cache ngram_cache_context;766 common_ngram_cache ngram_cache_dynamic;767 common_ngram_cache ngram_cache_static;768 769 size_t cache_size = 0; // number of tokens in n-gram cache770 771 common_speculative_state_ngram_cache(772 const enum common_speculative_type type,773 const std::string & path_static,774 const std::string & path_dynamic,775 uint16_t n_draft,776 bool save_dynamic,777 bool save_static)778 : common_speculative_state(type)779 , n_draft(n_draft)780 , save_dynamic(save_dynamic)781 , save_static(save_static)782 {783 if (!path_static.empty()) {784 try {785 ngram_cache_static = common_ngram_cache_load(path_static);786 } catch (...) {787 LOG_ERR("failed to open static lookup cache: %s", path_static.c_str());788 GGML_ABORT("Couldn't read static lookup cache");789 }790 }791 792 if (!path_dynamic.empty()) {793 try {794 ngram_cache_dynamic = common_ngram_cache_load(path_dynamic);795 } catch (...) {796 LOG_ERR("failed to open dynamic lookup cache: %s", path_dynamic.c_str());797 GGML_ABORT("Couldn't read dynamic lookup cache");798 }799 }800 }801 802 void begin(const llama_tokens & prompt) override {803 GGML_UNUSED(prompt);804 }805 806 void draft(807 const common_params_speculative & params,808 const llama_tokens & prompt_tgt,809 llama_token id_last,810 llama_tokens & result) override {811 GGML_UNUSED(params);812 813 if (cache_size < prompt_tgt.size() + 1) {814 llama_tokens tokens_new;815 tokens_new.reserve(prompt_tgt.size() + 1 - cache_size);816 for (size_t j = cache_size; j < prompt_tgt.size(); ++j) {817 tokens_new.push_back(prompt_tgt[j]);818 }819 tokens_new.push_back(id_last); // add the last token820 821 // Update context ngram cache with new prompt_tgt:822 common_ngram_cache_update(ngram_cache_context, LLAMA_NGRAM_MIN, LLAMA_NGRAM_MAX,823 tokens_new, tokens_new.size(), false);824 cache_size = prompt_tgt.size() + 1;825 }826 827 llama_tokens inp;828 inp.reserve(prompt_tgt.size() + 1);829 for (size_t j = 0; j < prompt_tgt.size(); ++j) {830 inp.push_back(prompt_tgt[j]);831 }832 inp.push_back(id_last);833 834 result.push_back(id_last);835 836 common_ngram_cache_draft(inp, result, n_draft, LLAMA_NGRAM_MIN, LLAMA_NGRAM_MAX,837 ngram_cache_context,838 ngram_cache_dynamic,839 ngram_cache_static);840 841 if (result.size() > 0) {842 // delete first token in result (which is the id_last token)843 result.erase(result.begin());844 }845 }846 847 void accept(uint16_t n_accepted) override {848 // TODO: noop849 GGML_UNUSED(n_accepted);850 }851};852 853struct common_speculative {854 std::vector<std::unique_ptr<common_speculative_state>> impls; // list of implementations to use and their states855 856 common_speculative_state * curr_impl = nullptr; // current implementation in use (for stats)857};858 859static common_ngram_map get_common_ngram_map(const common_speculative_config & config) {860 uint16_t size_key = config.params.ngram_size_n;861 uint16_t size_value = config.params.ngram_size_m;862 bool key_only = (config.type == COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K);863 uint16_t min_hits = config.params.ngram_min_hits;864 865 return common_ngram_map(size_key, size_value, key_only, min_hits);866}867 868static common_speculative_state_ngram_cache create_state_ngram_cache(869 const std::string & path_static, const std::string & path_dynamic,870 const common_speculative_config & config) {871 uint16_t n_draft = 8; // TODO get from config?872 873 // TODO bool param in common/common.h to set save_static/save_dynamic?874 bool save_static = false;875 bool save_dynamic = false;876 877 common_speculative_state_ngram_cache state(config.type, path_static, path_dynamic, n_draft, save_static, save_dynamic);878 879 return state;880}881 882std::string common_speculative_type_name_str() {883 std::string result;884 for (size_t i = 0; i < common_speculative_types.size(); i++) {885 if (i > 0) {886 result += ", ";887 }888 result += common_speculative_type_to_str(common_speculative_types[i]);889 }890 return result;891}892 893std::string common_speculative_type_to_str(enum common_speculative_type type) {894 switch (type) {895 case COMMON_SPECULATIVE_TYPE_NONE: return "none";896 case COMMON_SPECULATIVE_TYPE_DRAFT: return "draft";897 case COMMON_SPECULATIVE_TYPE_EAGLE3: return "eagle3";898 case COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE: return "ngram_simple";899 case COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K: return "ngram_map_k";900 case COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V: return "ngram_map_k4v";901 case COMMON_SPECULATIVE_TYPE_NGRAM_MOD: return "ngram_mod";902 case COMMON_SPECULATIVE_TYPE_NGRAM_CACHE: return "ngram_cache";903 default: return "unknown";904 }905}906 907enum common_speculative_type common_speculative_type_from_name(const std::string & name) {908 const auto it = common_speculative_type_from_name_map.find(name);909 if (it == common_speculative_type_from_name_map.end()) {910 return COMMON_SPECULATIVE_TYPE_COUNT;911 }912 return it->second;913}914 915// initialization of the speculative decoding system916//917common_speculative * common_speculative_init(918 common_params_speculative & params,919 llama_context * ctx_tgt) {920 llama_context * ctx_dft = nullptr;921 if (params.model_dft) {922 ctx_dft = llama_init_from_model(params.model_dft, params.cparams_dft);923 if (ctx_dft == nullptr) {924 LOG_ERR("%s", "failed to create draft context\n");925 return nullptr;926 }927 }928 929 // Compute the implementations to use based on the config and their order of preference930 std::vector<common_speculative_config> configs = {}; // list of speculative configs to try931 {932 bool has_draft = !params.mparams_dft.path.empty();933 bool has_draft_eagle3 = false; // TODO PR-18039: if params.speculative.eagle3934 935 bool has_ngram_cache = (params.type == COMMON_SPECULATIVE_TYPE_NGRAM_CACHE);936 bool has_ngram_simple = (params.type == COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE);937 bool has_ngram_map_k = (params.type == COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K);938 bool has_ngram_map_k4v = (params.type == COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V);939 bool has_ngram_mod = (params.type == COMMON_SPECULATIVE_TYPE_NGRAM_MOD);940 941 // In a more complex implementation we could use the same implementation but with different parameters.942 // This was initially used in PR-18471 but removed to simplify the code.943 if (has_ngram_simple) {944 // This implementation can guess a lot of tokens without any draft model.945 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE, params));946 }947 if (has_ngram_map_k) {948 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K, params));949 }950 if (has_ngram_map_k4v) {951 // This implementation can guess tokens with high acceptance rate but is more expensive.952 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V, params));953 }954 if (has_ngram_mod) {955 // shared instance for all speculative decoding contexts956 if (!params.ngram_mod) {957 params.ngram_mod = std::make_shared<common_ngram_mod>(params.ngram_size_n, 4*1024*1024);958 959 LOG_INF("%s: initialized ngram_mod with n=%d, size=%zu (%.3f MB)\n", __func__,960 params.ngram_size_n, params.ngram_mod->size(),961 (float)(params.ngram_mod->size_bytes())/1024/1024);962 963 if (params.ngram_size_n < 16) {964 LOG_WRN("%s: ngram_mod n=%d is too small - poor quality is possible, see: https://github.com/ggml-org/llama.cpp/pull/19164\n", __func__, params.ngram_size_n);965 }966 }967 968 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_MOD, params));969 }970 if (has_ngram_cache) {971 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_CACHE, params));972 }973 if (has_draft) {974 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_DRAFT, params));975 }976 if (has_draft_eagle3) {977 configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_EAGLE3, params));978 }979 }980 981 std::vector<std::unique_ptr<common_speculative_state>> impls = {};982 983 for (const common_speculative_config & config : configs) {984 LOG_DBG("%s: adding implementation %s\n", __func__, common_speculative_type_to_str(config.type).c_str());985 switch (config.type) {986 case COMMON_SPECULATIVE_TYPE_NONE:987 break;988 case COMMON_SPECULATIVE_TYPE_DRAFT: {989 const bool use_ckpt = common_context_can_seq_rm(ctx_dft) == COMMON_CONTEXT_SEQ_RM_TYPE_FULL;990 991 impls.push_back(std::make_unique<common_speculative_state_draft>(config.type,992 /* .ctx_tgt = */ ctx_tgt,993 /* .ctx_dft = */ ctx_dft,994 /* .replacements = */ params.replacements,995 /* .use_ckpt = */ use_ckpt996 ));997 break;998 }999 case COMMON_SPECULATIVE_TYPE_EAGLE3: {1000 impls.push_back(std::make_unique<common_speculative_state_eagle3>(config.type));1001 break;1002 }1003 case COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE: {1004 common_ngram_map ngram_map = get_common_ngram_map(config);1005 1006 uint16_t ngram_size_key = ngram_map.size_key;1007 uint16_t mgram_size_value = ngram_map.size_value;1008 1009 auto config_simple = common_ngram_simple_config {1010 /* .size_ngram = */ ngram_size_key,1011 /* .size_mgram = */ mgram_size_value1012 };1013 auto state = std::make_unique<common_speculative_state_ngram_simple>(1014 /* .type = */ config.type,1015 /* .state = */ config_simple1016 );1017 impls.push_back(std::move(state));1018 break;1019 }1020 case COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K:1021 case COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V: {1022 impls.push_back(std::make_unique<common_speculative_state_ngram_map_k>(1023 (config.type),1024 get_common_ngram_map(config)1025 ));1026 break;1027 }1028 case COMMON_SPECULATIVE_TYPE_NGRAM_MOD: {1029 GGML_ASSERT(config.params.ngram_mod);1030 impls.push_back(std::make_unique<common_speculative_state_ngram_mod>(config.type, *config.params.ngram_mod));1031 break;1032 }1033 case COMMON_SPECULATIVE_TYPE_NGRAM_CACHE: {1034 auto state = create_state_ngram_cache(1035 params.lookup_cache_static, params.lookup_cache_dynamic, config);1036 impls.push_back(std::make_unique<common_speculative_state_ngram_cache>(state));1037 break;1038 }1039 default:1040 break;1041 }1042 }1043 1044 if (impls.empty()) {1045 LOG_WRN("%s", "no implementations specified for speculative decoding\n");1046 return nullptr;1047 }1048 1049 auto * result = new common_speculative {1050 /* .impls = */ std::move(impls),1051 /* .curr_impl = */ nullptr,1052 };1053 1054 return result;1055}1056 1057void common_speculative_free(common_speculative * spec) {1058 if (spec == nullptr) {1059 return;1060 }1061 1062 delete spec;1063}1064 1065void common_speculative_begin(common_speculative * spec, const llama_tokens & prompt) {1066 if (spec == nullptr) {1067 return;1068 }1069 1070 for (auto & impl : spec->impls) {1071 common_time_meas tm(impl->t_begin_us, !impl->gen_perf);1072 impl->begin(prompt);1073 impl->n_call_begin++;1074 }1075}1076 1077llama_tokens common_speculative_draft(1078 common_speculative * spec,1079 const common_params_speculative & params,1080 const llama_tokens & prompt_tgt, // specified in target model vocab1081 llama_token id_last) {1082 llama_tokens result;1083 1084 spec->curr_impl = nullptr; // reset current implementation1085 1086 for (auto & impl : spec->impls) {1087 {1088 common_time_meas tm(impl->t_draft_us, !impl->gen_perf);1089 impl->draft(params, prompt_tgt, id_last, result);1090 impl->n_call_draft++;1091 }1092 1093 if (!result.empty()) {1094 LOG_DBG("%s: called impl %s, hist size = %zu, call_count = %zu, gen = %zu\n", __func__,1095 common_speculative_type_to_str(impl.get()->type).c_str(), prompt_tgt.size(),1096 impl.get()->n_call_draft, result.size());1097 1098 spec->curr_impl = impl.get(); // set current implementation for stats1099 impl->n_gen_drafts++;1100 impl->n_gen_tokens += result.size();1101 1102 break; // We have a draft, so break out of the loop and return it.1103 }1104 }1105 1106 return result;1107}1108 1109void common_speculative_accept(common_speculative * spec, uint16_t n_accepted) {1110 if (n_accepted == 0) {1111 return;1112 }1113 1114 common_speculative_state * impl = spec->curr_impl;1115 1116 GGML_ASSERT(impl);1117 1118 {1119 common_time_meas tm(impl->t_accept_us, !impl->gen_perf);1120 if (n_accepted > 0) {1121 impl->n_acc_drafts++;1122 impl->n_acc_tokens += n_accepted;1123 }1124 1125 impl->accept(n_accepted);1126 impl->n_call_accept++;1127 }1128}1129 1130void common_speculative_print_stats(const common_speculative * spec) {1131 if (spec == nullptr) {1132 return;1133 }1134 1135 for (const auto & impl : spec->impls) {1136 std::string str_perf;1137 if (impl->gen_perf) {1138 std::ostringstream oss;1139 oss << std::fixed << std::setprecision(3) << impl->t_begin_us / 1000.0 << ", ";1140 oss << std::fixed << std::setprecision(3) << impl->t_draft_us / 1000.0 << ", ";1141 oss << std::fixed << std::setprecision(3) << impl->t_accept_us / 1000.0;1142 str_perf = ", dur(b,g,a) = " + oss.str() + " ms";1143 } else {1144 str_perf = "";1145 }1146 1147 LOG_INF("statistics %s: #calls(b,g,a) = %zu %zu %zu, #gen drafts = %zu, #acc drafts = %zu, #gen tokens = %zu, #acc tokens = %zu%s\n",1148 common_speculative_type_to_str(impl->type).c_str(),1149 impl->n_call_begin, impl->n_call_draft, impl->n_call_accept,1150 impl->n_gen_drafts,1151 impl->n_acc_drafts,1152 impl->n_gen_tokens,1153 impl->n_acc_tokens,1154 str_perf.c_str());1155 }1156}1157 