Brunobkr/llama.cpp_AlgMor24_github
ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.
03.1k
1#include "llama-graph.h"2 3#include "llama-impl.h"4#include "llama-model.h"5#include "llama-batch.h"6#include "llama-cparams.h"7#include "llama-sampler.h"8 9#include "llama-kv-cache.h"10#include "llama-kv-cache-iswa.h"11#include "llama-kv-cache-dsa.h"12#include "llama-kv-cache-msa.h"13#include "llama-kv-cache-dsv4.h"14#include "llama-memory-hybrid.h"15#include "llama-memory-hybrid-iswa.h"16#include "llama-memory-recurrent.h"17 18#include <cassert>19#include <cmath>20#include <cstring>21#include <numeric>22#include <sstream>23#include <string>24#include <unordered_set>25 26// dedup helpers27 28static ggml_tensor * build_attn_inp_kq_mask(29 ggml_context * ctx,30 const llama_kv_cache_context * mctx,31 const llama_ubatch & ubatch,32 const llama_cparams & cparams) {33 const auto n_kv = mctx->get_n_kv();34 const auto n_tokens = ubatch.n_tokens;35 const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq;36 37 // flash attention requires an f16 mask38 const auto type = cparams.flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32;39 40 ggml_tensor * res = ggml_new_tensor_4d(ctx, type, n_kv, n_tokens/n_stream, 1, n_stream);41 ggml_set_input(res);42 ggml_set_name(res, "attn_inp_kq_mask");43 44 return res;45}46 47static bool can_reuse_kq_mask(48 ggml_tensor * kq_mask,49 const llama_kv_cache_context * mctx,50 const llama_ubatch & ubatch,51 const llama_cparams & cparams) {52 const auto n_kv = mctx->get_n_kv();53 const auto n_tokens = ubatch.n_tokens;54 const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq;55 56 bool res = true;57 58 res &= (kq_mask->ne[0] == n_kv);59 res &= (kq_mask->ne[1] == n_tokens/n_stream);60 res &= (kq_mask->ne[2] == 1);61 res &= (kq_mask->ne[3] == n_stream);62 63 return res;64}65 66// impl67 68void llm_graph_input_embd::set_input(const llama_ubatch * ubatch) {69 if (ubatch->token) {70 const int64_t n_tokens = ubatch->n_tokens;71 72 ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));73 }74 75 if (ubatch->embd) {76 GGML_ASSERT(n_embd == embd->ne[0]);77 78 const int64_t n_tokens = ubatch->n_tokens;79 80 ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(embd));81 }82}83 84bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {85 bool res = true;86 87 res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);88 res &= (!params.ubatch.embd) || (embd && embd->ne[1] == params.ubatch.n_tokens);89 90 return res;91}92 93void llm_graph_input_embd_h::set_input(const llama_ubatch * ubatch) {94 const int64_t n_tokens = ubatch->n_tokens;95 96 if (ubatch->token) {97 ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));98 } else {99 // note: mtmd embedding input goes through here100 GGML_ASSERT(ubatch->embd);101 GGML_ASSERT(n_embd == embd->ne[0]);102 103 ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(h));104 }105 106 // TODO: extend llama_ubatch to differentiate between token embeddings and hidden states107 // for now, we assume that the hidden state is always provided as an embedding108 // ref: https://github.com/ggml-org/llama.cpp/pull/23643109 if (ubatch->embd) {110 GGML_ASSERT(n_embd == h->ne[0]);111 112 ggml_backend_tensor_set(h, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(h));113 }114}115 116bool llm_graph_input_embd_h::can_reuse(const llm_graph_params & params) {117 bool res = true;118 119 res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);120 res &= (!params.ubatch.embd) || (embd && embd->ne[1] == params.ubatch.n_tokens);121 res &= (!params.ubatch.embd) || (h && h->ne[1] == params.ubatch.n_tokens);122 123 return res;124}125 126void llm_graph_input_pos::set_input(const llama_ubatch * ubatch) {127 if (ubatch->pos && pos) {128 const int64_t n_tokens = ubatch->n_tokens;129 130 if (ubatch->token && n_pos_per_embd == 4) {131 // in case we're using M-RoPE with text tokens, convert the 1D positions to 4D132 // the 3 first dims are the same, and 4th dim is all 0133 std::vector<llama_pos> pos_data(n_tokens*n_pos_per_embd);134 // copy the first dimension135 for (int i = 0; i < n_tokens; ++i) {136 pos_data[ i] = ubatch->pos[i];137 pos_data[ n_tokens + i] = ubatch->pos[i];138 pos_data[2 * n_tokens + i] = ubatch->pos[i];139 pos_data[3 * n_tokens + i] = 0; // 4th dim is 0140 }141 ggml_backend_tensor_set(pos, pos_data.data(), 0, pos_data.size()*ggml_element_size(pos));142 } else {143 ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));144 }145 }146}147 148bool llm_graph_input_pos::can_reuse(const llm_graph_params & params) {149 bool res = true;150 151 res &= pos->ne[0] == params.ubatch.n_tokens*n_pos_per_embd;152 153 return res;154}155 156void llm_graph_input_attn_temp::set_input(const llama_ubatch * ubatch) {157 if (ubatch->pos && attn_scale) {158 const int64_t n_tokens = ubatch->n_tokens;159 160 GGML_ASSERT(f_attn_temp_scale != 0.0f);161 GGML_ASSERT(n_attn_temp_floor_scale != 0);162 163 std::vector<float> attn_scale_data(n_tokens, 0.0f);164 for (int i = 0; i < n_tokens; ++i) {165 const float pos = ubatch->pos[i];166 attn_scale_data[i] = std::log(167 std::floor((pos + f_attn_temp_offset) / n_attn_temp_floor_scale) + 1.0168 ) * f_attn_temp_scale + 1.0;169 }170 171 ggml_backend_tensor_set(attn_scale, attn_scale_data.data(), 0, n_tokens*ggml_element_size(attn_scale));172 }173}174 175void llm_graph_input_pos_bucket::set_input(const llama_ubatch * ubatch) {176 if (pos_bucket) {177 const int64_t n_tokens = ubatch->n_tokens;178 179 GGML_ASSERT(ggml_backend_buffer_is_host(pos_bucket->buffer));180 GGML_ASSERT(!ubatch->equal_seqs()); // TODO: use ubatch->n_seqs instead of failing181 182 int32_t * data = (int32_t *) pos_bucket->data;183 184 for (int j = 0; j < n_tokens; ++j) {185 for (int i = 0; i < n_tokens; ++i) {186 data[j*n_tokens + i] = llama_relative_position_bucket(ubatch->pos[i], ubatch->pos[j], hparams.n_rel_attn_bkts, true);187 }188 }189 }190}191 192void llm_graph_input_pos_bucket_kv::set_input(const llama_ubatch * ubatch) {193 if (pos_bucket) {194 mctx->set_input_pos_bucket(pos_bucket, ubatch);195 }196}197 198void llm_graph_input_out_ids::set_input(const llama_ubatch * ubatch) {199 GGML_ASSERT(out_ids);200 201 const int64_t n_tokens = ubatch->n_tokens;202 203 GGML_ASSERT(ggml_backend_buffer_is_host(out_ids->buffer));204 int32_t * data = (int32_t *) out_ids->data;205 206 if (n_outputs == n_tokens) {207 for (int i = 0; i < n_tokens; ++i) {208 data[i] = i;209 }210 211 return;212 }213 214 GGML_ASSERT(ubatch->output);215 216 int n_outputs = 0;217 218 for (int i = 0; i < n_tokens; ++i) {219 if (ubatch->output[i]) {220 data[n_outputs++] = i;221 }222 }223}224 225bool llm_graph_input_out_ids::can_reuse(const llm_graph_params & params) {226 bool res = true;227 228 res &= n_outputs == params.n_outputs;229 230 return res;231}232 233void llm_graph_input_mean::set_input(const llama_ubatch * ubatch) {234 if (cparams.embeddings &&235 (cparams.pooling_type == LLAMA_POOLING_TYPE_MEAN ||236 cparams.pooling_type == LLAMA_POOLING_TYPE_RANK )) {237 238 const int64_t n_tokens = ubatch->n_tokens;239 const int64_t n_seq_tokens = ubatch->n_seq_tokens;240 const int64_t n_seqs_unq = ubatch->n_seqs_unq;241 242 GGML_ASSERT(mean);243 GGML_ASSERT(ggml_backend_buffer_is_host(mean->buffer));244 245 float * data = (float *) mean->data;246 memset(mean->data, 0, n_tokens*n_seqs_unq*ggml_element_size(mean));247 248 std::vector<uint64_t> sums(n_seqs_unq, 0);249 for (int i = 0; i < n_tokens; i += n_seq_tokens) {250 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {251 const llama_seq_id seq_id = ubatch->seq_id[i][s];252 const int32_t seq_idx = ubatch->seq_idx[seq_id];253 254 sums[seq_idx] += ubatch->n_seq_tokens;255 }256 }257 258 std::vector<float> div(n_seqs_unq, 0.0f);259 for (int s = 0; s < n_seqs_unq; ++s) {260 const uint64_t sum = sums[s];261 if (sum > 0) {262 div[s] = 1.0f/float(sum);263 }264 }265 266 for (int i = 0; i < n_tokens; i += n_seq_tokens) {267 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {268 const llama_seq_id seq_id = ubatch->seq_id[i][s];269 const int32_t seq_idx = ubatch->seq_idx[seq_id];270 271 for (int j = 0; j < n_seq_tokens; ++j) {272 data[seq_idx*n_tokens + i + j] = div[seq_idx];273 }274 }275 }276 }277}278 279void llm_graph_input_cls::set_input(const llama_ubatch * ubatch) {280 const int64_t n_tokens = ubatch->n_tokens;281 const int64_t n_seqs_unq = ubatch->n_seqs_unq;282 283 if (cparams.embeddings && (284 cparams.pooling_type == LLAMA_POOLING_TYPE_CLS ||285 cparams.pooling_type == LLAMA_POOLING_TYPE_RANK ||286 cparams.pooling_type == LLAMA_POOLING_TYPE_LAST287 )) {288 GGML_ASSERT(cls);289 GGML_ASSERT(ggml_backend_buffer_is_host(cls->buffer));290 291 uint32_t * data = (uint32_t *) cls->data;292 memset(cls->data, 0, n_seqs_unq*ggml_element_size(cls));293 294 std::vector<int> target_pos(n_seqs_unq, -1);295 std::vector<int> target_row(n_seqs_unq, -1);296 297 const bool last = (298 cparams.pooling_type == LLAMA_POOLING_TYPE_LAST ||299 (cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && (arch == LLM_ARCH_QWEN3 || arch == LLM_ARCH_QWEN3VL)) // qwen3 reranking & embedding models use last token300 );301 302 for (int i = 0; i < n_tokens; ++i) {303 const llama_pos pos = ubatch->pos[i];304 305 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {306 const llama_seq_id seq_id = ubatch->seq_id[i][s];307 const int32_t seq_idx = ubatch->seq_idx[seq_id];308 309 if (310 (target_pos[seq_idx] == -1) ||311 ( last && pos >= target_pos[seq_idx]) ||312 (!last && pos < target_pos[seq_idx])313 ) {314 target_pos[seq_idx] = pos;315 target_row[seq_idx] = i;316 }317 }318 }319 320 for (int s = 0; s < n_seqs_unq; ++s) {321 if (target_row[s] >= 0) {322 data[s] = target_row[s];323 }324 }325 }326}327 328void llm_graph_input_rs::set_input(const llama_ubatch * ubatch) {329 GGML_UNUSED(ubatch);330 331 const int64_t n_rs = mctx->get_n_rs();332 333 if (s_copy) {334 GGML_ASSERT(ggml_backend_buffer_is_host(s_copy->buffer));335 int32_t * data = (int32_t *) s_copy->data;336 337 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n338 for (uint32_t i = 0; i < n_rs; ++i) {339 data[i] = mctx->s_copy(i);340 }341 }342}343 344bool llm_graph_input_rs::can_reuse(const llm_graph_params & params) {345 const auto * mctx = static_cast<const llama_memory_recurrent_context *>(params.mctx);346 347 this->mctx = mctx;348 349 bool res = true;350 351 res &= s_copy->ne[0] == mctx->get_n_rs();352 353 res &= s_copy_main->ne[0] == params.ubatch.n_seqs;354 res &= s_copy_extra->ne[0] == mctx->get_n_rs() - params.ubatch.n_seqs;355 356 res &= head == mctx->get_head();357 res &= rs_z == mctx->get_rs_z();358 359 return res;360}361 362void llm_graph_input_cross_embd::set_input(const llama_ubatch * ubatch) {363 GGML_UNUSED(ubatch);364 365 if (cross_embd && !cross->v_embd.empty()) {366 assert(cross_embd->type == GGML_TYPE_F32);367 368 ggml_backend_tensor_set(cross_embd, cross->v_embd.data(), 0, ggml_nbytes(cross_embd));369 }370}371 372template <typename T>373static void print_mask(const T * data, int64_t n_tokens, int64_t n_kv, int64_t n_swa, llama_swa_type swa_type) {374 LLAMA_LOG_DEBUG("%s: === Attention mask ===\n", __func__);375 const char * swa_type_str = "unknown";376 377 switch (swa_type) {378 case LLAMA_SWA_TYPE_NONE: swa_type_str = "LLAMA_SWA_TYPE_NONE"; break;379 case LLAMA_SWA_TYPE_STANDARD: swa_type_str = "LLAMA_SWA_TYPE_STANDARD"; break;380 case LLAMA_SWA_TYPE_CHUNKED: swa_type_str = "LLAMA_SWA_TYPE_CHUNKED"; break;381 case LLAMA_SWA_TYPE_SYMMETRIC: swa_type_str = "LLAMA_SWA_TYPE_SYMMETRIC"; break;382 };383 384 LLAMA_LOG_DEBUG("%s: n_swa : %d, n_kv: %d, swa_type: %s\n", __func__, (int)n_swa, (int)n_kv, swa_type_str);385 LLAMA_LOG_DEBUG("%s: '0' = can attend, '∞' = masked\n", __func__);386 LLAMA_LOG_DEBUG("%s: Rows = query tokens, Columns = key/value tokens\n\n", __func__);387 388 LLAMA_LOG_DEBUG(" ");389 for (int j = 0; j < std::min((int64_t)20, n_kv); ++j) {390 LLAMA_LOG_DEBUG("%2d", j);391 }392 LLAMA_LOG_DEBUG("\n");393 394 for (int i = 0; i < std::min((int64_t)20, n_tokens); ++i) {395 LLAMA_LOG_DEBUG(" %2d ", i);396 for (int j = 0; j < std::min((int64_t)20, n_kv); ++j) {397 float val = llama_cast<float>(data[i * n_kv + j]);398 if (val == -INFINITY) {399 LLAMA_LOG_DEBUG(" ∞");400 } else {401 LLAMA_LOG_DEBUG(" 0");402 }403 }404 LLAMA_LOG_DEBUG("\n");405 }406}407 408void llm_graph_input_attn_no_cache::set_input(const llama_ubatch * ubatch) {409 const int64_t n_kv = ubatch->n_tokens;410 const int64_t n_tokens = ubatch->n_tokens;411 412 const auto fill_mask = [&](auto * data, int64_t ne, int n_swa, llama_swa_type swa_type) {413 using T = std::remove_reference_t<decltype(*data)>;414 std::fill(data, data + ne, llama_cast<T>(-INFINITY));415 416 for (int i1 = 0; i1 < n_tokens; ++i1) {417 const llama_seq_id s1 = ubatch->seq_id[i1][0];418 const llama_pos p1 = ubatch->pos[i1];419 420 const uint64_t idst = i1*n_kv;421 422 for (int i0 = 0; i0 < n_tokens; ++i0) {423 const llama_seq_id s0 = ubatch->seq_id[i0][0];424 const llama_pos p0 = ubatch->pos[i0];425 426 // mask different sequences427 if (s0 != s1) {428 continue;429 }430 431 // mask future tokens432 if (cparams.causal_attn && p0 > p1) {433 continue;434 }435 436 // apply SWA if any437 if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {438 continue;439 }440 441 data[idst + i0] = llama_cast<T>(hparams.use_alibi ? -std::abs(p0 - p1) : 0.0f);442 }443 }444 445 if (debug) {446 print_mask(data, n_tokens, n_kv, n_swa, swa_type);447 }448 };449 450 GGML_ASSERT(self_kq_mask);451 GGML_ASSERT(ggml_backend_buffer_is_host(self_kq_mask->buffer));452 if (self_kq_mask->type == GGML_TYPE_F16) {453 fill_mask((ggml_fp16_t *) self_kq_mask->data, ggml_nelements(self_kq_mask), 0, LLAMA_SWA_TYPE_NONE);454 } else {455 fill_mask((float *) self_kq_mask->data, ggml_nelements(self_kq_mask), 0, LLAMA_SWA_TYPE_NONE);456 }457 458 if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) {459 GGML_ASSERT(self_kq_mask_swa);460 GGML_ASSERT(ggml_backend_buffer_is_host(self_kq_mask_swa->buffer));461 if (self_kq_mask_swa->type == GGML_TYPE_F16) {462 fill_mask((ggml_fp16_t *) self_kq_mask_swa->data, ggml_nelements(self_kq_mask_swa), hparams.n_swa, hparams.swa_type);463 } else {464 fill_mask((float *) self_kq_mask_swa->data, ggml_nelements(self_kq_mask_swa), hparams.n_swa, hparams.swa_type);465 }466 }467}468 469void llm_graph_input_attn_kv::set_input(const llama_ubatch * ubatch) {470 mctx->set_input_k_idxs(self_k_idxs, ubatch);471 mctx->set_input_v_idxs(self_v_idxs, ubatch);472 473 // the mask is left unallocated when the graph only stores K/V without attending474 // (e.g. DFlash's KV-injection pass)475 if (self_kq_mask && self_kq_mask->buffer) {476 mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);477 }478 479 if (self_k_rot && self_k_rot->buffer) {480 mctx->set_input_k_rot(self_k_rot);481 }482 483 if (self_v_rot && self_v_rot->buffer) {484 mctx->set_input_v_rot(self_v_rot);485 }486}487 488bool llm_graph_input_attn_kv::can_reuse(const llm_graph_params & params) {489 const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);490 491 this->mctx = mctx;492 493 bool res = true;494 495 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;496 //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there497 498 res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams);499 500 return res;501}502 503void llm_graph_input_attn_k::set_input(const llama_ubatch * ubatch) {504 mctx->set_input_k_idxs(self_k_idxs, ubatch);505 506 mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);507}508 509bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) {510 const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);511 512 this->mctx = mctx;513 514 bool res = true;515 516 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;517 518 res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams);519 520 return res;521}522 523llm_graph_input_attn_kv_msa::llm_graph_input_attn_kv_msa(524 const llama_hparams & hparams,525 const llama_cparams & cparams,526 const llama_kv_cache_msa_context * mctx) :527 llm_graph_input_attn_kv(hparams, cparams, mctx->get_base()),528 mctx_msa(mctx) {529}530 531void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) {532 llm_graph_input_attn_kv::set_input(ubatch);533 534 if (self_k_idxs_idx) {535 mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch);536 }537}538 539bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) {540 mctx_msa = static_cast<const llama_kv_cache_msa_context *>(params.mctx);541 542 // the parent class operates on the base cache context543 this->mctx = mctx_msa->get_base();544 545 bool res = true;546 547 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;548 if (self_k_idxs_idx) {549 res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens;550 }551 552 res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams);553 554 return res;555}556 557void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) {558 mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch);559 560 mctx->get_mla()->set_input_kq_mask(self_kq_mask_mla, ubatch, cparams.causal_attn);561 562 mctx->get_lid()->set_input_k_idxs(self_k_idxs_lid, ubatch);563 564 mctx->get_lid()->set_input_kq_mask(self_kq_mask_lid, ubatch, cparams.causal_attn);565 566 mctx->get_lid()->set_input_k_rot(self_k_rot_lid);567}568 569bool llm_graph_input_attn_k_dsa::can_reuse(const llm_graph_params & params) {570 const auto * mctx = static_cast<const llama_kv_cache_dsa_context *>(params.mctx);571 572 this->mctx = mctx;573 574 bool res = true;575 576 res &= self_k_idxs_mla->ne[0] == params.ubatch.n_tokens;577 res &= self_k_idxs_lid->ne[0] == params.ubatch.n_tokens;578 579 res &= can_reuse_kq_mask(self_kq_mask_mla, mctx->get_mla(), params.ubatch, params.cparams);580 res &= can_reuse_kq_mask(self_kq_mask_lid, mctx->get_lid(), params.ubatch, params.cparams);581 582 return res;583}584 585void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) {586 // base tensors may not be allocated if there are no non-SWA attention layers587 if (self_k_idxs && self_k_idxs->buffer) {588 mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);589 if (self_v_idxs) {590 mctx->get_base()->set_input_v_idxs(self_v_idxs, ubatch);591 }592 }593 594 // the kq mask guards on its own buffer: shared cells leave idxs unbacked while the mask stays live595 if (self_kq_mask && self_kq_mask->buffer) {596 mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);597 }598 599 // swa tensors may not be allocated if there are no SWA attention layers600 if (self_k_idxs_swa && self_k_idxs_swa->buffer) {601 mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);602 if (self_v_idxs_swa) {603 mctx->get_swa()->set_input_v_idxs(self_v_idxs_swa, ubatch);604 }605 }606 607 if (self_kq_mask_swa && self_kq_mask_swa->buffer) {608 mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);609 }610 611 if (self_k_rot && self_k_rot->buffer) {612 mctx->get_base()->set_input_k_rot(self_k_rot);613 }614 615 if (self_v_rot && self_v_rot->buffer) {616 mctx->get_base()->set_input_v_rot(self_v_rot);617 }618 619 if (self_k_rot_swa && self_k_rot_swa->buffer) {620 mctx->get_swa()->set_input_k_rot(self_k_rot_swa);621 }622 623 if (self_v_rot_swa && self_v_rot_swa->buffer) {624 mctx->get_swa()->set_input_v_rot(self_v_rot_swa);625 }626}627 628bool llm_graph_input_attn_kv_iswa::can_reuse(const llm_graph_params & params) {629 const auto * mctx = static_cast<const llama_kv_cache_iswa_context *>(params.mctx);630 631 this->mctx = mctx;632 633 bool res = true;634 635 // base tensors may not be allocated if there are no non-SWA attention layers636 if (self_k_idxs && self_k_idxs->buffer) {637 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;638 //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there639 }640 641 if (self_kq_mask && self_kq_mask->buffer) {642 res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);643 }644 645 // swa tensors may not be allocated if there are no SWA attention layers646 if (self_k_idxs_swa && self_k_idxs_swa->buffer) {647 res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;648 //res &= self_v_idxs_swa->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there649 }650 651 if (self_kq_mask_swa && self_kq_mask_swa->buffer) {652 res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);653 }654 655 return res;656}657 658void llm_graph_input_attn_k_iswa::set_input(const llama_ubatch * ubatch) {659 // base tensors may not be allocated if there are no non-SWA attention layers660 if (self_k_idxs && self_k_idxs->buffer) {661 mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);662 }663 664 // the kq mask guards on its own buffer: shared cells leave idxs unbacked while the mask stays live665 if (self_kq_mask && self_kq_mask->buffer) {666 mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);667 }668 669 // swa tensors may not be allocated if there are no SWA attention layers670 if (self_k_idxs_swa && self_k_idxs_swa->buffer) {671 mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);672 }673 674 if (self_kq_mask_swa && self_kq_mask_swa->buffer) {675 mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);676 }677 678 if (self_k_rot && self_k_rot->buffer) {679 mctx->get_base()->set_input_k_rot(self_k_rot);680 }681 682 if (self_k_rot_swa && self_k_rot_swa->buffer) {683 mctx->get_swa()->set_input_k_rot(self_k_rot_swa);684 }685}686 687bool llm_graph_input_attn_k_iswa::can_reuse(const llm_graph_params & params) {688 const auto * mctx = static_cast<const llama_kv_cache_iswa_context *>(params.mctx);689 690 this->mctx = mctx;691 692 bool res = true;693 694 // base tensors may not be allocated if there are no non-SWA attention layers695 if (self_k_idxs && self_k_idxs->buffer) {696 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;697 }698 699 if (self_kq_mask && self_kq_mask->buffer) {700 res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);701 }702 703 // swa tensors may not be allocated if there are no SWA attention layers704 if (self_k_idxs_swa && self_k_idxs_swa->buffer) {705 res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;706 }707 708 if (self_kq_mask_swa && self_kq_mask_swa->buffer) {709 res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);710 }711 712 return res;713}714 715static void dsv4_set_i64(ggml_tensor * dst, const std::vector<int64_t> & src) {716 if (!dst || !dst->buffer) {717 return;718 }719 720 GGML_ASSERT(dst->ne[0] == (int64_t) src.size());721 ggml_backend_tensor_set(dst, src.data(), 0, src.size()*ggml_element_size(dst));722}723 724static void dsv4_set_i32(ggml_tensor * dst, const std::vector<int32_t> & src) {725 if (!dst || !dst->buffer) {726 return;727 }728 729 GGML_ASSERT(dst->ne[0] == (int64_t) src.size());730 ggml_backend_tensor_set(dst, src.data(), 0, src.size()*ggml_element_size(dst));731}732 733static void dsv4_set_kq_mask(734 ggml_tensor * dst,735 const llama_kv_cache_dsv4_context::comp_plan & plan,736 uint32_t n_tokens,737 int64_t n_stream) {738 if (!dst || !dst->buffer) {739 return;740 }741 742 GGML_ASSERT(dst->type == GGML_TYPE_F32 || dst->type == GGML_TYPE_F16);743 GGML_ASSERT(n_stream > 0);744 GGML_ASSERT(n_tokens%n_stream == 0);745 GGML_ASSERT(dst->ne[0] == plan.n_kv);746 GGML_ASSERT(dst->ne[1] == (int64_t) n_tokens/n_stream);747 GGML_ASSERT(dst->ne[2] == 1);748 GGML_ASSERT(dst->ne[3] == n_stream);749 GGML_ASSERT((int64_t) plan.n_visible.size() == (int64_t) n_tokens);750 GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));751 752 if (dst->type == GGML_TYPE_F32) {753 float * data = (float *) dst->data;754 755 for (int64_t i = 0; i < (int64_t) n_tokens; ++i) {756 const int32_t n_visible = plan.n_visible[i];757 758 for (int64_t j = 0; j < dst->ne[0]; ++j) {759 data[i*dst->ne[0] + j] = j < n_visible ? 0.0f : -INFINITY;760 }761 }762 } else if (dst->type == GGML_TYPE_F16) {763 ggml_fp16_t * data = (ggml_fp16_t *) dst->data;764 const ggml_fp16_t fp16_ninf = llama_cast<ggml_fp16_t>(-INFINITY);765 const ggml_fp16_t fp16_zero = llama_cast<ggml_fp16_t>(0.0f);766 767 for (int64_t i = 0; i < (int64_t) n_tokens; ++i) {768 const int32_t n_visible = plan.n_visible[i];769 770 for (int64_t j = 0; j < dst->ne[0]; ++j) {771 data[i*dst->ne[0] + j] = j < n_visible ? fp16_zero : fp16_ninf;772 }773 }774 }775}776 777static ggml_tensor * dsv4_build_raw_kq_mask(778 ggml_context * ctx,779 const llama_kv_cache_dsv4_raw_context * mctx,780 const llama_ubatch & ubatch,781 const llama_cparams & cparams,782 int64_t n_stream) {783 const auto n_kv = mctx->get_n_kv();784 const auto n_tokens = ubatch.n_tokens;785 786 GGML_ASSERT(n_stream > 0);787 GGML_ASSERT(n_tokens%n_stream == 0);788 789 const auto type = cparams.flash_attn ? GGML_TYPE_F16 : GGML_TYPE_F32;790 791 ggml_tensor * res = ggml_new_tensor_4d(ctx, type, n_kv, n_tokens/n_stream, 1, n_stream);792 ggml_set_input(res);793 ggml_set_name(res, "attn_inp_kq_mask");794 795 return res;796}797 798static bool dsv4_can_reuse_raw_kq_mask(799 ggml_tensor * kq_mask,800 const llama_kv_cache_dsv4_raw_context * mctx,801 const llama_ubatch & ubatch,802 int64_t n_stream) {803 const auto n_kv = mctx->get_n_kv();804 const auto n_tokens = ubatch.n_tokens;805 806 GGML_ASSERT(n_stream > 0);807 808 bool res = true;809 810 res &= (kq_mask->ne[0] == n_kv);811 res &= (kq_mask->ne[1] == n_tokens/n_stream);812 res &= (kq_mask->ne[2] == 1);813 res &= (kq_mask->ne[3] == n_stream);814 815 return res;816}817 818static std::string dsv4_plan_positions(const std::vector<int32_t> & values) {819 std::ostringstream ss;820 ss << "[";821 for (size_t i = 0; i < values.size(); ++i) {822 if (i > 0) {823 ss << ", ";824 }825 ss << values[i];826 }827 ss << "]";828 return ss.str();829}830 831static bool dsv4_compress_debug() {832 static const bool debug = []() {833 const char * env = getenv("LLAMA_DSV4_COMPRESS_DEBUG");834 return env && atoi(env) > 0;835 }();836 837 return debug;838}839 840static void dsv4_set_comp_inputs(841 const llm_graph_input_dsv4::comp_input & inp,842 const llama_kv_cache_dsv4_context::comp_plan & plan,843 const char * name,844 bool debug,845 uint32_t n_tokens,846 int64_t n_stream) {847 dsv4_set_i32(inp.state_pos, plan.state_pos);848 dsv4_set_i32(inp.state_persist_src_idxs, plan.state_persist_src_idxs);849 dsv4_set_i32(inp.state_persist_dst_idxs, plan.state_persist_dst_idxs);850 dsv4_set_i32(inp.state_restore_src_idxs, plan.state_restore_src_idxs);851 dsv4_set_i32(inp.state_restore_dst_idxs, plan.state_restore_dst_idxs);852 dsv4_set_i32(inp.state_snapshot_src_idxs, plan.state_snapshot_src_idxs);853 dsv4_set_i32(inp.state_snapshot_dst_idxs, plan.state_snapshot_dst_idxs);854 dsv4_set_i32(inp.state_read_idxs, plan.state_read_idxs);855 dsv4_set_i64(inp.state_write_idxs, plan.state_write_idxs);856 dsv4_set_i32(inp.state_write_pos, plan.state_write_pos);857 dsv4_set_kq_mask(inp.kq_mask, plan, n_tokens, n_stream);858 859 if (debug || dsv4_compress_debug()) {860 LLAMA_LOG_INFO("%s: %s n_tokens=%u, n_stream=%d, state_persist_dst=%s, state_write_pos=%s\n",861 __func__, name, n_tokens, (int) n_stream,862 dsv4_plan_positions(plan.state_persist_dst_idxs).c_str(),863 dsv4_plan_positions(plan.state_write_pos).c_str());864 }865}866 867static bool dsv4_can_reuse_tensor_1d(ggml_tensor * t, int64_t ne0) {868 return (t == nullptr && ne0 == 0) || (t != nullptr && t->ne[0] == ne0);869}870 871static bool dsv4_can_reuse_kq_mask(872 ggml_tensor * t,873 const llama_kv_cache_dsv4_context::comp_plan & plan,874 uint32_t n_tokens,875 int64_t n_stream) {876 if (plan.n_kv == 0) {877 return t == nullptr;878 }879 880 GGML_ASSERT(n_stream > 0);881 882 return t != nullptr &&883 t->ne[0] == plan.n_kv &&884 t->ne[1] == (int64_t) n_tokens/n_stream &&885 t->ne[2] == 1 &&886 t->ne[3] == n_stream;887}888 889static bool dsv4_can_reuse_comp_input(890 const llm_graph_input_dsv4::comp_input & inp,891 const llama_kv_cache_dsv4_context::comp_plan & plan,892 uint32_t n_tokens,893 int64_t n_stream) {894 bool res = true;895 res &= dsv4_can_reuse_tensor_1d(inp.state_pos, plan.state_pos.size());896 res &= dsv4_can_reuse_tensor_1d(inp.state_persist_src_idxs, plan.state_persist_src_idxs.size());897 res &= dsv4_can_reuse_tensor_1d(inp.state_persist_dst_idxs, plan.state_persist_dst_idxs.size());898 res &= dsv4_can_reuse_tensor_1d(inp.state_restore_src_idxs, plan.state_restore_src_idxs.size());899 res &= dsv4_can_reuse_tensor_1d(inp.state_restore_dst_idxs, plan.state_restore_dst_idxs.size());900 res &= dsv4_can_reuse_tensor_1d(inp.state_snapshot_src_idxs, plan.state_snapshot_src_idxs.size());901 res &= dsv4_can_reuse_tensor_1d(inp.state_snapshot_dst_idxs, plan.state_snapshot_dst_idxs.size());902 res &= dsv4_can_reuse_tensor_1d(inp.state_read_idxs, plan.state_read_idxs.size());903 res &= dsv4_can_reuse_tensor_1d(inp.state_write_idxs, plan.state_write_idxs.size());904 res &= dsv4_can_reuse_tensor_1d(inp.state_write_pos, plan.state_write_pos.size());905 res &= dsv4_can_reuse_kq_mask(inp.kq_mask, plan, n_tokens, n_stream);906 907 return res;908}909 910static ggml_tensor * dsv4_build_input_1d(911 ggml_context * ctx,912 ggml_type type,913 int64_t ne0,914 const std::string & name) {915 if (ne0 == 0) {916 return nullptr;917 }918 919 ggml_tensor * res = ggml_new_tensor_1d(ctx, type, ne0);920 ggml_set_input(res);921 ggml_set_name(res, name.c_str());922 923 return res;924}925 926static void dsv4_build_comp_inputs(927 ggml_context * ctx,928 llm_graph_input_dsv4::comp_input & inp,929 const llama_kv_cache_dsv4_context::comp_plan & plan,930 const char * name,931 const llama_cparams & cparams,932 int64_t n_stream) {933 inp.state_pos = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_pos.size(), std::string("dsv4_") + name + "_state_pos");934 inp.state_persist_src_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_persist_src_idxs.size(), std::string("dsv4_") + name + "_state_persist_src_idxs");935 inp.state_persist_dst_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_persist_dst_idxs.size(), std::string("dsv4_") + name + "_state_persist_dst_idxs");936 inp.state_restore_src_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_restore_src_idxs.size(), std::string("dsv4_") + name + "_state_restore_src_idxs");937 inp.state_restore_dst_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_restore_dst_idxs.size(), std::string("dsv4_") + name + "_state_restore_dst_idxs");938 inp.state_snapshot_src_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_snapshot_src_idxs.size(), std::string("dsv4_") + name + "_state_snapshot_src_idxs");939 inp.state_snapshot_dst_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_snapshot_dst_idxs.size(), std::string("dsv4_") + name + "_state_snapshot_dst_idxs");940 inp.state_read_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_read_idxs.size(), std::string("dsv4_") + name + "_state_read_idxs");941 inp.state_write_idxs = dsv4_build_input_1d(ctx, GGML_TYPE_I64, plan.state_write_idxs.size(), std::string("dsv4_") + name + "_state_write_idxs");942 inp.state_write_pos = dsv4_build_input_1d(ctx, GGML_TYPE_I32, plan.state_write_pos.size(), std::string("dsv4_") + name + "_state_write_pos");943 944 if (plan.n_kv > 0) {945 const int64_t n_tokens = (int64_t) plan.n_visible.size();946 947 GGML_ASSERT(n_stream > 0);948 GGML_ASSERT(n_tokens%n_stream == 0);949 950 inp.kq_mask = ggml_new_tensor_4d(ctx, (strcmp(name, "lid") != 0 && cparams.flash_attn) || (strcmp(name, "lid") == 0 && cparams.fused_lid) ? GGML_TYPE_F16 : GGML_TYPE_F32, plan.n_kv, n_tokens/n_stream, 1, n_stream);951 ggml_set_input(inp.kq_mask);952 ggml_set_name(inp.kq_mask, (std::string("dsv4_") + name + "_kq_mask").c_str());953 }954}955 956void llm_graph_input_dsv4_raw::set_input(const llama_ubatch * ubatch) {957 if (self_k_idxs && self_k_idxs->buffer) {958 mctx->set_input_k_idxs(self_k_idxs);959 }960 961 if (self_kq_mask && self_kq_mask->buffer) {962 mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);963 }964 965 if (self_k_rot) {966 mctx->set_input_k_rot(self_k_rot);967 }968}969 970void llm_graph_input_dsv4::set_input(const llama_ubatch * ubatch) {971 const auto & plan_csa = mctx->get_csa_plan(*ubatch);972 const auto & plan_hca = mctx->get_hca_plan(*ubatch);973 const auto & plan_lid = mctx->get_lid_plan(*ubatch);974 const int64_t n_stream = plan_csa.n_stream;975 976 inp_raw->mctx = mctx->get_raw();977 inp_raw->set_input(ubatch);978 979 dsv4_set_comp_inputs(inp_csa, plan_csa, "csa", debug > 0, ubatch->n_tokens, n_stream);980 dsv4_set_comp_inputs(inp_hca, plan_hca, "hca", debug > 0, ubatch->n_tokens, n_stream);981 dsv4_set_comp_inputs(inp_lid, plan_lid, "lid", debug > 0, ubatch->n_tokens, n_stream);982 983 if (inp_csa.k_rot && inp_csa.k_rot->buffer) {984 mctx->get_csa()->set_input_k_rot(inp_csa.k_rot);985 }986 987 if (inp_hca.k_rot && inp_hca.k_rot->buffer) {988 mctx->get_hca()->set_input_k_rot(inp_hca.k_rot);989 }990 991 if (inp_lid.k_rot && inp_lid.k_rot->buffer) {992 mctx->get_lid()->set_input_k_rot(inp_lid.k_rot);993 }994}995 996bool llm_graph_input_dsv4::can_reuse(const llm_graph_params & params) {997 const auto * mctx = static_cast<const llama_kv_cache_dsv4_context *>(params.mctx);998 999 this->mctx = mctx;1000 inp_raw->mctx = mctx->get_raw();1001 1002 bool res = true;1003 1004 const auto & plan_csa = mctx->get_csa_plan(params.ubatch);1005 const auto & plan_hca = mctx->get_hca_plan(params.ubatch);1006 const auto & plan_lid = mctx->get_lid_plan(params.ubatch);1007 const int64_t n_stream = plan_csa.n_stream;1008 1009 const auto * raw_ctx = mctx->get_raw();1010 inp_raw->mctx = raw_ctx;1011 1012 if (inp_raw->self_k_idxs && inp_raw->self_k_idxs->buffer) {1013 res &= inp_raw->self_k_idxs->ne[0] == raw_ctx->get_n_write();1014 }1015 if (inp_raw->self_kq_mask && inp_raw->self_kq_mask->buffer) {1016 res &= dsv4_can_reuse_raw_kq_mask(inp_raw->self_kq_mask, raw_ctx, params.ubatch, n_stream);1017 }1018 1019 res &= dsv4_can_reuse_comp_input(inp_csa, plan_csa, params.ubatch.n_tokens, n_stream);1020 res &= dsv4_can_reuse_comp_input(inp_hca, plan_hca, params.ubatch.n_tokens, n_stream);1021 res &= dsv4_can_reuse_comp_input(inp_lid, plan_lid, params.ubatch.n_tokens, n_stream);1022 1023 return res;1024}1025 1026void llm_graph_input_attn_cross::set_input(const llama_ubatch * ubatch) {1027 GGML_ASSERT(cross_kq_mask);1028 1029 const int64_t n_enc = cross_kq_mask->ne[0];1030 const int64_t n_tokens = ubatch->n_tokens;1031 1032 GGML_ASSERT(ggml_backend_buffer_is_host(cross_kq_mask->buffer));1033 GGML_ASSERT(!ubatch->equal_seqs()); // TODO: use ubatch->n_seqs instead of failing1034 1035 const auto fill_mask = [&](auto * data) {1036 using T = std::remove_reference_t<decltype(*data)>;1037 for (int i = 0; i < n_tokens; ++i) {1038 GGML_ASSERT(!cross->seq_ids_enc.empty() && "llama_encode must be called first");1039 for (int j = 0; j < n_enc; ++j) {1040 float f = -INFINITY;1041 1042 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {1043 const llama_seq_id seq_id = ubatch->seq_id[i][s];1044 1045 if (cross->seq_ids_enc[j].find(seq_id) != cross->seq_ids_enc[j].end()) {1046 f = 0.0f;1047 }1048 }1049 1050 data[i*n_enc + j] = llama_cast<T>(f);1051 }1052 }1053 };1054 1055 if (cross_kq_mask->type == GGML_TYPE_F16) {1056 fill_mask((ggml_fp16_t *) cross_kq_mask->data);1057 } else {1058 fill_mask((float *) cross_kq_mask->data);1059 }1060}1061 1062void llm_graph_input_mem_hybrid::set_input(const llama_ubatch * ubatch) {1063 mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);1064 mctx->get_attn()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch);1065 1066 mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);1067 1068 if (inp_attn->self_k_rot) {1069 mctx->get_attn()->set_input_k_rot(inp_attn->self_k_rot);1070 }1071 1072 if (inp_attn->self_v_rot) {1073 mctx->get_attn()->set_input_v_rot(inp_attn->self_v_rot);1074 }1075 1076 const int64_t n_rs = mctx->get_recr()->get_n_rs();1077 1078 if (inp_rs->s_copy) {1079 GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));1080 int32_t * data = (int32_t *) inp_rs->s_copy->data;1081 1082 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n1083 for (uint32_t i = 0; i < n_rs; ++i) {1084 data[i] = mctx->get_recr()->s_copy(i);1085 }1086 }1087}1088 1089bool llm_graph_input_mem_hybrid::can_reuse(const llm_graph_params & params) {1090 const auto * mctx = static_cast<const llama_memory_hybrid_context *>(params.mctx);1091 1092 this->mctx = mctx;1093 1094 bool res = true;1095 1096 res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;1097 //res &= inp_attn->self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there1098 1099 res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams);1100 1101 res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();1102 1103 res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;1104 res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;1105 1106 res &= inp_rs->head == mctx->get_recr()->get_head();1107 res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();1108 1109 return res;1110}1111 1112// TODO: Hybrid input classes are a bit redundant.1113// Instead of creating a hybrid input, the graph can simply create 2 separate inputs.1114// Refactoring is required in the future.1115void llm_graph_input_mem_hybrid_k::set_input(const llama_ubatch * ubatch) {1116 mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);1117 1118 mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);1119 1120 const int64_t n_rs = mctx->get_recr()->get_n_rs();1121 1122 if (inp_rs->s_copy) {1123 GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));1124 int32_t * data = (int32_t *) inp_rs->s_copy->data;1125 1126 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n1127 for (uint32_t i = 0; i < n_rs; ++i) {1128 data[i] = mctx->get_recr()->s_copy(i);1129 }1130 }1131}1132 1133bool llm_graph_input_mem_hybrid_k::can_reuse(const llm_graph_params & params) {1134 const auto * mctx = static_cast<const llama_memory_hybrid_context *>(params.mctx);1135 1136 this->mctx = mctx;1137 1138 bool res = true;1139 1140 res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;1141 1142 res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams);1143 1144 res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();1145 1146 res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;1147 res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;1148 1149 res &= inp_rs->head == mctx->get_recr()->get_head();1150 res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();1151 1152 return res;1153}1154 1155void llm_graph_input_mem_hybrid_iswa::set_input(const llama_ubatch * ubatch) {1156 const auto * attn_ctx = mctx->get_attn();1157 1158 // base tensors may not be allocated if there are no non-SWA attention layers1159 if (inp_attn->self_k_idxs && inp_attn->self_k_idxs->buffer) {1160 attn_ctx->get_base()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);1161 attn_ctx->get_base()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch);1162 }1163 1164 if (inp_attn->self_kq_mask && inp_attn->self_kq_mask->buffer) {1165 attn_ctx->get_base()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);1166 }1167 1168 // swa tensors may not be allocated if there are no SWA attention layers1169 if (inp_attn->self_k_idxs_swa && inp_attn->self_k_idxs_swa->buffer) {1170 attn_ctx->get_swa()->set_input_k_idxs(inp_attn->self_k_idxs_swa, ubatch);1171 attn_ctx->get_swa()->set_input_v_idxs(inp_attn->self_v_idxs_swa, ubatch);1172 }1173 1174 if (inp_attn->self_kq_mask_swa && inp_attn->self_kq_mask_swa->buffer) {1175 attn_ctx->get_swa()->set_input_kq_mask(inp_attn->self_kq_mask_swa, ubatch, cparams.causal_attn);1176 }1177 1178 if (inp_attn->self_k_rot) {1179 attn_ctx->get_base()->set_input_k_rot(inp_attn->self_k_rot);1180 }1181 1182 if (inp_attn->self_v_rot) {1183 attn_ctx->get_base()->set_input_v_rot(inp_attn->self_v_rot);1184 }1185 1186 if (inp_attn->self_k_rot_swa) {1187 attn_ctx->get_swa()->set_input_k_rot(inp_attn->self_k_rot_swa);1188 }1189 1190 if (inp_attn->self_v_rot_swa) {1191 attn_ctx->get_swa()->set_input_v_rot(inp_attn->self_v_rot_swa);1192 }1193 1194 const int64_t n_rs = mctx->get_recr()->get_n_rs();1195 1196 if (inp_rs->s_copy) {1197 GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));1198 int32_t * data = (int32_t *) inp_rs->s_copy->data;1199 1200 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n