Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes479downloads
speculative.cpp1157 linesDownload Raw Back to common
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