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