KBaba7/llama.cpp
0
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 