Team Ai
Datasetpublic

Brunobkr/llama.cpp_AlgMor24_github

ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes3kdownloads
speculative.cpp2775 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 "../src/llama-ext.h" // staging API: llama_set_embeddings_nextn / llama_get_embeddings_nextn_ith (used by MTP)13 14#include <algorithm>15#include <cassert>16#include <cstring>17#include <iomanip>18#include <map>19#include <cinttypes>20 21#define SPC_DBG(fmt, ...) LOG_DBG("spec %12.*s: " fmt, 12, __func__, __VA_ARGS__)22#define SPC_TRC(fmt, ...) LOG_TRC("spec %12.*s: " fmt, 12, __func__, __VA_ARGS__)23#define SPC_INF(fmt, ...) LOG_INF("spec %12.*s: " fmt, 12, __func__, __VA_ARGS__)24#define SPC_WRN(fmt, ...) LOG_WRN("spec %12.*s: " fmt, 12, __func__, __VA_ARGS__)25#define SPC_ERR(fmt, ...) LOG_ERR("spec %12.*s: " fmt, 12, __func__, __VA_ARGS__)26#define SPC_CNT(fmt, ...) LOG_CNT(""              fmt,               __VA_ARGS__)27 28#define SPEC_VOCAB_MAX_SIZE_DIFFERENCE  12829#define SPEC_VOCAB_CHECK_START_TOKEN_ID 530 31const std::map<std::string, common_speculative_type> common_speculative_type_from_name_map = {32    {"none",          COMMON_SPECULATIVE_TYPE_NONE},33    {"draft-simple",  COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE},34    {"draft-eagle3",  COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3},35    {"draft-mtp",     COMMON_SPECULATIVE_TYPE_DRAFT_MTP},36    {"draft-dflash",  COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH},37    {"draft-dspark",  COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK},38    {"ngram-simple",  COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE},39    {"ngram-map-k",   COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K},40    {"ngram-map-k4v", COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V},41    {"ngram-mod",     COMMON_SPECULATIVE_TYPE_NGRAM_MOD},42    {"ngram-cache",   COMMON_SPECULATIVE_TYPE_NGRAM_CACHE}43};44 45static std::string common_speculative_get_devices_str(const std::vector<ggml_backend_dev_t> & devices) {46    std::string result;47    for (size_t i = 0; i < devices.size(); i++) {48        if (devices[i] == nullptr) {49            continue;50        }51        if (!result.empty()) result += ", ";52        result += ggml_backend_dev_name(devices[i]);53    }54    return result.empty() ? "default" : result;55}56 57struct common_speculative_config {58    common_speculative_type type;59    common_params_speculative params;60 61    common_speculative_config(common_speculative_type t,62            const common_params_speculative & p = common_params_speculative{}) : type(t), params(p) {}63};64 65static bool common_speculative_are_compatible(66    const llama_model * model_tgt,67    const llama_model * model_dft) {68    const llama_vocab * vocab_tgt = llama_model_get_vocab(model_tgt);69    const llama_vocab * vocab_dft = llama_model_get_vocab(model_dft);70 71    const auto vocab_type_tgt = llama_vocab_type(vocab_tgt);72    SPC_DBG("vocab_type tgt: %d\n", vocab_type_tgt);73 74    const auto vocab_type_dft = llama_vocab_type(vocab_dft);75    SPC_DBG("vocab_type dft: %d\n", vocab_type_dft);76 77    if (vocab_type_tgt != vocab_type_dft) {78        SPC_WRN("draft model vocab type must match target model to use speculation but "79                "vocab_type_dft = %d while vocab_type_tgt = %d\n", vocab_type_dft, vocab_type_tgt);80        return false;81    }82 83    if (llama_vocab_get_add_bos(vocab_tgt) != llama_vocab_get_add_bos(vocab_dft) ||84        (llama_vocab_get_add_bos(vocab_tgt) && llama_vocab_bos(vocab_tgt) != llama_vocab_bos(vocab_dft))) {85        SPC_WRN("draft model bos tokens must match target model to use speculation. add: %d - %d, id: %d - %d)\n",86                llama_vocab_get_add_bos(vocab_tgt), llama_vocab_get_add_bos(vocab_dft),87                llama_vocab_bos(vocab_tgt), llama_vocab_bos(vocab_dft));88        return false;89    }90 91    if (llama_vocab_get_add_eos(vocab_tgt) != llama_vocab_get_add_eos(vocab_dft) ||92        (llama_vocab_get_add_eos(vocab_tgt) && llama_vocab_eos(vocab_tgt) != llama_vocab_eos(vocab_dft))) {93        SPC_WRN("draft model eos tokens must match target model to use speculation. add: %d - %d, id: %d - %d)\n",94                llama_vocab_get_add_eos(vocab_tgt), llama_vocab_get_add_eos(vocab_dft),95                llama_vocab_eos(vocab_tgt), llama_vocab_eos(vocab_dft));96        return false;97    }98 99    {100        const int n_vocab_tgt = llama_vocab_n_tokens(vocab_tgt);101        const int n_vocab_dft = llama_vocab_n_tokens(vocab_dft);102        const int vocab_diff  = n_vocab_tgt > n_vocab_dft103            ? n_vocab_tgt - n_vocab_dft104            : n_vocab_dft - n_vocab_tgt;105 106        if (vocab_diff > SPEC_VOCAB_MAX_SIZE_DIFFERENCE) {107            SPC_DBG("draft model vocab must closely match target model to use speculation but "108                    "target vocab size %d does not match draft vocab size %d - difference %d, max allowed %d\n",109                    n_vocab_tgt, llama_vocab_n_tokens(vocab_dft), vocab_diff, SPEC_VOCAB_MAX_SIZE_DIFFERENCE);110            return false;111        }112 113        for (int i = SPEC_VOCAB_CHECK_START_TOKEN_ID; i < std::min(n_vocab_tgt, n_vocab_dft); ++i) {114            const char * token_text_tgt = llama_vocab_get_text(vocab_tgt, i);115            const char * token_text_dft = llama_vocab_get_text(vocab_dft, i);116 117            if (std::strcmp(token_text_tgt, token_text_dft) != 0) {118                SPC_DBG("draft model vocab must match target model to use speculation but "119                        "token %d content differs - target '%s', draft '%s'\n", i,120                        common_token_to_piece(vocab_tgt, i).c_str(),121                        common_token_to_piece(vocab_dft, i).c_str());122                return false;123            }124        }125    }126 127    return true;128}129 130using common_speculative_draft_params_vec = std::vector<common_speculative_draft_params>;131 132// state of an implementation of speculative decoding133//134// each implementation has a unique type and a state that is implementation-specific135// in a subclass of common_speculative_impl136struct common_speculative_impl {137    const common_speculative_type type;138 139    uint32_t n_seq;140 141    size_t n_call_begin  = 0; // number of times this implementation was called for refresh.142    size_t n_call_draft  = 0; // number of times this implementation was called for generation.143    size_t n_call_accept = 0; // number of times this implementation was called for accumulation.144 145    size_t n_gen_drafts = 0; // number of times a draft or part was generated by this implementation.146    size_t n_acc_drafts = 0; // number of times a draft or part was accepted by the target model.147    size_t n_gen_tokens = 0; // number of tokens generated by this implementation.148    size_t n_acc_tokens = 0; // number of tokens accepted by the target model.149 150    std::vector<size_t> n_acc_tokens_per_pos; // number of tokens accepted per draft position.151 152    // TODO: track performance of most recent calls153    const bool gen_perf = true; // whether to generate performance stats.154 155    int64_t t_begin_us  = 0; // total time spent in refresh of this implementation in microseconds.156    int64_t t_draft_us  = 0; // total time spent in generating drafts in this implementation in microseconds.157    int64_t t_accept_us = 0; // total time spent in accumulation of this implementation in microseconds.158 159    common_speculative_impl(common_speculative_type type, uint32_t n_seq) : type(type), n_seq(n_seq) {}160 161    virtual ~common_speculative_impl() = default;162 163    virtual void begin(llama_seq_id seq_id, const llama_tokens & prompt) = 0;164 165    virtual bool process(const llama_batch & batch) = 0;166 167    virtual void draft(common_speculative_draft_params_vec & dparams) = 0;168 169    virtual void accept(llama_seq_id seq_id, uint16_t n_accepted, bool is_other) = 0;170 171    // (optional) serialize/restore per-seq internal state (e.g. eagle3's deferred boundary).172    virtual bool get_state(llama_seq_id /*seq_id*/, std::vector<uint8_t> & /*data*/) const { return false; }173    virtual void set_state(llama_seq_id /*seq_id*/, const std::vector<uint8_t> & /*data*/) {}174 175    // true if this implementation requires the target context to extract post-norm embeddings176    virtual bool need_embd() const = 0;177 178    // true if this implementation requires the target context to extract pre-norm embeddings179    virtual bool need_embd_nextn() const { return false; }180};181 182struct common_speculative_impl_draft_simple : public common_speculative_impl {183    common_params_speculative_draft params;184 185    llama_batch batch;186 187    std::vector<common_sampler_ptr> smpls;188 189    common_speculative_impl_draft_simple(const common_params_speculative & params, uint32_t n_seq)190        : common_speculative_impl(COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE, n_seq)191        , params(params.draft)192    {193        auto * ctx_dft = this->params.ctx_dft;194        auto * ctx_tgt = this->params.ctx_tgt;195 196        SPC_TRC("%s", "adding speculative implementation 'draft-simple'\n");197        SPC_TRC("- n_max=%d, n_min=%d, p_min=%f\n", this->params.n_max, this->params.n_min, this->params.p_min);198        SPC_TRC("- gpu_layers=%d, cache_k=%s, cache_v=%s, ctx_tgt=%s, ctx_dft=%s, devices=[%s]\n",199                this->params.n_gpu_layers,200                ggml_type_name(this->params.cache_type_k),201                ggml_type_name(this->params.cache_type_v),202                ctx_tgt ? "yes" : "no",203                ctx_dft ? "yes" : "no",204                common_speculative_get_devices_str(this->params.devices).c_str());205 206        batch = llama_batch_init(llama_n_batch(ctx_dft), 0, 1);207 208        // TODO: optimize or pass from outside?209        // {210        //     common_params_sampling params;211        //     params.no_perf = false;212        //213        //     params.top_k = 40;214        //     params.top_p = 0.9;215        //216        //     params.samplers = {217        //         COMMON_SAMPLER_TYPE_TOP_K,218        //         COMMON_SAMPLER_TYPE_TOP_P,219        //         COMMON_SAMPLER_TYPE_INFILL,220        //     };221        //222        //     result->smpl = common_sampler_init(llama_get_model(ctx_dft), params);223        // }224 225        smpls.resize(n_seq);226        for (auto & smpl : smpls) {227            common_params_sampling params;228            params.no_perf = false;229            params.top_k = 10;230            params.samplers = {231                COMMON_SAMPLER_TYPE_TOP_K,232            };233 234            smpl.reset(common_sampler_init(llama_get_model(ctx_dft), params));235        }236 237        const bool vocab_cmpt = common_speculative_are_compatible(llama_get_model(ctx_tgt), llama_get_model(ctx_dft));238        SPC_DBG("vocab_cmpt = %d\n", vocab_cmpt);239 240        if (!vocab_cmpt) {241            SPC_ERR("%s", "the target and draft vocabs are not compatible\n");242 243            throw std::runtime_error("draft model vocab type must match target model to use speculation");244        }245 246        if (n_seq != llama_n_seq_max(ctx_dft)) {247            SPC_ERR("n_seq mismatch: %d != %d\n", n_seq, llama_n_seq_max(ctx_dft));248 249            throw std::runtime_error("the draft model number of sequences is incompatible with the speculative n_seq");250        }251    }252 253    ~common_speculative_impl_draft_simple() override {254        llama_batch_free(batch);255    }256 257    void begin(llama_seq_id /*seq_id*/, const llama_tokens & /*prompt*/) override {258        // noop259    }260 261    bool process(const llama_batch & batch) override {262        auto * ctx_dft = params.ctx_dft;263 264        llama_batch batch_dft = batch;265        batch_dft.logits = nullptr;266 267        const int ret = llama_decode(ctx_dft, batch_dft);268 269        if (ret != 0) {270            SPC_ERR("failed to decode draft batch, ret = %d\n", ret);271 272            return false;273        }274 275        return true;276    }277 278    void draft(common_speculative_draft_params_vec & dparams) override {279        auto & ctx_dft = params.ctx_dft;280 281        common_batch_clear(batch);282 283        // keep track of which sequences are still drafting284        int n_drafting = 0;285        std::vector<bool> drafting(n_seq);286 287        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {288            auto & dp = dparams[seq_id];289 290            if (!dp.drafting) {291                continue;292            }293 294            n_drafting++;295            drafting[seq_id] = true;296            common_sampler_reset(smpls[seq_id].get());297 298            common_batch_add(batch, dp.id_last, dp.n_past, { seq_id }, true);299        }300 301        int ret = llama_decode(ctx_dft, batch);302        if (ret != 0) {303            SPC_ERR("llama_decode returned %d\n", ret);304            return;305        }306 307        int i = 0;308 309        while (n_drafting > 0) {310            int i_batch = 0;311 312            common_batch_clear(batch);313 314            for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {315                if (!drafting[seq_id]) {316                    continue;317                }318 319                auto * smpl = smpls[seq_id].get();320 321                common_sampler_sample(smpl, ctx_dft, i_batch, true);322                ++i_batch;323 324                const auto * cur_p = common_sampler_get_candidates(smpl, true);325 326                for (int k = 0; k < std::min(3, (int) cur_p->size); ++k) {327                    SPC_DBG(" - seq_id %d, draft candidate %3d, pos %3d: %6d (%8.3f) '%s'\n",328                            seq_id, k, i, cur_p->data[k].id, cur_p->data[k].p,329                            common_token_to_piece(ctx_dft, cur_p->data[k].id).c_str());330                }331 332                // add drafted token for each sequence333                const llama_token id = cur_p->data[0].id;334 335                // only collect very high-confidence draft tokens336                if (cur_p->data[0].p < params.p_min) {337                    drafting[seq_id] = false;338                    n_drafting--;339 340                    continue;341                }342 343                common_sampler_accept(smpl, id, true);344 345                auto & dp = dparams.at(seq_id);346                auto & result = *dp.result;347 348                result.push_back(id);349 350                if ((params.n_max <= (int) result.size()) ||351                    (dp.n_max > 0 && dp.n_max <= (int) result.size())) {352                    drafting[seq_id] = false;353                    n_drafting--;354                    continue;355                }356 357                common_batch_add(batch, id, dp.n_past + i + 1, { seq_id }, true);358            }359 360            if (batch.n_tokens == 0) {361                break;362            }363 364            // evaluate the drafted tokens on the draft model365            ret = llama_decode(ctx_dft, batch);366            if (ret != 0) {367                SPC_ERR("llama_decode[%d] returned %d\n", i, ret);368                break;369            }370 371            ++i;372        }373 374        for (auto & dp : dparams) {375            if (!dp.drafting) {376                continue;377            }378 379            if (dp.result->size() < (size_t) params.n_min) {380                dp.result->clear();381            }382        }383    }384 385    void accept(llama_seq_id /*seq_id*/, uint16_t /*n_accepted*/, bool /*is_other*/) override {386        // noop387    }388 389    bool need_embd() const override {390        return false;391    }392};393 394 395// EAGLE3 speculative decoding state396//397// Input of draft decoder: (This is different compared to MTP)398//   At "pos P", the decoder takes input pair (t_{P+1}, g_P), with RoPE at P.399//     - t_{P+1} = token at sequence pos P+1 (the *next* token after P)400//     - g_P     = encoder output = projection of target's extracted hidden states at P401//402// Deferred boundary (MTP doesn't have this issue):403//   Within a single process() call with n_tokens, we can only write decoder KV for404//   training pos 0..n_tokens-2. The last training pos (n_tokens-1) needs t_{n_tokens}405//   which lies *outside* this batch — it is the token target will sample next or the first token from next ubatch.406//   So the last training pos of each process() call is *deferred* to whichever next call has407//   the missing token in hand:408//     - multi-ubatch prefill: the next process()'s first token completes the pair409//                              (handled by the per-seq "cross-ubatch bridge")410//     - single-ubatch prefill / after verify: draft()'s seed step uses "dp.id_last"411//                              (target's freshest sample) to complete the pair412//413// Per-seq carry-over state:414//   pending_g_last    [n_embd_dec]  ┐  the deferred boundary's (g, pos). Set by415//   pending_pos_last  llama_pos     ┘  process() at end of ubatch (= last row);416//                                       rebased by accept() to first-non-accepted pos.417//   verify_g          [N × n_embd_dec] snapshot of process()'s encoder output;418//   verify_pos_first  llama_pos         consumed by accept() to recover the right419//   verify_g_rows     int32_t           pending_g_last row for any n_accepted value.420//421// Performance is overall good but there is waste in verify cycle:422//   process() runs encoder + decoder on the *full* verify batch including rows for423//   rejected drafts. The KV at those positions is then dropped.424//425// TODO: Not sure if we need optimization for this waste?426// If so we may need hybrid stash:427//      in verify mode, have process() only stash features and let draft() seed run428//      encoder+decoder on n_accepted+1 rows).429struct common_speculative_impl_draft_eagle3 : public common_speculative_impl {430    common_params_speculative_draft params;431    llama_batch batch;432 433    std::vector<common_sampler_ptr> smpls;434 435    // backend sampler chain per seq, attached to ctx_dft436    std::vector<llama_sampler *> backend_chains;437 438    int32_t n_embd_dec = 0;       // draft hidden size439    int32_t n_embd_enc = 0;       // target_layer_ids_n * target_hidden_size440    int32_t n_embd_tgt = 0;       // target model hidden size441    int32_t n_layer_tgt = 0;      // target model layer count442 443    const int32_t * target_layer_ids   = nullptr; // model_dft's extract layer indices444    uint32_t        target_layer_ids_n = 0;445 446    // [per-seq] deferred boundary state447    std::vector<std::vector<float>> pending_g_last;448    std::vector<llama_pos>          pending_pos_last;449 450    // [per-seq] snapshot of the most recent process()'s encoder output451    std::vector<std::vector<float>> verify_g;         // [n_seq][n_rows * n_embd_dec]452    std::vector<llama_pos>          verify_pos_first; // [n_seq] — pos of verify_g[seq][0]453    std::vector<int32_t>            verify_g_rows;    // [n_seq] — number of rows454 455    // scratch buffer for concatenated target features [n_tokens, n_embd_enc]456    std::vector<float> features_buf;457    std::vector<float> g_embd_buf;458 459    common_speculative_impl_draft_eagle3(const common_params_speculative & params, uint32_t n_seq)460        : common_speculative_impl(COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3, n_seq)461        , params(params.draft)462    {463        SPC_TRC("%s", "adding speculative implementation 'draft-eagle3'\n");464        SPC_TRC("- n_max=%d, n_min=%d, p_min=%f, backend_sampling=%d\n", params.draft.n_max, params.draft.n_min, params.draft.p_min, (int) params.draft.backend_sampling);465 466        auto * ctx_tgt = this->params.ctx_tgt;467        auto * ctx_dft = this->params.ctx_dft;468        GGML_ASSERT(ctx_tgt && ctx_dft && "EAGLE3 requires ctx_tgt and ctx_dft to be set");469 470        const llama_model * model_dft = llama_get_model(ctx_dft);471        const llama_model * model_tgt = llama_get_model(ctx_tgt);472 473        target_layer_ids   = llama_model_target_layer_ids  (model_dft);474        target_layer_ids_n = llama_model_target_layer_ids_n(model_dft);475        if (target_layer_ids_n != 3) {476            throw std::runtime_error("draft model is not eagle3 (expected 3 extract layers, got " +477                                     std::to_string(target_layer_ids_n) + ")");478        }479 480        n_embd_tgt = llama_model_n_embd(model_tgt);481        n_embd_dec = llama_model_n_embd(model_dft);482        n_embd_enc = (int32_t) target_layer_ids_n * n_embd_tgt;483        n_layer_tgt = llama_model_n_layer(model_tgt);484 485        const int32_t n_b = (int32_t) llama_n_batch(ctx_dft);486        batch = llama_batch_init(/*n_tokens=*/ n_b, /*embd=*/ n_embd_dec, /*n_seq_max=*/ 1);487        // llama_batch_init allocates only one of token/embd; eagle3 decoder needs both.488        // TODO: fix, how to call without malloc489        batch.token = (llama_token *) malloc(sizeof(llama_token) * n_b);490 491        smpls.resize(n_seq);492        for (auto & s : smpls) {493            common_params_sampling sparams;494            sparams.no_perf  = false;495            sparams.top_k    = 10;496            sparams.samplers = { COMMON_SAMPLER_TYPE_TOP_K };497            s.reset(common_sampler_init(llama_get_model(ctx_dft), sparams));498        }499 500        // offload draft sampling to the backend501        backend_chains.assign(n_seq, nullptr);502        if (this->params.backend_sampling) {503            for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {504                llama_sampler * chain = llama_sampler_chain_init(llama_sampler_chain_default_params());505                llama_sampler_chain_add(chain, llama_sampler_init_top_k(10));506 507                if (!llama_set_sampler(ctx_dft, seq_id, chain)) {508                    SPC_WRN("backend offload failed for seq_id=%d; using CPU sampler\n", (int) seq_id);509                    llama_sampler_free(chain);510                    chain = nullptr;511                }512                backend_chains[seq_id] = chain;513            }514        }515 516        // turn on extraction of the target layers' hidden states517        for (uint32_t k = 0; k < target_layer_ids_n; ++k) {518            if (target_layer_ids[k] < n_layer_tgt) {519                llama_set_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k], true);520            } else if (target_layer_ids[k] == n_layer_tgt) {521                llama_set_embeddings_nextn(ctx_tgt, true, /*masked*/ false);522            } else {523                GGML_ABORT("EAGLE3: target layer id %d exceeds target n_layer %d", target_layer_ids[k], n_layer_tgt);524            }525        }526 527        // turn on extraction of the draft model's pre-norm hidden state528        // (used both for the encoder output g_embd and the decoder pre-norm output).529        llama_set_embeddings_nextn(ctx_dft, true, /*masked*/ true);530 531        pending_g_last.assign(n_seq, std::vector<float>(n_embd_dec, 0.0f));532        pending_pos_last.assign(n_seq, -1);533 534        verify_g.assign(n_seq, std::vector<float>());535        verify_pos_first.assign(n_seq, -1);536        verify_g_rows.assign(n_seq, 0);537    }538 539    ~common_speculative_impl_draft_eagle3() override {540        auto * ctx_dft = this->params.ctx_dft;541        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) backend_chains.size(); ++seq_id) {542            if (backend_chains[seq_id] == nullptr) {543                continue;544            }545            if (ctx_dft) {546                llama_set_sampler(ctx_dft, seq_id, nullptr);547            }548            llama_sampler_free(backend_chains[seq_id]);549        }550        backend_chains.clear();551 552        if (batch.token != nullptr) {553            free(batch.token);554            batch.token = nullptr;555        }556        llama_batch_free(batch);557    }558 559    void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {560        const int32_t N = (int32_t) prompt.size();561        if (N <= 0) {562            return;563        }564        // expected state after prefill: ctx_dft has pos 0..N-2 (last position is deferred to565        // draft()'s seed step). Warn only if more than one position is missing.566        auto * ctx_dft = this->params.ctx_dft;567        const llama_pos pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id);568        if (pos_max < N - 2) {569            SPC_WRN("ctx_dft pos_max=%d < N-2=%d — process() did not run on every prefill ubatch. "570                    "Drafts may degrade.\n",571                    (int) pos_max, N - 2);572        }573    }574 575    bool process(const llama_batch & batch_in) override {576        if (batch_in.n_tokens <= 0) {577            return true;578        }579 580        if (batch_in.token == nullptr || batch_in.embd != nullptr) {581            return true;582        }583 584        const int32_t n_tokens = batch_in.n_tokens;585 586        // i_batch_beg[seq] / i_batch_end[seq]: inclusive batch indices of this seq's587        // first/last token in batch_in. Assumes per-seq tokens are contiguous within588        // the ubatch (server's default ordering).589        std::vector<int32_t> i_batch_beg(n_seq, -1);590        std::vector<int32_t> i_batch_end(n_seq, -1);591        for (int k = 0; k < n_tokens; ++k) {592            GGML_ASSERT(batch_in.n_seq_id[k] == 1);593            const llama_seq_id seq_id = batch_in.seq_id[k][0];594            if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {595                continue;596            }597            i_batch_end[seq_id] = k;598            if (i_batch_beg[seq_id] < 0) {599                i_batch_beg[seq_id] = k;600            }601        }602 603        auto * ctx_tgt = this->params.ctx_tgt;604        auto * ctx_dft = this->params.ctx_dft;605 606        // Interleave each extract_layer's hidden state into a contiguous buffer of607        // shape [n_tokens, target_layer_ids_n * n_embd_tgt]. Then run EAGLE3 encoder608        // to get one g_embd row per token.609        features_buf.resize((size_t) n_tokens * n_embd_enc, 0.0f);610 611        for (uint32_t k = 0; k < target_layer_ids_n; ++k) {612            const float * layer = target_layer_ids[k] < n_layer_tgt613                ? llama_get_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k])614                : llama_get_embeddings_nextn(ctx_tgt);615            if (!layer) {616                GGML_ABORT("EAGLE3: target layer %d input not extracted.", target_layer_ids[k]);617            }618            for (int32_t i = 0; i < n_tokens; ++i) {619                float * dst = features_buf.data() + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;620                const float * src = layer + (size_t) i * n_embd_tgt;621                std::memcpy(dst, src, (size_t) n_embd_tgt * sizeof(float));622            }623        }624 625        g_embd_buf.resize((size_t) n_tokens * n_embd_dec);626 627        // llama_encode() requires the full encoder batch to fit in n_ubatch.628        // Allow batch > ubatch: eagle3's per-token encoder can be chunked safely.629        const int32_t n_ubatch_dft = (int32_t) llama_n_ubatch(ctx_dft);630        for (int32_t i = 0; i < n_tokens; i += n_ubatch_dft) {631            const int32_t n_chunk = std::min(n_ubatch_dft, n_tokens - i);632 633            llama_batch enc_batch = {634                /*.n_tokens =*/ n_chunk,635                /*.token    =*/ nullptr,636                /*.embd     =*/ features_buf.data() + (size_t) i * n_embd_enc,637                /*.pos      =*/ nullptr,638                /*.n_seq_id =*/ nullptr,639                /*.seq_id   =*/ nullptr,640                /*.logits   =*/ nullptr,641            };642            const int32_t rc = llama_encode(ctx_dft, enc_batch);643            if (rc != 0) {644                SPC_ERR("llama_encode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",645                        rc, (int) n_chunk, (int) i);646                return false;647            }648 649            // g_embd has shape [n_chunk, n_embd_dec] in ctx_dft's pre-norm embeddings buffer.650            const float * g_embd_chunk = llama_get_embeddings_nextn(ctx_dft);651            GGML_ASSERT(g_embd_chunk && "EAGLE3 encoder produced no output.");652            std::memcpy(g_embd_buf.data() + (size_t) i * n_embd_dec,653                        g_embd_chunk,654                        (size_t) n_chunk * n_embd_dec * sizeof(float));655        }656 657        const float * g_embd = g_embd_buf.data();658 659        const size_t row_bytes = (size_t) n_embd_dec * sizeof(float);660 661        // EAGLE3 decoder input convention: at memory pos P the input pair is662        // (token[P+1], g_embd[P]). This shifts the token index "left by one" relative to g_embd.663        //664        // Per seq, in order:665        //   (a) cross-ubatch bridge — when applicable, write the previously-deferred666        //       pos using this ubatch's first token + pending_g_last.667        //   (b) main write loop — for k in [beg, end-1], write (token[k+1], g_embd[k])668        //       at pos[k]. The last training pos (k=end) is left unwritten = new669        //       deferred boundary, completed by the next process() or draft() call.670        //   (c) refresh deferred state — stash this ubatch's full g_embd into verify_g,671        //       update pending_g_last / pending_pos_last to the last row.672        common_batch_clear(batch);673 674        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {675            const int32_t beg = i_batch_beg[seq_id];676            const int32_t end = i_batch_end[seq_id];677            if (beg < 0 || end < 0) {678                continue;679            }680 681            // cross-ubatch bridge — complete the prior ubatch's deferred boundary.682            // Fires iff all three preconditions hold:683            //   1) pending_pos_last >= 0684            //   2) pending_pos_last + 1 == pos[beg]685            //   3) pending_pos_last > dft_pos_max // TODO: is this check needed?686            const llama_pos pending_pos = pending_pos_last[seq_id];687            if (pending_pos >= 0 && pending_pos + 1 == batch_in.pos[beg]) {688                const llama_pos dft_pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id);689                if (pending_pos > dft_pos_max) {690                    common_batch_add(batch, batch_in.token[beg], pending_pos, { seq_id }, /*logits=*/ false);691                    std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,692                                pending_g_last[seq_id].data(), row_bytes);693                }694            }695 696            for (int32_t k = beg; k < end; ++k) {697                common_batch_add(batch, batch_in.token[k + 1], batch_in.pos[k], { seq_id }, /*logits=*/ false);698                std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,699                            g_embd + (size_t) k * n_embd_dec, row_bytes);700            }701 702            // refresh deferred state703            const int32_t n_rows = end - beg + 1;704            verify_pos_first[seq_id] = batch_in.pos[beg];705            pending_pos_last[seq_id] = batch_in.pos[end];706            verify_g_rows[seq_id]    = n_rows;707            verify_g[seq_id].resize((size_t) n_rows * n_embd_dec, 0.0f);708            std::memcpy(verify_g[seq_id].data(),       g_embd + (size_t) beg * n_embd_dec, row_bytes * n_rows);709            std::memcpy(pending_g_last[seq_id].data(), g_embd + (size_t) end * n_embd_dec, row_bytes);710        }711 712        if (batch.n_tokens > 0) {713            const int32_t rc = llama_decode(ctx_dft, batch);714            if (rc != 0) {715                SPC_ERR("llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, ubatch_pos[0]=%d)\n",716                        rc, (int) batch.n_tokens, (int) batch_in.pos[0]);717                return false;718            }719        }720 721        return true;722    }723 724    void draft(common_speculative_draft_params_vec & dparams) override {725        auto & ctx_dft = params.ctx_dft;726 727        common_batch_clear(batch);728 729        // keep track of which sequences are still drafting730        int n_drafting = 0;731        std::vector<bool> drafting(n_seq);732 733        const size_t row_bytes = (size_t) n_embd_dec * sizeof(float);734 735        // Complete the deferred boundary pair (dp.id_last, pending_g_last) at memory736        // pos pending_pos_last. dp.id_last is target's freshest sample (= corrected737        // token after verify, or first generated token after prefill), matching the738        // EAGLE3 input convention (token[P+1], g_embd[P]) at pos P.739        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {740            auto & dp = dparams[seq_id];741 742            if (!dp.drafting) {743                continue;744            }745            if (pending_pos_last[seq_id] < 0) {746                continue;747            }748 749            n_drafting++;750            drafting[seq_id] = true;751            common_sampler_reset(smpls[seq_id].get());752 753            llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, pending_pos_last[seq_id], -1);754 755            common_batch_add(batch, dp.id_last, pending_pos_last[seq_id], { seq_id }, true);756            std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec,757                        pending_g_last[seq_id].data(),758                        row_bytes);759        }760 761        if (batch.n_tokens == 0) {762            return;763        }764 765        int ret = llama_decode(ctx_dft, batch);766        if (ret != 0) {767            SPC_ERR("llama_decode returned %d\n", ret);768            return;769        }770 771        int i = 0;772 773        while (n_drafting > 0) {774            int i_batch = 0;775 776            common_batch_clear(batch);777 778            for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {779                if (!drafting[seq_id]) {780                    continue;781                }782 783                auto * smpl = smpls[seq_id].get();784 785                common_sampler_sample(smpl, ctx_dft, i_batch, true);786                // pre-norm hidden state of this position becomes g_embd for the next step787                const float * prenorm = llama_get_embeddings_nextn_ith(ctx_dft, i_batch);788                ++i_batch;789 790                const auto * cur_p = common_sampler_get_candidates(smpl, true);791 792                for (int k = 0; k < std::min(3, (int) cur_p->size); ++k) {793                    SPC_DBG(" - seq_id %d, draft candidate %3d, pos %3d: %6d (%8.3f) '%s'\n",794                            seq_id, k, i, cur_p->data[k].id, cur_p->data[k].p,795                            common_token_to_piece(ctx_dft, cur_p->data[k].id).c_str());796                }797 798                const llama_token id = cur_p->data[0].id;799 800                // only collect very high-confidence draft tokens801                // (configurable via --spec-draft-p-min, set to 0.0 to disable early-stop)802                if (cur_p->data[0].p < params.p_min) {803                    drafting[seq_id] = false;804                    n_drafting--;805 806                    continue;807                }808 809                common_sampler_accept(smpl, id, true);810 811                auto & dp = dparams.at(seq_id);812                auto & result = *dp.result;813 814                result.push_back(id);815 816                if (params.n_max <= (int) result.size()) {817                    drafting[seq_id] = false;818                    n_drafting--;819                    continue;820                }821 822                common_batch_add(batch, id, pending_pos_last[seq_id] + (i + 1), { seq_id }, true);823                std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd_dec, prenorm, row_bytes);824            }825 826            if (batch.n_tokens == 0) {827                break;828            }829 830            ret = llama_decode(ctx_dft, batch);831            if (ret != 0) {832                SPC_ERR("llama_decode[%d] returned %d\n", i, ret);833                break;834            }835 836            ++i;837        }838 839        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {840            auto & dp = dparams[seq_id];841            if (!dp.drafting) {842                continue;843            }844 845            if (dp.result->size() < (size_t) params.n_min) {846                dp.result->clear();847            }848        }849    }850 851    void accept(llama_seq_id seq_id, uint16_t n_accepted, bool /*is_other*/) override {852        if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {853            return;854        }855 856        const int32_t n_rows = verify_g_rows[seq_id];857        if (n_rows <= 0) {858            return;859        }860 861        const int32_t i_g = std::min<int32_t>(n_accepted, n_rows - 1);862        pending_pos_last[seq_id] = verify_pos_first[seq_id] + i_g;863        std::memcpy(pending_g_last[seq_id].data(),864                    verify_g[seq_id].data() + (size_t) i_g * n_embd_dec,865                    (size_t) n_embd_dec * sizeof(float));866    }867 868    // we only need to stash the deferred boundary's g_embd row for recurrent/hybrid targets:869    // their single-position checkpoints drop it on restore870    bool need_boundary_stash() const {871        const llama_model * model_tgt = llama_get_model(params.ctx_tgt);872        return llama_model_is_recurrent(model_tgt) || llama_model_is_hybrid(model_tgt);873    }874 875    bool get_state(llama_seq_id seq_id, std::vector<uint8_t> & data) const override {876        if (!need_boundary_stash()) {877            return false;878        }879        if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq || pending_pos_last[seq_id] < 0) {880            return false;881        }882 883        const llama_pos          pos = pending_pos_last[seq_id];884        const std::vector<float> & g = pending_g_last[seq_id];885 886        data.resize(sizeof(llama_pos) + g.size() * sizeof(float));887        std::memcpy(data.data(),                     &pos,     sizeof(llama_pos));888        std::memcpy(data.data() + sizeof(llama_pos), g.data(), g.size() * sizeof(float));889        return true;890    }891 892    void set_state(llama_seq_id seq_id, const std::vector<uint8_t> & data) override {893        if (!need_boundary_stash()) {894            return;895        }896        if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {897            return;898        }899        if (data.size() != sizeof(llama_pos) + (size_t) n_embd_dec * sizeof(float)) {900            return;901        }902 903        llama_pos pos = -1;904        std::memcpy(&pos, data.data(), sizeof(llama_pos));905 906        pending_pos_last[seq_id] = pos;907        pending_g_last[seq_id].resize(n_embd_dec);908        std::memcpy(pending_g_last[seq_id].data(), data.data() + sizeof(llama_pos), (size_t) n_embd_dec * sizeof(float));909    }910 911    bool need_embd() const override {912        return false;913    }914};915 916// DFlash: block-diffusion drafting with a draft-side KV cache injection917struct common_speculative_impl_draft_dflash : public common_speculative_impl {918    common_params_speculative_draft params;919 920    llama_batch batch;        // noise tokens921    llama_batch batch_inject; // target features for KV cache injection922 923    std::vector<common_sampler_ptr> smpls;924 925    int32_t n_embd_dec = 0;  // draft hidden size926    int32_t n_embd_enc = 0;  // target_layer_ids_n * target_hidden_size927    int32_t n_embd_tgt = 0;  // target model hidden size928 929    int32_t     block_size    = 0;930    llama_token mask_token_id = 0;931 932    // draft-dspark: the draft carries a Markov head and uses an anchor-first block layout933    const bool is_dspark;934 935    const int32_t * target_layer_ids   = nullptr; // model_dft's extract layer indices936    uint32_t        target_layer_ids_n = 0;937 938    // scratch buffer for concatenated target features [n_tokens, n_embd_enc]939    std::vector<float> features_buf;940 941    common_speculative_impl_draft_dflash(const common_params_speculative & params, uint32_t n_seq,942            common_speculative_type type = COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH)943        : common_speculative_impl(type, n_seq)944        , params(params.draft)945        , is_dspark(type == COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK)946    {947        auto * ctx_tgt = this->params.ctx_tgt;948        auto * ctx_dft = this->params.ctx_dft;949        GGML_ASSERT(ctx_tgt && ctx_dft && "DFlash requires ctx_tgt and ctx_dft to be set");950 951        const llama_model * model_dft = llama_get_model(ctx_dft);952        const llama_model * model_tgt = llama_get_model(ctx_tgt);953 954        target_layer_ids   = llama_model_target_layer_ids  (model_dft);955        target_layer_ids_n = llama_model_target_layer_ids_n(model_dft);956        GGML_ASSERT(target_layer_ids_n > 0 && "DFlash model has no target_layer_ids");957 958        n_embd_tgt    = llama_model_n_embd(model_tgt);959        n_embd_dec    = llama_model_n_embd(model_dft);960        n_embd_enc    = (int32_t) target_layer_ids_n * n_embd_tgt;961 962        // read the trained block size from the dflash.block_size metadata key963        block_size = 16;964        {965            char buf[32] = {};966            if (llama_model_meta_val_str(model_dft, "dflash.block_size", buf, sizeof(buf)) >= 0) {967                block_size = std::atoi(buf);968            }969        }970        mask_token_id = llama_vocab_mask(llama_model_get_vocab(model_dft));971 972        LOG_INF("%s: adding speculative implementation '%s'\n", __func__, common_speculative_type_to_str(type).c_str());973        LOG_INF("%s: - n_max=%d, n_min=%d, p_min=%.2f\n", __func__, this->params.n_max, this->params.n_min, this->params.p_min);974        LOG_INF("%s: - block_size=%d, mask_token_id=%d, n_extract=%u\n", __func__, block_size, mask_token_id, target_layer_ids_n);975 976        // DFlash input is [id_last, <mask> * (block_size-1)]: in-place denoising yields at most977        // block_size-1 draft tokens, DSpark yield a full block_size draft tokens978        const int32_t n_draft_max = is_dspark ? block_size : block_size - 1;979        if (this->params.n_max > n_draft_max || this->params.n_min > n_draft_max) {980            LOG_WRN("%s: requested draft size (n_max=%d, n_min=%d) exceeds the trained block size %d -- clamping to %d\n",981                    __func__, this->params.n_max, this->params.n_min, block_size, n_draft_max);982            this->params.n_max = std::min(this->params.n_max, n_draft_max);983            this->params.n_min = std::min(this->params.n_min, n_draft_max);984        }985 986        batch        = llama_batch_init(llama_n_batch(ctx_dft), 0,          n_seq);987        batch_inject = llama_batch_init(llama_n_batch(ctx_dft), n_embd_dec, n_seq);988 989        smpls.resize(n_seq);990        for (auto & s : smpls) {991            common_params_sampling sparams;992            sparams.no_perf  = false;993            sparams.top_k    = 10;994            sparams.samplers = { COMMON_SAMPLER_TYPE_TOP_K };995            s.reset(common_sampler_init(model_dft, sparams));996        }997 998        // turn on extraction of the target layers' input embeddings999        for (uint32_t k = 0; k < target_layer_ids_n; ++k) {1000            llama_set_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k], true);1001        }1002 1003        llama_set_embeddings_nextn(ctx_dft, true, /*masked*/ true);1004        llama_set_causal_attn(ctx_dft, false); // DFlash needs non-causal attention1005    }1006 1007    ~common_speculative_impl_draft_dflash() override {1008        llama_batch_free(batch);1009        llama_batch_free(batch_inject);1010    }1011 1012    void begin(llama_seq_id seq_id, const llama_tokens & prompt) override {1013        if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {1014            return;1015        }1016 1017        const int32_t N = (int32_t) prompt.size();1018        if (N <= 0) {1019            return;1020        }1021 1022        const llama_pos pos_max = llama_memory_seq_pos_max(llama_get_memory(params.ctx_dft), seq_id);1023        if (pos_max < N - 1) {1024            LOG_WRN("%s: ctx_dft pos_max=%d < N-1=%d - process() did not run on every prefill ubatch. "1025                    "Drafts may degrade.\n",1026                    __func__, (int) pos_max, N - 1);1027        }1028    }1029 1030    bool process(const llama_batch & batch_in) override {1031        if (batch_in.n_tokens <= 0) {1032            return true;1033        }1034 1035        // Target prefill may contain token IDs or multimodal embeddings. Both1036        // produce the target-layer features used to seed the draft KV cache, so1037        // skipping the embedding batches leaves a hole in the draft's cache and1038        // the next injection fails to initialize.1039        // TODO: revisit after https://github.com/ggml-org/llama.cpp/pull/24669 is merged1040        const bool has_tokens     = batch_in.token != nullptr;1041        const bool has_embeddings = batch_in.embd  != nullptr;1042        if (has_tokens == has_embeddings) {1043            return true;1044        }1045 1046        const int32_t n_tokens = batch_in.n_tokens;1047 1048        // per-seq inclusive batch range (assumes each seq's tokens are contiguous in the batch)1049        std::vector<int32_t> i_batch_beg(n_seq, -1);1050        std::vector<int32_t> i_batch_end(n_seq, -1);1051        for (int32_t k = 0; k < n_tokens; ++k) {1052            GGML_ASSERT(batch_in.n_seq_id[k] == 1);1053            const llama_seq_id seq_id = batch_in.seq_id[k][0];1054            if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) {1055                continue;1056            }1057            i_batch_end[seq_id] = k;1058            if (i_batch_beg[seq_id] < 0) {1059                i_batch_beg[seq_id] = k;1060            }1061        }1062 1063        auto * ctx_tgt = this->params.ctx_tgt;1064        auto * ctx_dft = this->params.ctx_dft;1065 1066        const int32_t n_ubatch = (int32_t) llama_n_ubatch(ctx_dft);1067 1068        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {1069            if (i_batch_beg[seq_id] < 0) {1070                continue;1071            }1072            const int32_t n_rows = i_batch_end[seq_id] - i_batch_beg[seq_id] + 1;1073 1074            for (int32_t offset = 0; offset < n_rows; offset += n_ubatch) {1075                const int32_t n_chunk = std::min(n_ubatch, n_rows - offset);1076 1077                // gather this chunk's target features, interleaved by extract layer1078                features_buf.resize((size_t) n_chunk * n_embd_enc);1079                for (uint32_t k = 0; k < target_layer_ids_n; ++k) {1080                    const float * layer = llama_get_embeddings_layer_inp(ctx_tgt, (uint32_t) target_layer_ids[k]);1081                    if (!layer) {1082                        GGML_ABORT("DFlash: target layer %d input not extracted.", target_layer_ids[k]);1083                    }1084                    for (int32_t i = 0; i < n_chunk; ++i) {1085                        float       * dst = features_buf.data() + (size_t) i * n_embd_enc + k * (size_t) n_embd_tgt;1086                        const float * src = layer + (size_t) (i_batch_beg[seq_id] + offset + i) * n_embd_tgt;1087                        std::memcpy(dst, src, (size_t) n_embd_tgt * sizeof(float));1088                    }1089                }1090 1091                // fuse extracted features through DFlash encoder1092                llama_batch enc_batch = {1093                    /*.n_tokens =*/ n_chunk,1094                    /*.token    =*/ nullptr,1095                    /*.embd     =*/ features_buf.data(),1096                    /*.pos      =*/ nullptr,1097                    /*.n_seq_id =*/ nullptr,1098                    /*.seq_id   =*/ nullptr,1099                    /*.logits   =*/ nullptr,1100                };1101 1102                int32_t rc = llama_encode(ctx_dft, enc_batch);1103                if (rc != 0) {1104                    LOG_ERR("%s: llama_encode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",1105                            __func__, rc, (int) n_chunk, (int) offset);1106                    return false;1107                }1108 1109                const float * inp_g = llama_get_embeddings_nextn(ctx_dft);1110                GGML_ASSERT(inp_g && "DFlash encoder produced no output.");1111 1112                // inject the DFlash decoder K/V cache at the tokens' target positions1113                batch_inject.n_tokens = n_chunk;1114                std::memcpy(batch_inject.embd, inp_g, (size_t) n_chunk * n_embd_dec * sizeof(float));1115 1116                for (int32_t i = 0; i < n_chunk; ++i) {1117                    batch_inject.pos[i]       = batch_in.pos[i_batch_beg[seq_id] + offset + i];1118                    batch_inject.n_seq_id[i]  = 1;1119                    batch_inject.seq_id[i][0] = seq_id;1120                    batch_inject.logits[i]    = false;1121                }1122                rc = llama_decode(ctx_dft, batch_inject);1123                if (rc != 0) {1124                    LOG_ERR("%s: llama_decode(ctx_dft) failed rc=%d (n_tokens=%d, offset=%d)\n",1125                            __func__, rc, (int) n_chunk, (int) offset);1126                    return false;1127                }1128            }1129        }1130 1131        return true;1132    }1133 1134    void draft(common_speculative_draft_params_vec & dparams) override {1135        auto & ctx_dft = params.ctx_dft;1136 1137        common_batch_clear(batch);1138 1139        // build one batch holding every drafting sequence's noise block into a single decode)1140        // record where each block starts and its size1141        std::vector<int32_t> i_block_beg(n_seq, -1);1142        std::vector<int32_t> n_block    (n_seq,  0);1143 1144        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {1145            auto & dp = dparams[seq_id];1146            if (!dp.drafting) {1147                continue;1148            }1149 1150            common_sampler_reset(smpls[seq_id].get());1151 1152            const int32_t n = (int32_t) dp.n_past;1153 1154            const int32_t n_draft = params.n_max;1155 1156            const int32_t n_block_tokens = n_draft + (is_dspark ? 0 : 1);1157            i_block_beg[seq_id] = batch.n_tokens;1158            n_block    [seq_id] = n_block_tokens;1159            for (int32_t i = 0; i < n_block_tokens; ++i) {1160                common_batch_add(batch, i == 0 ? dp.id_last : mask_token_id, n + i, { seq_id }, true);1161            }1162        }1163 1164        if (batch.n_tokens == 0) {1165            return;1166        }1167 1168        // decode all sequence's noise block in a single batch1169        int ret = llama_decode(ctx_dft, batch);1170        if (ret != 0) {1171            LOG_WRN("%s: llama_decode returned %d\n", __func__, ret);1172            return;1173        }1174 1175        for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) {1176            if (i_block_beg[seq_id] < 0) {1177                continue;1178            }1179            auto & dp = dparams[seq_id];1180 1181            const int32_t beg            = i_block_beg[seq_id];1182            const int32_t n_block_tokens = n_block[seq_id];1183 1184            auto * smpl = smpls[seq_id].get();1185 1186            auto & result = *dp.result;1187 1188            if (is_dspark) {1189                // DSpark predicts the next token from position 0 and optionally truncates1190                // at the first position below the confidence threshold.1191                const float * conf = params.p_min > 0.0f ? llama_get_embeddings_nextn(ctx_dft) : nullptr;1192 1193                for (int32_t i = 0; i < n_block_tokens; ++i) {1194                    const int32_t idx = beg + i;1195 1196                    if (conf && conf[(size_t) idx * n_embd_dec] < params.p_min) {1197                        break;1198                    }1199 1200                    common_sampler_sample(smpl, ctx_dft, idx, true);

Showing the first 1,200 of 2775 lines. Download the file for the rest.

Brunobkr/llama.cpp_AlgMor24_github · Team Ai