Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
llama-kv-cache.h465 linesDownload Raw Back to src
1#pragma once2 3#include "llama-batch.h"4#include "llama-graph.h"5#include "llama-kv-cells.h"6#include "llama-memory.h"7 8#include <unordered_map>9#include <vector>10 11struct llama_cparams;12struct llama_hparams;13struct llama_model;14struct llama_context;15 16//17// llama_kv_cache18//19 20class llama_kv_cache : public llama_memory_i {21public:22    struct stream_copy_info {23        bool empty() const {24            assert(ssrc.size() == sdst.size());25            return ssrc.empty();26        }27 28        std::vector<uint32_t> ssrc;29        std::vector<uint32_t> sdst;30    };31 32    // for each ubatch, create a slot_info that contains information about where the ubatch should be inserted in the33    //   KV cells. for example, cell indices for each token, such that: token[i] -> goes to cells[idxs[i]]34    struct slot_info {35        // data for ggml_set_rows36        using idx_vec_t = std::vector<uint32_t>;37 38        // number of streams: ns = s1 - s0 + 139        uint32_t s0;40        uint32_t s1;41 42        std::vector<llama_seq_id> strm; // [ns]43        std::vector<idx_vec_t>    idxs; // [ns]44 45        uint32_t head() const {46            GGML_ASSERT(idxs.size() == 1);47            GGML_ASSERT(!idxs[0].empty());48 49            return idxs[0][0];50        }51 52        void resize(size_t n) {53            strm.resize(n);54            idxs.resize(n);55        }56 57        size_t size() const {58            GGML_ASSERT(idxs.size() == strm.size());59            GGML_ASSERT(!idxs.empty());60 61            return idxs[0].size();62        }63 64        size_t n_stream() const {65            return strm.size();66        }67 68        bool empty() const {69            return idxs.empty();70        }71 72        void clear() {73            idxs.clear();74        }75 76        // check if indices are contiguous starting from head()77        bool is_contiguous() const {78            if (idxs.empty() || idxs[0].empty()) {79                return true;80            }81            if (idxs.size() > 1) {82                return false;83            }84            const uint32_t h = idxs[0][0];85            for (size_t i = 0; i < idxs[0].size(); ++i) {86                if (idxs[0][i] != h + i) {87                    return false;88                }89            }90            return true;91        }92    };93 94    using slot_info_vec_t = std::vector<slot_info>;95 96    // TODO: refactor the memory instances to not depend on `llama_model`97    //       instead pass all necessary info (e.g. hparams, dev layers, arch, etc.) directly98    //       likely through `struct llama_memory_params`99    llama_kv_cache(100            const llama_model & model,101          const llama_hparams & hparams,102                    ggml_type   type_k,103                    ggml_type   type_v,104                         bool   v_trans,105                         bool   offload,106                         bool   unified,107                     uint32_t   kv_size,108                     uint32_t   n_seq_max,109                     uint32_t   n_pad,110                     uint32_t   n_swa,111               llama_swa_type   swa_type,112               llama_memory_t   mem_other,113        const layer_filter_cb & filter,114        const  layer_reuse_cb & reuse,115        const  layer_share_cb & share,116        // a model can hold more than one cache, so the tensor names have to stay unique117                 const char *   name_tag = "");118 119    ~llama_kv_cache() = default;120 121    //122    // llama_memory_i123    //124 125    llama_memory_context_ptr init_batch(126            llama_batch_allocr & balloc,127            uint32_t n_ubatch,128            bool embd_all) override;129 130    llama_memory_context_ptr init_full() override;131 132    llama_memory_context_ptr init_update(llama_context * lctx, bool optimize) override;133 134    bool get_can_shift() const override;135 136    void clear(bool data) override;137 138    bool seq_rm  (llama_seq_id seq_id,                              llama_pos p0, llama_pos p1) override;139    void seq_cp  (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) override;140    void seq_keep(llama_seq_id seq_id)                                                          override;141    void seq_add (llama_seq_id seq_id,                              llama_pos p0, llama_pos p1, llama_pos shift) override;142    void seq_div (llama_seq_id seq_id,                              llama_pos p0, llama_pos p1, int d) override;143 144    llama_pos seq_pos_min(llama_seq_id seq_id) const override;145    llama_pos seq_pos_max(llama_seq_id seq_id) const override;146 147    std::map<ggml_backend_buffer_type_t, size_t> memory_breakdown() const override;148 149    // state write/load150 151    void state_write(llama_io_write_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) const override;152    void state_read (llama_io_read_i  & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) override;153 154    //155    // llama_kv_cache specific API156    //157 158    uint32_t get_size()     const;159    uint32_t get_n_stream() const;160 161    bool get_has_shift() const;162 163    ggml_type type_k() const;164    ggml_type type_v() const;165 166    std::vector<uint32_t> get_layer_ids() const;167    ggml_tensor * get_k_storage(int32_t il) const;168 169    const llama_kv_cells & get_cells(llama_seq_id seq_id) const;170 171    // state_read, plus the cells the restored tokens were placed in172    // a cache that mirrors another one (the qwen4exp indexer) must not search for its own cells: two searches agree only by luck173    //   sinfos_out: if set, filled with the layout used; a stream with no cells leaves an empty entry174    //   sinfos_in : if set, the layout to use instead of searching. one entry per stream, cell count must match the blob175    void state_read_sinfo(176            llama_io_read_i & io,177               llama_seq_id   seq_id,178      llama_state_seq_flags   flags,179          slot_info_vec_t *   sinfos_out,180    const slot_info_vec_t *   sinfos_in);181 182    //183    // graph_build API184    //185 186    uint32_t get_n_kv(const slot_info & sinfo) const;187 188    // get views of the current state of the cache189    ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;190    ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;191 192    // store k_cur and v_cur in the cache based on the provided head location193    ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const;194    ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const;195 196    //197    // preparation API198    //199 200    // find places for the provided ubatches in the cache, returns the slot infos201    // return empty vector on failure202    slot_info_vec_t prepare(const std::vector<llama_ubatch> & ubatches);203 204    bool update(llama_context * lctx, bool do_shift, const stream_copy_info & sc_info);205 206    // find a slot of kv cells that can hold the ubatch207    // if cont == true, then the slot must be continuous208    // return empty slot_info on failure209    slot_info find_slot(const llama_ubatch & ubatch, bool cont) const;210 211    // emplace the ubatch context into slot: [sinfo.idxs[0...ubatch.n_tokens - 1]]212    void apply_ubatch(const slot_info & sinfo, const llama_ubatch & ubatch);213 214    //215    // input API216    //217 218    ggml_tensor * build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const;219    ggml_tensor * build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const;220 221    ggml_tensor * build_input_k_rot(ggml_context * ctx) const;222    ggml_tensor * build_input_v_rot(ggml_context * ctx) const;223 224    void set_input_k_idxs(ggml_tensor * dst, const llama_ubatch * ubatch, const slot_info & sinfo) const;225    void set_input_v_idxs(ggml_tensor * dst, const llama_ubatch * ubatch, const slot_info & sinfo) const;226 227    void set_input_k_shift(ggml_tensor * dst) const;228 229    void set_input_kq_mask   (ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn) const;230    void set_input_pos_bucket(ggml_tensor * dst, const llama_ubatch * ubatch) const;231 232    void set_input_k_rot(ggml_tensor * dst) const;233    void set_input_v_rot(ggml_tensor * dst) const;234 235    // true if llama_kv_cell_ext holds information that has to survive a state save/restore236    bool has_cell_ext() const;237 238    // for every token of the ubatch, the ids of the n tokens that precede it in its sequence239    // example for M-RoPE image case: tokens A B X X X C, where X is a 3-token image at pos 2 spanning positions 2..4:240    //   tok: A B X X X C241    //   pos: 0 1 2 2 2 5242    //   prev, n=2: A -> [NULL, NULL], B -> [NULL, A], 3rd X -> [X, X], C -> [X, X]243    // note: used by n-gram input embeddings244    void get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector<llama_token> & res) const;245 246private:247    const llama_model & model;248    const llama_hparams & hparams;249 250    struct kv_layer {251        // layer index in the model252        // note: can be different from the layer index in the KV cache253        uint32_t il;254 255        ggml_tensor * k;256        ggml_tensor * v;257 258        std::vector<ggml_tensor *> k_stream;259        std::vector<ggml_tensor *> v_stream;260    };261 262    bool v_trans = true;  // the value tensor is transposed263 264    const uint32_t n_seq_max = 1;265    const uint32_t n_stream  = 1;266 267    // required padding268    const uint32_t n_pad = 1;269 270    // SWA271    const uint32_t n_swa = 0;272 273    // env: LLAMA_ATTN_ROT_DISABLE274    bool attn_rot_k = false;275    bool attn_rot_v = false;276 277    // if all layers participating in the cache have constant head size, the value is stored here278    // otherwise the value is -1279    int32_t n_embd_head_k_all = 0;280    int32_t n_embd_head_v_all = 0;281 282    // pre-computed hadamard martrices283    std::unordered_map<int64_t, std::vector<float>> attn_rot_hadamard;284 285    // env: LLAMA_KV_CACHE_DEBUG286    int debug = 0;287 288    // this is the SWA type of the cache - not to be confused with the model SWA type289    const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE;290 291    // ggml contexts for the KV cache along with the allocated backend buffers:292    std::vector<std::pair<ggml_context_ptr, ggml_backend_buffer_ptr>> ctxs_bufs;293 294    // the current index from where we start searching for a free slot in the ring buffer of KV cells (see find_slot())295    // note: this is not part of the KV state and it's only used to speed-up the find_slot() method296    std::vector<uint32_t> v_heads;297 298    // TODO: temporary until we refactor to be able to share the same cells between 2 kv caches [TAG_KV_CACHE_SHARE_CELLS]299    llama_kv_cache * other;300 301    std::shared_ptr<llama_kv_cells_vec> v_cells_impl;302 303    llama_kv_cells_vec & v_cells;304 305    // maps from a sequence id to a stream id306    std::vector<uint32_t> seq_to_stream;307 308    // pending stream copies that will be applied during the next update309    stream_copy_info sc_info;310 311    std::vector<kv_layer> layers;312 313    // model layer id -> KV cache layer id314    std::unordered_map<int32_t, int32_t> map_layer_ids;315 316    size_t total_size() const;317 318    size_t size_k_bytes() const;319    size_t size_v_bytes() const;320 321    ggml_tensor * build_rope_shift(322            const llama_cparams & cparams,323                   ggml_context * ctx,324                    ggml_tensor * cur,325                    ggml_tensor * shift,326                    ggml_tensor * rot,327                    ggml_tensor * factors,328                          float   freq_base,329                          float   freq_scale,330                       uint32_t   il) const;331 332    ggml_cgraph * build_graph_shift(333               llm_graph_result * res,334                  llama_context * lctx) const;335 336    struct cell_ranges_t {337        uint32_t strm;338 339        std::vector<std::pair<uint32_t, uint32_t>> data; // ranges, from inclusive, to exclusive340    };341 342    void state_write_meta(llama_io_write_i & io, const cell_ranges_t & cr, llama_seq_id seq_id = -1) const;343    void state_write_data(llama_io_write_i & io, const cell_ranges_t & cr) const;344 345    // sinfo_in, when set, replaces the find_slot call: the cells are given by the caller346    bool state_read_meta(llama_io_read_i & io, uint32_t strm, uint32_t cell_count,       slot_info & sinfo, llama_seq_id dest_seq_id = -1, const slot_info * sinfo_in = nullptr);347    bool state_read_data(llama_io_read_i & io, uint32_t strm, uint32_t cell_count, const slot_info & sinfo);348};349 350class llama_kv_cache_context : public llama_memory_context_i {351public:352    // some shorthands353    using slot_info_vec_t  = llama_kv_cache::slot_info_vec_t;354    using stream_copy_info = llama_kv_cache::stream_copy_info;355 356    // used for errors357    llama_kv_cache_context(llama_memory_status status);358 359    // used to create a full-cache context360    llama_kv_cache_context(361            llama_kv_cache * kv);362 363    // used to create an update context364    llama_kv_cache_context(365            llama_kv_cache * kv,366            llama_context * lctx,367            bool do_shift,368            stream_copy_info sc_info);369 370    // used to create a batch processing context from a batch371    llama_kv_cache_context(372            llama_kv_cache * kv,373            slot_info_vec_t sinfos,374            std::vector<llama_ubatch> ubatches);375 376    virtual ~llama_kv_cache_context();377 378    //379    // llama_memory_context_i380    //381 382    bool next()  override;383    bool apply() override;384 385    llama_memory_status  get_status() const override;386    const llama_ubatch & get_ubatch() const override;387 388    //389    // llama_kv_cache_context specific API390    //391 392    uint32_t get_n_kv() const;393 394    ggml_type type_k() const;395    ggml_type type_v() const;396 397    // get views of the current state of the cache398    ggml_tensor * get_k(ggml_context * ctx, int32_t il) const;399    ggml_tensor * get_v(ggml_context * ctx, int32_t il) const;400 401    // store k_cur and v_cur in the cache based on the provided head location402    // note: the heads in k_cur and v_cur should be laid out contiguously in memory403    //   - k_cur  [n_embd_head_k, n_head_k, n_tokens]404    //   - k_idxs [n_tokens]405    //   - v_cur  [n_embd_head_v, n_head_v, n_tokens]406    //   - v_idxs [n_tokens] or [n_tokens*n_embd_v_gqa] depending if V cache is transposed407    ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const;408    ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il) const;409 410    // create destination indices for each head of the current batch for where it would be written in the KV cache411    // the indices address the global KV cache (not per stream) - this is not relevant for the user of this API, but412    //   helps understand the implementation logic of cpy_k and cpy_v413    ggml_tensor * build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const;414    ggml_tensor * build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const;415 416    ggml_tensor * build_input_k_rot(ggml_context * ctx) const;417    ggml_tensor * build_input_v_rot(ggml_context * ctx) const;418 419    void set_input_k_idxs(ggml_tensor * dst, const llama_ubatch * ubatch) const;420    void set_input_v_idxs(ggml_tensor * dst, const llama_ubatch * ubatch) const;421 422    void set_input_k_shift   (ggml_tensor * dst) const;423    void set_input_kq_mask   (ggml_tensor * dst, const llama_ubatch * ubatch, bool causal_attn) const;424    void set_input_pos_bucket(ggml_tensor * dst, const llama_ubatch * ubatch) const;425 426    void set_input_k_rot(ggml_tensor * dst) const;427    void set_input_v_rot(ggml_tensor * dst) const;428 429    // see llama_kv_cache::get_prev_tokens()430    void get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector<llama_token> & res) const;431 432private:433    llama_memory_status status;434 435    llama_kv_cache * kv;436    llama_context * lctx;437 438    //439    // update context440    //441 442    bool do_shift = false;443 444    stream_copy_info sc_info;445 446    //447    // batch processing context448    //449 450    // the index of the cur ubatch to process451    size_t i_cur = 0;452 453    slot_info_vec_t sinfos;454 455    std::vector<llama_ubatch> ubatches;456 457    //458    // data needed for building the compute graph for the current ubatch:459    //460 461    // a heuristic, to avoid attending the full cache if it is not yet utilized462    // as the cache gets filled, the benefit from this heuristic disappears463    int32_t n_kv;464};465