Felipe97/llama-cpp-compiled
01.1k
1#include "llama-kv-cache-dsa.h"2 3#include "llama-impl.h"4#include "llama-batch.h"5#include "llama-model.h"6 7#include <algorithm>8#include <cassert>9 10//11// llama_kv_cache_dsa12//13 14llama_kv_cache_dsa::llama_kv_cache_dsa(15 const llama_model & model,16 ggml_type type_k,17 ggml_type type_v,18 bool v_trans,19 bool offload,20 bool unified,21 uint32_t kv_size,22 uint32_t n_seq_max,23 uint32_t n_pad,24 uint32_t n_swa,25 llama_swa_type swa_type,26 const layer_filter_cb & filter_mla,27 const layer_filter_cb & filter_lid,28 const layer_reuse_cb & reuse) :29 hparams_lid(model.hparams), n_stream(unified ? 1 : n_seq_max) {30 31 LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size);32 33 kv_mla = std::make_unique<llama_kv_cache>(34 model, model.hparams, type_k, type_v,35 v_trans, offload, unified, kv_size, n_seq_max, n_pad,36 n_swa, swa_type, nullptr, filter_mla, reuse, nullptr);37 38 // we use llama_kv_cache for caching indexer keys39 // by hand-tweaking some hparams we fool it to create40 // indexer key cache tensors with correct dimensions41 // https://github.com/ggml-org/llama.cpp/pull/21149#discussion_r301594082342 43 // DSA lightning indexer uses MQA with single key head44 std::fill(hparams_lid.n_head_kv_arr.begin(), hparams_lid.n_head_kv_arr.end(), 1);45 hparams_lid.n_embd_head_k_full = model.hparams.indexer_head_size;46 hparams_lid.rope_type = LLAMA_ROPE_TYPE_NEOX;47 48 LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size);49 50 kv_lid = std::make_unique<llama_kv_cache>(51 model, hparams_lid, type_k, type_v,52 v_trans, offload, unified, kv_size, n_seq_max, n_pad,53 n_swa, swa_type, nullptr, filter_lid, reuse, nullptr);54}55 56void llama_kv_cache_dsa::clear(bool data) {57 kv_mla->clear(data);58 kv_lid->clear(data);59}60 61bool llama_kv_cache_dsa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {62 bool res = true;63 64 res = res & kv_mla->seq_rm(seq_id, p0, p1);65 res = res & kv_lid->seq_rm(seq_id, p0, p1);66 67 return res;68}69 70void llama_kv_cache_dsa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {71 kv_mla->seq_cp(seq_id_src, seq_id_dst, p0, p1);72 kv_lid->seq_cp(seq_id_src, seq_id_dst, p0, p1);73}74 75void llama_kv_cache_dsa::seq_keep(llama_seq_id seq_id) {76 kv_mla->seq_keep(seq_id);77 kv_lid->seq_keep(seq_id);78}79 80void llama_kv_cache_dsa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {81 kv_mla->seq_add(seq_id, p0, p1, shift);82 kv_lid->seq_add(seq_id, p0, p1, shift);83}84 85void llama_kv_cache_dsa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {86 kv_mla->seq_div(seq_id, p0, p1, d);87 kv_lid->seq_div(seq_id, p0, p1, d);88}89 90llama_pos llama_kv_cache_dsa::seq_pos_min(llama_seq_id seq_id) const {91 return kv_mla->seq_pos_min(seq_id);92}93 94llama_pos llama_kv_cache_dsa::seq_pos_max(llama_seq_id seq_id) const {95 return kv_mla->seq_pos_max(seq_id);96}97 98std::map<ggml_backend_buffer_type_t, size_t> llama_kv_cache_dsa::memory_breakdown() const {99 std::map<ggml_backend_buffer_type_t, size_t> mb = kv_mla->memory_breakdown();100 for (const auto & buft_size : kv_lid->memory_breakdown()) {101 mb[buft_size.first] += buft_size.second;102 }103 return mb;104}105 106llama_memory_context_ptr llama_kv_cache_dsa::init_batch(107 llama_batch_allocr & balloc,108 uint32_t n_ubatch,109 bool embd_all) {110 GGML_UNUSED(embd_all);111 112 do {113 balloc.split_reset();114 115 std::vector<llama_ubatch> ubatches;116 while (true) {117 auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0);118 119 if (ubatch.n_tokens == 0) {120 break;121 }122 123 ubatches.push_back(std::move(ubatch)); // NOLINT124 }125 126 if (balloc.get_n_used() < balloc.get_n_tokens()) {127 // failed to find a suitable split128 break;129 }130 131 auto sinfos_mla = kv_mla->prepare(ubatches);132 if (sinfos_mla.empty()) {133 break;134 }135 136 auto sinfos_lid = kv_lid->prepare(ubatches);137 if (sinfos_lid.empty()) {138 break;139 }140 141 assert(sinfos_mla.size() == sinfos_lid.size());142 143 return std::make_unique<llama_kv_cache_dsa_context>(144 this, std::move(sinfos_mla), std::move(sinfos_lid), std::move(ubatches));145 } while (false);146 147 return std::make_unique<llama_kv_cache_dsa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);148}149 150llama_memory_context_ptr llama_kv_cache_dsa::init_full() {151 return std::make_unique<llama_kv_cache_dsa_context>(this);152}153 154llama_memory_context_ptr llama_kv_cache_dsa::init_update(llama_context * lctx, bool optimize) {155 return std::make_unique<llama_kv_cache_dsa_context>(this, lctx, optimize);156}157 158bool llama_kv_cache_dsa::get_can_shift() const {159 return kv_mla->get_can_shift() &&160 kv_lid->get_can_shift() &&161 kv_mla->get_size() == kv_lid->get_size();162}163 164void llama_kv_cache_dsa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {165 kv_mla->state_write(io, seq_id, flags);166 kv_lid->state_write(io, seq_id, flags);167}168 169void llama_kv_cache_dsa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {170 kv_mla->state_read(io, seq_id, flags);171 kv_lid->state_read(io, seq_id, flags);172}173 174llama_kv_cache * llama_kv_cache_dsa::get_mla() const {175 return kv_mla.get();176}177 178llama_kv_cache * llama_kv_cache_dsa::get_lid() const {179 return kv_lid.get();180}181 182//183// llama_kv_cache_dsa_context184//185 186llama_kv_cache_dsa_context::llama_kv_cache_dsa_context(llama_memory_status status) : status(status) {}187 188llama_kv_cache_dsa_context::llama_kv_cache_dsa_context(189 llama_kv_cache_dsa * kv) :190 ctx_mla(kv->get_mla()->init_full()),191 ctx_lid(kv->get_lid()->init_full()),192 status(llama_memory_status_combine(ctx_mla->get_status(), ctx_lid->get_status())) {193}194 195llama_kv_cache_dsa_context::llama_kv_cache_dsa_context(196 llama_kv_cache_dsa * kv,197 llama_context * lctx,198 bool optimize) :199 ctx_mla(kv->get_mla()->init_update(lctx, optimize)),200 ctx_lid(kv->get_lid()->init_update(lctx, optimize)),201 status(llama_memory_status_combine(ctx_mla->get_status(), ctx_lid->get_status())) {202}203 204llama_kv_cache_dsa_context::llama_kv_cache_dsa_context(205 llama_kv_cache_dsa * kv,206 slot_info_vec_t sinfos_mla,207 slot_info_vec_t sinfos_lid,208 std::vector<llama_ubatch> ubatches) :209 ubatches(std::move(ubatches)),210 // note: here we copy the ubatches. not sure if this is ideal211 ctx_mla(new llama_kv_cache_context(kv->get_mla(), std::move(sinfos_mla), this->ubatches)),212 ctx_lid(new llama_kv_cache_context(kv->get_lid(), std::move(sinfos_lid), this->ubatches)),213 status(llama_memory_status_combine(ctx_mla->get_status(), ctx_lid->get_status())) {214}215 216llama_kv_cache_dsa_context:: ~llama_kv_cache_dsa_context() = default;217 218bool llama_kv_cache_dsa_context::next() {219 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);220 221 ctx_mla->next();222 ctx_lid->next();223 224 if (++i_next >= ubatches.size()) {225 return false;226 }227 228 return true;229}230 231bool llama_kv_cache_dsa_context::apply() {232 assert(!llama_memory_status_is_fail(status));233 234 bool res = true;235 236 res = res & ctx_mla->apply();237 res = res & ctx_lid->apply();238 239 return res;240}241 242llama_memory_status llama_kv_cache_dsa_context::get_status() const {243 return status;244}245 246const llama_ubatch & llama_kv_cache_dsa_context::get_ubatch() const {247 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);248 249 return ubatches[i_next];250}251 252const llama_kv_cache_context * llama_kv_cache_dsa_context::get_mla() const {253 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);254 255 return static_cast<const llama_kv_cache_context *>(ctx_mla.get());256}257 258const llama_kv_cache_context * llama_kv_cache_dsa_context::get_lid() const {259 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);260 261 return static_cast<const llama_kv_cache_context *>(ctx_lid.get());262}263 