Felipe97/llama-cpp-compiled
01.1k
1#include "llama-kv-cache-msa.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#include <cmath>10 11// llama_kv_cache_msa12 13llama_kv_cache_msa::llama_kv_cache_msa(14 const llama_model & model,15 ggml_type type_k,16 ggml_type type_v,17 bool v_trans,18 bool offload,19 bool unified,20 uint32_t kv_size,21 uint32_t n_seq_max,22 uint32_t n_pad,23 uint32_t n_swa,24 llama_swa_type swa_type,25 const layer_filter_cb & filter,26 const layer_filter_cb & filter_idx,27 const layer_reuse_cb & reuse) :28 hparams_idx(model.hparams),29 n_stream(unified ? 1 : n_seq_max), n_seq_max(n_seq_max), n_pad(n_pad),30 n_swa(n_swa), swa_type(swa_type) {31 32 LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size);33 34 kv_base = std::make_unique<llama_kv_cache>(35 model, model.hparams, type_k, type_v,36 v_trans, offload, unified, kv_size, n_seq_max, n_pad,37 n_swa, swa_type, nullptr, filter, reuse, nullptr);38 39 // the MSA indexer uses a single key head per layer40 std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1);41 hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size;42 // the rope parameters are kept identical to the main cache43 44 LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size);45 46 kv_idx = std::make_unique<llama_kv_cache>(47 model, hparams_idx, type_k, type_v,48 v_trans, offload, unified, kv_size, n_seq_max, n_pad,49 n_swa, swa_type, nullptr, filter_idx, reuse, nullptr);50}51 52void llama_kv_cache_msa::clear(bool data) {53 kv_base->clear(data);54 kv_idx ->clear(data);55}56 57bool llama_kv_cache_msa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {58 bool res = true;59 60 res = res & kv_base->seq_rm(seq_id, p0, p1);61 res = res & kv_idx ->seq_rm(seq_id, p0, p1);62 63 return res;64}65 66void llama_kv_cache_msa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {67 kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1);68 kv_idx ->seq_cp(seq_id_src, seq_id_dst, p0, p1);69}70 71void llama_kv_cache_msa::seq_keep(llama_seq_id seq_id) {72 kv_base->seq_keep(seq_id);73 kv_idx ->seq_keep(seq_id);74}75 76void llama_kv_cache_msa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {77 kv_base->seq_add(seq_id, p0, p1, shift);78 kv_idx ->seq_add(seq_id, p0, p1, shift);79}80 81void llama_kv_cache_msa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {82 kv_base->seq_div(seq_id, p0, p1, d);83 kv_idx ->seq_div(seq_id, p0, p1, d);84}85 86llama_pos llama_kv_cache_msa::seq_pos_min(llama_seq_id seq_id) const {87 return kv_base->seq_pos_min(seq_id);88}89 90llama_pos llama_kv_cache_msa::seq_pos_max(llama_seq_id seq_id) const {91 return kv_base->seq_pos_max(seq_id);92}93 94std::map<ggml_backend_buffer_type_t, size_t> llama_kv_cache_msa::memory_breakdown() const {95 std::map<ggml_backend_buffer_type_t, size_t> mb = kv_base->memory_breakdown();96 for (const auto & buft_size : kv_idx->memory_breakdown()) {97 mb[buft_size.first] += buft_size.second;98 }99 return mb;100}101 102llama_memory_context_ptr llama_kv_cache_msa::init_batch(103 llama_batch_allocr & balloc,104 uint32_t n_ubatch,105 bool embd_all) {106 GGML_UNUSED(embd_all);107 108 do {109 balloc.split_reset();110 111 std::vector<llama_ubatch> ubatches;112 while (true) {113 auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0);114 115 if (ubatch.n_tokens == 0) {116 break;117 }118 119 ubatches.push_back(std::move(ubatch));120 }121 122 if (balloc.get_n_used() < balloc.get_n_tokens()) {123 // failed to find a suitable split124 break;125 }126 127 auto sinfos_base = kv_base->prepare(ubatches);128 if (sinfos_base.empty()) {129 break;130 }131 132 auto sinfos_idx = kv_idx->prepare(ubatches);133 if (sinfos_idx.empty()) {134 break;135 }136 137 assert(sinfos_base.size() == sinfos_idx.size());138 139 return std::make_unique<llama_kv_cache_msa_context>(140 this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches));141 } while (false);142 143 return std::make_unique<llama_kv_cache_msa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);144}145 146llama_memory_context_ptr llama_kv_cache_msa::init_full() {147 return std::make_unique<llama_kv_cache_msa_context>(this);148}149 150llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) {151 return std::make_unique<llama_kv_cache_msa_context>(this, lctx, optimize);152}153 154bool llama_kv_cache_msa::get_can_shift() const {155 return kv_base->get_can_shift() &&156 kv_idx ->get_can_shift() &&157 kv_base->get_size() == kv_idx->get_size();158}159 160void llama_kv_cache_msa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {161 kv_base->state_write(io, seq_id, flags);162 kv_idx ->state_write(io, seq_id, flags);163}164 165void llama_kv_cache_msa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {166 kv_base->state_read(io, seq_id, flags);167 kv_idx ->state_read(io, seq_id, flags);168}169 170llama_kv_cache * llama_kv_cache_msa::get_base() const {171 return kv_base.get();172}173 174llama_kv_cache * llama_kv_cache_msa::get_idx() const {175 return kv_idx.get();176}177 178// llama_kv_cache_msa_context179 180llama_kv_cache_msa_context::llama_kv_cache_msa_context(llama_memory_status status) :181 kv(nullptr), status(status) {}182 183llama_kv_cache_msa_context::llama_kv_cache_msa_context(184 llama_kv_cache_msa * kv) :185 kv(kv),186 ctx_base(kv->get_base()->init_full()),187 ctx_idx (kv->get_idx ()->init_full()),188 status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {189}190 191llama_kv_cache_msa_context::llama_kv_cache_msa_context(192 llama_kv_cache_msa * kv,193 llama_context * lctx,194 bool optimize) :195 kv(kv),196 ctx_base(kv->get_base()->init_update(lctx, optimize)),197 ctx_idx (kv->get_idx ()->init_update(lctx, optimize)),198 status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {199}200 201llama_kv_cache_msa_context::llama_kv_cache_msa_context(202 llama_kv_cache_msa * kv,203 slot_info_vec_t sinfos_base,204 slot_info_vec_t sinfos_idx,205 std::vector<llama_ubatch> ubatches) :206 kv(kv),207 ubatches(std::move(ubatches)),208 // here we copy the ubatches. not sure if this is ideal209 ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)),210 ctx_idx (new llama_kv_cache_context(kv->get_idx (), std::move(sinfos_idx), this->ubatches)),211 status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {212}213 214llama_kv_cache_msa_context::~llama_kv_cache_msa_context() = default;215 216bool llama_kv_cache_msa_context::next() {217 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);218 219 ctx_base->next();220 ctx_idx ->next();221 222 if (++i_next >= ubatches.size()) {223 return false;224 }225 226 return true;227}228 229bool llama_kv_cache_msa_context::apply() {230 assert(!llama_memory_status_is_fail(status));231 232 bool res = true;233 234 res = res & ctx_base->apply();235 res = res & ctx_idx ->apply();236 237 return res;238}239 240llama_memory_status llama_kv_cache_msa_context::get_status() const {241 return status;242}243 244const llama_ubatch & llama_kv_cache_msa_context::get_ubatch() const {245 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);246 247 return ubatches[i_next];248}249 250const llama_kv_cache_context * llama_kv_cache_msa_context::get_base() const {251 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);252 253 return static_cast<const llama_kv_cache_context *>(ctx_base.get());254}255 256const llama_kv_cache_context * llama_kv_cache_msa_context::get_idx() const {257 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);258 259 return static_cast<const llama_kv_cache_context *>(ctx_idx.get());260}261 262uint32_t llama_kv_cache_msa_context::get_n_pos() const {263 // pad the value so that the graph remains constant across batches and can be reused264 const uint32_t n_pad_cur = std::max(kv->get_n_pad(), 256u);265 266 llama_pos pos_max = -1;267 268 for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) kv->get_n_seq_max(); ++seq_id) {269 pos_max = std::max(pos_max, kv->seq_pos_max(seq_id));270 }271 272 return std::max(n_pad_cur, GGML_PAD((uint32_t) (pos_max + 1), n_pad_cur));273}274 275void llama_kv_cache_msa_context::set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const {276 GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));277 GGML_ASSERT(dst->type == GGML_TYPE_I32);278 GGML_ASSERT(div > 0);279 280 const int64_t n_tokens = ubatch->n_tokens;281 const int64_t n_kv = dst->ne[0];282 const int64_t n_stream_ub = dst->ne[1];283 284 GGML_ASSERT(n_tokens % n_stream_ub == 0);285 const int64_t n_tps = n_tokens/n_stream_ub;286 287 int32_t * data = (int32_t *) dst->data;288 289 for (int64_t s = 0; s < n_stream_ub; ++s) {290 const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];291 292 const auto & cells = kv->get_base()->get_cells(seq_id);293 294 for (int64_t j = 0; j < n_kv; ++j) {295 // the value for empty or other-sequence cells is irrelevant as consumers mask them296 data[s*n_kv + j] =297 cells.is_empty(j) || !cells.seq_has(j, seq_id)298 ? 0299 : (int32_t) (cells.pos_get(j)/div);300 }301 }302}303 304void llama_kv_cache_msa_context::set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const {305 GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));306 GGML_ASSERT(dst->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_F32);307 308 const int64_t n_tokens = ubatch->n_tokens;309 const int64_t n_pos = dst->ne[0];310 const int64_t n_stream_ub = dst->ne[1];311 312 GGML_ASSERT(n_tokens % n_stream_ub == 0);313 const int64_t n_tps = n_tokens/n_stream_ub;314 315 for (int64_t s = 0; s < n_stream_ub; ++s) {316 const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];317 318 const auto & cells = kv->get_base()->get_cells(seq_id);319 320 std::vector<int32_t> map(n_pos, 0);321 322 for (uint32_t j = 0; j < cells.size(); ++j) {323 if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {324 continue;325 }326 327 const llama_pos p0 = cells.pos_get(j);328 329 if (p0 < 0 || p0 >= n_pos) {330 continue;331 }332 333 map[p0] = (int32_t) j;334 }335 336 if (dst->type == GGML_TYPE_I32) {337 int32_t * data = (int32_t *) dst->data + s*n_pos;338 std::copy(map.begin(), map.end(), data);339 } else {340 float * data = (float *) dst->data + s*n_pos;341 for (int64_t p = 0; p < n_pos; ++p) {342 data[p] = (float) map[p];343 }344 }345 }346}347 348void llama_kv_cache_msa_context::set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const {349 GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));350 GGML_ASSERT(dst->type == GGML_TYPE_F32);351 352 const int64_t n_tokens = ubatch->n_tokens;353 const int64_t n_pos = dst->ne[0];354 355 GGML_ASSERT(dst->ne[1] == n_tokens);356 357 const uint32_t n_swa = kv->get_n_swa();358 const llama_swa_type swa_type = kv->get_swa_type();359 360 float * data = (float *) dst->data;361 362 std::fill(data, data + n_pos*n_tokens, -INFINITY);363 364 for (int64_t i = 0; i < n_tokens; ++i) {365 const llama_seq_id seq_id = ubatch->seq_id[i][0];366 367 const auto & cells = kv->get_base()->get_cells(seq_id);368 369 const llama_pos p1 = ubatch->pos[i];370 371 for (uint32_t j = 0; j < cells.size(); ++j) {372 if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {373 continue;374 }375 376 const llama_pos p0 = cells.pos_get(j);377 378 if (p0 < 0 || p0 >= n_pos) {379 continue;380 }381 382 // causal mask383 if (p0 > p1) {384 continue;385 }386 387 // apply SWA if any388 if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {389 continue;390 }391 392 data[i*n_pos + p0] = 0.0f;393 }394 }395}396 