Felipe97/llama-cpp-compiled
01.2k
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 