Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
llama-kv-cache.h219 linesDownload Raw Back to src
1#pragma once2 3#include "llama.h"4 5#include "ggml-cpp.h"6 7#include <set>8#include <vector>9 10struct llama_kv_cell {11    llama_pos pos   = -1;12    llama_pos delta = 0;13    int32_t   src   = -1; // used by recurrent state models to copy states14    int32_t   tail  = -1;15 16    std::set<llama_seq_id> seq_id;17 18    bool has_seq_id(const llama_seq_id & id) const {19        return seq_id.find(id) != seq_id.end();20    }21 22    bool is_empty() const {23        return seq_id.empty();24    }25 26    bool is_same_seq(const llama_kv_cell & other) const {27        return seq_id == other.seq_id;28    }29};30 31// ring-buffer of cached KV data32struct llama_kv_cache {33    bool has_shift = false;34    bool do_defrag = false;35    bool recurrent = false; // with recurrent state models, a cell can hold the state for more than one past token36    bool v_trans   = true;  // the value tensor is transposed37    bool can_shift = false;38 39    // Note: The value of head isn't only used to optimize searching40    // for a free KV slot. llama_decode_internal also uses it, so it41    // cannot be freely changed after a slot has been allocated.42    uint32_t head = 0;43    uint32_t size = 0;44    uint32_t used = 0; // used cells (i.e. at least one seq_id)45 46    // computed before each graph build47    uint32_t n = 0;48 49    ggml_type type_k = GGML_TYPE_F16;50    ggml_type type_v = GGML_TYPE_F16;51 52    std::vector<llama_kv_cell> cells;53 54    std::vector<struct ggml_tensor *> k_l; // per layer55    std::vector<struct ggml_tensor *> v_l;56 57    std::vector<ggml_context_ptr> ctxs;58    std::vector<ggml_backend_buffer_ptr> bufs;59 60    size_t total_size() const {61        size_t size = 0;62        for (const auto & buf : bufs) {63            size += ggml_backend_buffer_get_size(buf.get());64        }65 66        return size;67    }68 69    // TODO: better data structures to reduce the cost of this operation70    llama_pos max_pos() const {71        llama_pos max_pos = -1;72        for (const auto & cell : cells) {73            max_pos = std::max(max_pos, cell.pos);74        }75 76        return max_pos;77    }78};79 80// a structure holds information about the slot found in llama_kv_cache_find_slot81struct llama_kv_cache_slot_info {82    std::pair<uint32_t, uint32_t> boundaries; // slot boundaries [begin, end)83    bool found = false;                       // the slot was found84 85    explicit llama_kv_cache_slot_info(bool found_) : found{found_} {}86    llama_kv_cache_slot_info(uint32_t begin, uint32_t end) : boundaries{begin, end}, found{true} {}87 88    operator bool() const { return found; }89};90 91// TODO: maybe not needed92uint32_t llama_kv_cache_get_padding(const struct llama_cparams & cparams);93 94bool llama_kv_cache_init(95        struct llama_kv_cache & cache,96            const llama_model & model,97          const llama_cparams & cparams,98                    ggml_type   type_k,99                    ggml_type   type_v,100                     uint32_t   kv_size,101                         bool   offload);102 103// find an empty slot of size "n_tokens" in the cache104// updates the cache head105// returns a structure holding information about the slot found106// Note: On success, it's important that cache.head points107// to the first cell of the slot.108struct llama_kv_cache_slot_info llama_kv_cache_find_slot(109           struct llama_kv_cache & cache,110       const struct llama_ubatch & batch);111 112// find how many cells are currently in use113uint32_t llama_kv_cache_cell_max(const struct llama_kv_cache & cache);114 115void llama_kv_cache_clear(struct llama_kv_cache & cache);116 117bool llama_kv_cache_seq_rm(118        struct llama_kv_cache & cache,119                 llama_seq_id   seq_id,120                    llama_pos   p0,121                    llama_pos   p1);122 123void llama_kv_cache_seq_cp(124        struct llama_kv_cache & cache,125                 llama_seq_id   seq_id_src,126                 llama_seq_id   seq_id_dst,127                    llama_pos   p0,128                    llama_pos   p1);129 130void llama_kv_cache_seq_keep(131        struct llama_kv_cache & cache,132                 llama_seq_id   seq_id);133 134void llama_kv_cache_seq_add(135        struct llama_kv_cache & cache,136                 llama_seq_id   seq_id,137                    llama_pos   p0,138                    llama_pos   p1,139                    llama_pos   delta);140 141void llama_kv_cache_seq_div(142        struct llama_kv_cache & cache,143                 llama_seq_id   seq_id,144                    llama_pos   p0,145                    llama_pos   p1,146                          int   d);147 148llama_pos llama_kv_cache_seq_pos_max(149        struct llama_kv_cache & cache,150                 llama_seq_id   seq_id);151 152void llama_kv_cache_defrag(struct llama_kv_cache & cache);153 154int32_t llama_get_kv_cache_token_count(const struct llama_kv_cache & kv);155 156int32_t llama_get_kv_cache_used_cells(const struct llama_kv_cache & kv);157 158bool llama_kv_cache_can_shift(const struct llama_kv_cache & kv);159 160//161// kv cache view162//163 164struct llama_kv_cache_view llama_kv_cache_view_init(const struct llama_kv_cache & kv, int32_t n_seq_max);165 166void llama_kv_cache_view_update(struct llama_kv_cache_view * view, const struct llama_kv_cache & kv);167 168//169// kv cache restore170//171 172// saves the kv_cache state for future recovery.173// used to rollback llama_kv_cache_find_slot changes.174struct llama_kv_slot_restorer {175    struct llama_kv_cache_state {176        uint32_t head = 0;177        uint32_t n    = 0;178    } old_state;179 180    // for non-recurrent models only181    // list of slots to restore182    std::vector<std::pair<uint32_t, uint32_t>> slot_boundaries;183 184    bool do_restore = false;185 186    explicit llama_kv_slot_restorer(const struct llama_kv_cache & cache) {187        old_state.head = cache.head;188        old_state.n    = cache.n;189    }190 191    // saves a slot information for future restoration192    void save(const struct llama_kv_cache_slot_info & slot) {193        if (slot) {194            do_restore = true;195            if (slot.boundaries.first != slot.boundaries.second) {196                slot_boundaries.push_back(slot.boundaries);197            }198        }199    }200 201    // must be explicitly called to restore the kv_cache state202    // and rollback changes from all llama_kv_cache_find_slot calls203    void restore(struct llama_kv_cache & cache) {204        if (do_restore) {205            cache.head = old_state.head;206            cache.n    = old_state.n;207 208            if (cache.recurrent) { // recurrent models like Mamba or RWKV can't have a state partially erased209                llama_kv_cache_seq_rm(cache, -1, -1, -1);210            } else {211                for (auto & slot : slot_boundaries) {212                    llama_kv_cache_seq_rm(cache, -1, slot.first, slot.second);213                }214            }215        }216    }217};218 219