echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0604
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 8#include "llama-kv-cache.h"9#include "llama-kv-cache-iswa.h"10#include "llama-memory-hybrid.h"11#include "llama-memory-hybrid-iswa.h"12#include "llama-memory-recurrent.h"13 14#include <cassert>15#include <cmath>16#include <cstring>17#include <numeric>18#include <sstream>19#include <unordered_set>20 21// dedup helpers22 23static ggml_tensor * build_attn_inp_kq_mask(24 ggml_context * ctx,25 const llama_kv_cache_context * mctx,26 const llama_ubatch & ubatch,27 const llama_cparams & cparams) {28 const auto n_kv = mctx->get_n_kv();29 const auto n_tokens = ubatch.n_tokens;30 const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq;31 32 ggml_tensor * res = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, n_kv, n_tokens/n_stream, 1, n_stream);33 ggml_set_input(res);34 ggml_set_name(res, "attn_inp_kq_mask");35 36 return res;37}38 39static bool can_reuse_kq_mask(40 ggml_tensor * kq_mask,41 const llama_kv_cache_context * mctx,42 const llama_ubatch & ubatch,43 const llama_cparams & cparams) {44 const auto n_kv = mctx->get_n_kv();45 const auto n_tokens = ubatch.n_tokens;46 const auto n_stream = cparams.kv_unified ? 1 : ubatch.n_seqs_unq;47 48 bool res = true;49 50 res &= (kq_mask->ne[0] == n_kv);51 res &= (kq_mask->ne[1] == n_tokens/n_stream);52 res &= (kq_mask->ne[2] == 1);53 res &= (kq_mask->ne[3] == n_stream);54 55 return res;56}57 58// impl59 60static ggml_tensor * ggml_mul_mat_aux(61 ggml_context * ctx,62 ggml_tensor * cur,63 ggml_tensor * rot) {64 const auto n = rot->ne[0];65 66 ggml_tensor * res;67 68 res = ggml_reshape_2d(ctx, cur, n, ggml_nelements(cur)/n);69 res = ggml_mul_mat (ctx, rot, res);70 res = ggml_reshape_4d(ctx, res, cur->ne[0], cur->ne[1], cur->ne[2], cur->ne[3]);71 72 return res;73}74 75void llm_graph_input_embd::set_input(const llama_ubatch * ubatch) {76 if (ubatch->token) {77 const int64_t n_tokens = ubatch->n_tokens;78 79 ggml_backend_tensor_set(tokens, ubatch->token, 0, n_tokens*ggml_element_size(tokens));80 }81 82 if (ubatch->embd) {83 GGML_ASSERT(n_embd == embd->ne[0]);84 85 const int64_t n_tokens = ubatch->n_tokens;86 87 ggml_backend_tensor_set(embd, ubatch->embd, 0, n_tokens*n_embd*ggml_element_size(embd));88 }89}90 91bool llm_graph_input_embd::can_reuse(const llm_graph_params & params) {92 bool res = true;93 94 res &= (!params.ubatch.token) || (tokens && tokens->ne[0] == params.ubatch.n_tokens);95 res &= (!params.ubatch.embd) || (embd && embd->ne[1] == params.ubatch.n_tokens);96 97 return res;98}99 100void llm_graph_input_pos::set_input(const llama_ubatch * ubatch) {101 if (ubatch->pos && pos) {102 const int64_t n_tokens = ubatch->n_tokens;103 104 if (ubatch->token && n_pos_per_embd == 4) {105 // in case we're using M-RoPE with text tokens, convert the 1D positions to 4D106 // the 3 first dims are the same, and 4th dim is all 0107 std::vector<llama_pos> pos_data(n_tokens*n_pos_per_embd);108 // copy the first dimension109 for (int i = 0; i < n_tokens; ++i) {110 pos_data[ i] = ubatch->pos[i];111 pos_data[ n_tokens + i] = ubatch->pos[i];112 pos_data[2 * n_tokens + i] = ubatch->pos[i];113 pos_data[3 * n_tokens + i] = 0; // 4th dim is 0114 }115 ggml_backend_tensor_set(pos, pos_data.data(), 0, pos_data.size()*ggml_element_size(pos));116 } else {117 ggml_backend_tensor_set(pos, ubatch->pos, 0, n_tokens*n_pos_per_embd*ggml_element_size(pos));118 }119 }120}121 122bool llm_graph_input_pos::can_reuse(const llm_graph_params & params) {123 bool res = true;124 125 res &= pos->ne[0] == params.ubatch.n_tokens*n_pos_per_embd;126 127 return res;128}129 130void llm_graph_input_attn_temp::set_input(const llama_ubatch * ubatch) {131 if (ubatch->pos && attn_scale) {132 const int64_t n_tokens = ubatch->n_tokens;133 134 GGML_ASSERT(f_attn_temp_scale != 0.0f);135 GGML_ASSERT(n_attn_temp_floor_scale != 0);136 137 std::vector<float> attn_scale_data(n_tokens, 0.0f);138 for (int i = 0; i < n_tokens; ++i) {139 const float pos = ubatch->pos[i];140 attn_scale_data[i] = std::log(141 std::floor((pos + f_attn_temp_offset) / n_attn_temp_floor_scale) + 1.0142 ) * f_attn_temp_scale + 1.0;143 }144 145 ggml_backend_tensor_set(attn_scale, attn_scale_data.data(), 0, n_tokens*ggml_element_size(attn_scale));146 }147}148 149void llm_graph_input_pos_bucket::set_input(const llama_ubatch * ubatch) {150 if (pos_bucket) {151 const int64_t n_tokens = ubatch->n_tokens;152 153 GGML_ASSERT(ggml_backend_buffer_is_host(pos_bucket->buffer));154 GGML_ASSERT(!ubatch->equal_seqs()); // TODO: use ubatch->n_seqs instead of failing155 156 int32_t * data = (int32_t *) pos_bucket->data;157 158 for (int j = 0; j < n_tokens; ++j) {159 for (int i = 0; i < n_tokens; ++i) {160 data[j*n_tokens + i] = llama_relative_position_bucket(ubatch->pos[i], ubatch->pos[j], hparams.n_rel_attn_bkts, true);161 }162 }163 }164}165 166void llm_graph_input_pos_bucket_kv::set_input(const llama_ubatch * ubatch) {167 if (pos_bucket) {168 mctx->set_input_pos_bucket(pos_bucket, ubatch);169 }170}171 172void llm_graph_input_out_ids::set_input(const llama_ubatch * ubatch) {173 GGML_ASSERT(out_ids);174 175 const int64_t n_tokens = ubatch->n_tokens;176 177 GGML_ASSERT(ggml_backend_buffer_is_host(out_ids->buffer));178 int32_t * data = (int32_t *) out_ids->data;179 180 if (n_outputs == n_tokens) {181 for (int i = 0; i < n_tokens; ++i) {182 data[i] = i;183 }184 185 return;186 }187 188 GGML_ASSERT(ubatch->output);189 190 int n_outputs = 0;191 192 for (int i = 0; i < n_tokens; ++i) {193 if (ubatch->output[i]) {194 data[n_outputs++] = i;195 }196 }197}198 199bool llm_graph_input_out_ids::can_reuse(const llm_graph_params & params) {200 bool res = true;201 202 res &= n_outputs == params.n_outputs;203 204 return res;205}206 207void llm_graph_input_mean::set_input(const llama_ubatch * ubatch) {208 if (cparams.embeddings &&209 (cparams.pooling_type == LLAMA_POOLING_TYPE_MEAN ||210 cparams.pooling_type == LLAMA_POOLING_TYPE_RANK )) {211 212 const int64_t n_tokens = ubatch->n_tokens;213 const int64_t n_seq_tokens = ubatch->n_seq_tokens;214 const int64_t n_seqs_unq = ubatch->n_seqs_unq;215 216 GGML_ASSERT(mean);217 GGML_ASSERT(ggml_backend_buffer_is_host(mean->buffer));218 219 float * data = (float *) mean->data;220 memset(mean->data, 0, n_tokens*n_seqs_unq*ggml_element_size(mean));221 222 std::vector<uint64_t> sums(n_seqs_unq, 0);223 for (int i = 0; i < n_tokens; i += n_seq_tokens) {224 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {225 const llama_seq_id seq_id = ubatch->seq_id[i][s];226 const int32_t seq_idx = ubatch->seq_idx[seq_id];227 228 sums[seq_idx] += ubatch->n_seq_tokens;229 }230 }231 232 std::vector<float> div(n_seqs_unq, 0.0f);233 for (int s = 0; s < n_seqs_unq; ++s) {234 const uint64_t sum = sums[s];235 if (sum > 0) {236 div[s] = 1.0f/float(sum);237 }238 }239 240 for (int i = 0; i < n_tokens; i += n_seq_tokens) {241 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {242 const llama_seq_id seq_id = ubatch->seq_id[i][s];243 const int32_t seq_idx = ubatch->seq_idx[seq_id];244 245 for (int j = 0; j < n_seq_tokens; ++j) {246 data[seq_idx*n_tokens + i + j] = div[seq_idx];247 }248 }249 }250 }251}252 253void llm_graph_input_cls::set_input(const llama_ubatch * ubatch) {254 const int64_t n_tokens = ubatch->n_tokens;255 const int64_t n_seqs_unq = ubatch->n_seqs_unq;256 257 if (cparams.embeddings && (258 cparams.pooling_type == LLAMA_POOLING_TYPE_CLS ||259 cparams.pooling_type == LLAMA_POOLING_TYPE_RANK ||260 cparams.pooling_type == LLAMA_POOLING_TYPE_LAST261 )) {262 GGML_ASSERT(cls);263 GGML_ASSERT(ggml_backend_buffer_is_host(cls->buffer));264 265 uint32_t * data = (uint32_t *) cls->data;266 memset(cls->data, 0, n_seqs_unq*ggml_element_size(cls));267 268 std::vector<int> target_pos(n_seqs_unq, -1);269 std::vector<int> target_row(n_seqs_unq, -1);270 271 const bool last = (272 cparams.pooling_type == LLAMA_POOLING_TYPE_LAST ||273 (cparams.pooling_type == LLAMA_POOLING_TYPE_RANK && (arch == LLM_ARCH_QWEN3 || arch == LLM_ARCH_QWEN3VL)) // qwen3 reranking & embedding models use last token274 );275 276 for (int i = 0; i < n_tokens; ++i) {277 const llama_pos pos = ubatch->pos[i];278 279 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {280 const llama_seq_id seq_id = ubatch->seq_id[i][s];281 const int32_t seq_idx = ubatch->seq_idx[seq_id];282 283 if (284 (target_pos[seq_idx] == -1) ||285 ( last && pos >= target_pos[seq_idx]) ||286 (!last && pos < target_pos[seq_idx])287 ) {288 target_pos[seq_idx] = pos;289 target_row[seq_idx] = i;290 }291 }292 }293 294 for (int s = 0; s < n_seqs_unq; ++s) {295 if (target_row[s] >= 0) {296 data[s] = target_row[s];297 }298 }299 }300}301 302void llm_graph_input_rs::set_input(const llama_ubatch * ubatch) {303 GGML_UNUSED(ubatch);304 305 const int64_t n_rs = mctx->get_n_rs();306 307 if (s_copy) {308 GGML_ASSERT(ggml_backend_buffer_is_host(s_copy->buffer));309 int32_t * data = (int32_t *) s_copy->data;310 311 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n312 for (uint32_t i = 0; i < n_rs; ++i) {313 data[i] = mctx->s_copy(i);314 }315 }316}317 318bool llm_graph_input_rs::can_reuse(const llm_graph_params & params) {319 const auto * mctx = static_cast<const llama_memory_recurrent_context *>(params.mctx);320 321 this->mctx = mctx;322 323 bool res = true;324 325 res &= s_copy->ne[0] == mctx->get_n_rs();326 327 res &= s_copy_main->ne[0] == params.ubatch.n_seqs;328 res &= s_copy_extra->ne[0] == mctx->get_n_rs() - params.ubatch.n_seqs;329 330 res &= head == mctx->get_head();331 res &= rs_z == mctx->get_rs_z();332 333 return res;334}335 336void llm_graph_input_cross_embd::set_input(const llama_ubatch * ubatch) {337 GGML_UNUSED(ubatch);338 339 if (cross_embd && !cross->v_embd.empty()) {340 assert(cross_embd->type == GGML_TYPE_F32);341 342 ggml_backend_tensor_set(cross_embd, cross->v_embd.data(), 0, ggml_nbytes(cross_embd));343 }344}345 346static void print_mask(const float * data, int64_t n_tokens, int64_t n_kv, int64_t n_swa, llama_swa_type swa_type) {347 LLAMA_LOG_DEBUG("%s: === Attention mask ===\n", __func__);348 const char * swa_type_str = "unknown";349 350 switch (swa_type) {351 case LLAMA_SWA_TYPE_NONE: swa_type_str = "LLAMA_SWA_TYPE_NONE"; break;352 case LLAMA_SWA_TYPE_STANDARD: swa_type_str = "LLAMA_SWA_TYPE_STANDARD"; break;353 case LLAMA_SWA_TYPE_CHUNKED: swa_type_str = "LLAMA_SWA_TYPE_CHUNKED"; break;354 case LLAMA_SWA_TYPE_SYMMETRIC: swa_type_str = "LLAMA_SWA_TYPE_SYMMETRIC"; break;355 };356 357 LLAMA_LOG_DEBUG("%s: n_swa : %d, n_kv: %d, swq_type: %s\n", __func__, (int)n_swa, (int)n_kv, swa_type_str);358 LLAMA_LOG_DEBUG("%s: '0' = can attend, 'โ' = masked\n", __func__);359 LLAMA_LOG_DEBUG("%s: Rows = query tokens, Columns = key/value tokens\n\n", __func__);360 361 LLAMA_LOG_DEBUG(" ");362 for (int j = 0; j < std::min((int64_t)20, n_kv); ++j) {363 LLAMA_LOG_DEBUG("%2d", j);364 }365 LLAMA_LOG_DEBUG("\n");366 367 for (int i = 0; i < std::min((int64_t)20, n_tokens); ++i) {368 LLAMA_LOG_DEBUG(" %2d ", i);369 for (int j = 0; j < std::min((int64_t)20, n_kv); ++j) {370 float val = data[i * n_kv + j];371 if (val == -INFINITY) {372 LLAMA_LOG_DEBUG(" โ");373 } else {374 LLAMA_LOG_DEBUG(" 0");375 }376 }377 LLAMA_LOG_DEBUG("\n");378 }379}380 381void llm_graph_input_attn_no_cache::set_input(const llama_ubatch * ubatch) {382 const int64_t n_kv = ubatch->n_tokens;383 const int64_t n_tokens = ubatch->n_tokens;384 385 const auto fill_mask = [&](float * data, int n_swa, llama_swa_type swa_type) {386 for (int i1 = 0; i1 < n_tokens; ++i1) {387 const llama_seq_id s1 = ubatch->seq_id[i1][0];388 const llama_pos p1 = ubatch->pos[i1];389 390 const uint64_t idst = i1*n_kv;391 392 for (int i0 = 0; i0 < n_tokens; ++i0) {393 const llama_seq_id s0 = ubatch->seq_id[i0][0];394 const llama_pos p0 = ubatch->pos[i0];395 396 // mask different sequences397 if (s0 != s1) {398 continue;399 }400 401 // mask future tokens402 if (cparams.causal_attn && p0 > p1) {403 continue;404 }405 406 // apply SWA if any407 if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {408 continue;409 }410 411 data[idst + i0] = hparams.use_alibi ? -std::abs(p0 - p1) : 0.0f;412 }413 }414 };415 416 {417 GGML_ASSERT(self_kq_mask);418 GGML_ASSERT(ggml_backend_buffer_is_host(self_kq_mask->buffer));419 420 float * data = (float *) self_kq_mask->data;421 422 std::fill(data, data + ggml_nelements(self_kq_mask), -INFINITY);423 424 fill_mask(data, 0, LLAMA_SWA_TYPE_NONE);425 426 if (debug) {427 print_mask(data, n_tokens, n_kv, 0, LLAMA_SWA_TYPE_NONE);428 }429 }430 431 if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) {432 GGML_ASSERT(self_kq_mask_swa);433 GGML_ASSERT(ggml_backend_buffer_is_host(self_kq_mask_swa->buffer));434 435 float * data = (float *) self_kq_mask_swa->data;436 437 std::fill(data, data + ggml_nelements(self_kq_mask_swa), -INFINITY);438 439 fill_mask(data, hparams.n_swa, hparams.swa_type);440 441 if (debug) {442 print_mask(data, n_tokens, n_kv, hparams.n_swa, hparams.swa_type);443 }444 }445}446 447void llm_graph_input_attn_kv::set_input(const llama_ubatch * ubatch) {448 mctx->set_input_k_idxs(self_k_idxs, ubatch);449 mctx->set_input_v_idxs(self_v_idxs, ubatch);450 451 mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);452 453 if (self_k_rot) {454 mctx->set_input_k_rot(self_k_rot);455 }456 457 if (self_v_rot) {458 mctx->set_input_v_rot(self_v_rot);459 }460}461 462bool llm_graph_input_attn_kv::can_reuse(const llm_graph_params & params) {463 const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);464 465 this->mctx = mctx;466 467 bool res = true;468 469 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;470 //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there471 472 res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams);473 474 return res;475}476 477void llm_graph_input_attn_k::set_input(const llama_ubatch * ubatch) {478 mctx->set_input_k_idxs(self_k_idxs, ubatch);479 480 mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);481}482 483bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) {484 const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);485 486 this->mctx = mctx;487 488 bool res = true;489 490 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;491 492 res &= can_reuse_kq_mask(self_kq_mask, mctx, params.ubatch, params.cparams);493 494 return res;495}496 497void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) {498 mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);499 mctx->get_base()->set_input_v_idxs(self_v_idxs, ubatch);500 501 mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);502 503 mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);504 mctx->get_swa()->set_input_v_idxs(self_v_idxs_swa, ubatch);505 506 mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);507 508 if (self_k_rot) {509 mctx->get_base()->set_input_k_rot(self_k_rot);510 }511 512 if (self_v_rot) {513 mctx->get_base()->set_input_v_rot(self_v_rot);514 }515 516 if (self_k_rot_swa) {517 mctx->get_swa()->set_input_k_rot(self_k_rot_swa);518 }519 520 if (self_v_rot_swa) {521 mctx->get_swa()->set_input_v_rot(self_v_rot_swa);522 }523}524 525bool llm_graph_input_attn_kv_iswa::can_reuse(const llm_graph_params & params) {526 const auto * mctx = static_cast<const llama_kv_cache_iswa_context *>(params.mctx);527 528 this->mctx = mctx;529 530 bool res = true;531 532 res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;533 //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there534 535 res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;536 //res &= self_v_idxs_swa->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there537 538 res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);539 res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);540 541 return res;542}543 544void llm_graph_input_attn_cross::set_input(const llama_ubatch * ubatch) {545 GGML_ASSERT(cross_kq_mask);546 547 const int64_t n_enc = cross_kq_mask->ne[0];548 const int64_t n_tokens = ubatch->n_tokens;549 550 GGML_ASSERT(ggml_backend_buffer_is_host(cross_kq_mask->buffer));551 GGML_ASSERT(!ubatch->equal_seqs()); // TODO: use ubatch->n_seqs instead of failing552 553 float * data = (float *) cross_kq_mask->data;554 555 for (int i = 0; i < n_tokens; ++i) {556 GGML_ASSERT(!cross->seq_ids_enc.empty() && "llama_encode must be called first");557 for (int j = 0; j < n_enc; ++j) {558 float f = -INFINITY;559 560 for (int s = 0; s < ubatch->n_seq_id[i]; ++s) {561 const llama_seq_id seq_id = ubatch->seq_id[i][s];562 563 if (cross->seq_ids_enc[j].find(seq_id) != cross->seq_ids_enc[j].end()) {564 f = 0.0f;565 }566 }567 568 data[i*n_enc + j] = f;569 }570 }571}572 573void llm_graph_input_mem_hybrid::set_input(const llama_ubatch * ubatch) {574 mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);575 mctx->get_attn()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch);576 577 mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);578 579 if (inp_attn->self_k_rot) {580 mctx->get_attn()->set_input_k_rot(inp_attn->self_k_rot);581 }582 583 if (inp_attn->self_v_rot) {584 mctx->get_attn()->set_input_v_rot(inp_attn->self_v_rot);585 }586 587 const int64_t n_rs = mctx->get_recr()->get_n_rs();588 589 if (inp_rs->s_copy) {590 GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));591 int32_t * data = (int32_t *) inp_rs->s_copy->data;592 593 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n594 for (uint32_t i = 0; i < n_rs; ++i) {595 data[i] = mctx->get_recr()->s_copy(i);596 }597 }598}599 600bool llm_graph_input_mem_hybrid::can_reuse(const llm_graph_params & params) {601 const auto * mctx = static_cast<const llama_memory_hybrid_context *>(params.mctx);602 603 this->mctx = mctx;604 605 bool res = true;606 607 res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;608 //res &= inp_attn->self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there609 610 res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams);611 612 res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();613 614 res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;615 res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;616 617 res &= inp_rs->head == mctx->get_recr()->get_head();618 res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();619 620 return res;621}622 623// TODO: Hybrid input classes are a bit redundant.624// Instead of creating a hybrid input, the graph can simply create 2 separate inputs.625// Refactoring is required in the future.626void llm_graph_input_mem_hybrid_k::set_input(const llama_ubatch * ubatch) {627 mctx->get_attn()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);628 629 mctx->get_attn()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);630 631 const int64_t n_rs = mctx->get_recr()->get_n_rs();632 633 if (inp_rs->s_copy) {634 GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));635 int32_t * data = (int32_t *) inp_rs->s_copy->data;636 637 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n638 for (uint32_t i = 0; i < n_rs; ++i) {639 data[i] = mctx->get_recr()->s_copy(i);640 }641 }642}643 644bool llm_graph_input_mem_hybrid_k::can_reuse(const llm_graph_params & params) {645 const auto * mctx = static_cast<const llama_memory_hybrid_context *>(params.mctx);646 647 this->mctx = mctx;648 649 bool res = true;650 651 res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;652 653 res &= can_reuse_kq_mask(inp_attn->self_kq_mask, mctx->get_attn(), params.ubatch, params.cparams);654 655 res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();656 657 res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;658 res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;659 660 res &= inp_rs->head == mctx->get_recr()->get_head();661 res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();662 663 return res;664}665 666void llm_graph_input_mem_hybrid_iswa::set_input(const llama_ubatch * ubatch) {667 const auto * attn_ctx = mctx->get_attn();668 669 // base tensors may not be allocated if there are no non-SWA attention layers670 if (inp_attn->self_k_idxs && inp_attn->self_k_idxs->buffer) {671 attn_ctx->get_base()->set_input_k_idxs(inp_attn->self_k_idxs, ubatch);672 attn_ctx->get_base()->set_input_v_idxs(inp_attn->self_v_idxs, ubatch);673 674 attn_ctx->get_base()->set_input_kq_mask(inp_attn->self_kq_mask, ubatch, cparams.causal_attn);675 }676 677 // swa tensors may not be allocated if there are no SWA attention layers678 if (inp_attn->self_k_idxs_swa && inp_attn->self_k_idxs_swa->buffer) {679 attn_ctx->get_swa()->set_input_k_idxs(inp_attn->self_k_idxs_swa, ubatch);680 attn_ctx->get_swa()->set_input_v_idxs(inp_attn->self_v_idxs_swa, ubatch);681 682 attn_ctx->get_swa()->set_input_kq_mask(inp_attn->self_kq_mask_swa, ubatch, cparams.causal_attn);683 }684 685 if (inp_attn->self_k_rot) {686 attn_ctx->get_base()->set_input_k_rot(inp_attn->self_k_rot);687 }688 689 if (inp_attn->self_v_rot) {690 attn_ctx->get_base()->set_input_v_rot(inp_attn->self_v_rot);691 }692 693 if (inp_attn->self_k_rot_swa) {694 attn_ctx->get_swa()->set_input_k_rot(inp_attn->self_k_rot_swa);695 }696 697 if (inp_attn->self_v_rot_swa) {698 attn_ctx->get_swa()->set_input_v_rot(inp_attn->self_v_rot_swa);699 }700 701 const int64_t n_rs = mctx->get_recr()->get_n_rs();702 703 if (inp_rs->s_copy) {704 GGML_ASSERT(ggml_backend_buffer_is_host(inp_rs->s_copy->buffer));705 int32_t * data = (int32_t *) inp_rs->s_copy->data;706 707 // assuming copy destinations ALWAYS happen ONLY on the cells between head and head+n708 for (uint32_t i = 0; i < n_rs; ++i) {709 data[i] = mctx->get_recr()->s_copy(i);710 }711 }712}713 714bool llm_graph_input_mem_hybrid_iswa::can_reuse(const llm_graph_params & params) {715 const auto * mctx = static_cast<const llama_memory_hybrid_iswa_context *>(params.mctx);716 717 this->mctx = mctx;718 719 bool res = true;720 721 const auto * attn_ctx = mctx->get_attn();722 723 // base tensors may not be allocated if there are no non-SWA attention layers724 if (inp_attn->self_k_idxs && inp_attn->self_k_idxs->buffer) {725 res &= inp_attn->self_k_idxs->ne[0] == params.ubatch.n_tokens;726 //res &= inp_attn->self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there727 728 res &= can_reuse_kq_mask(inp_attn->self_kq_mask, attn_ctx->get_base(), params.ubatch, params.cparams);729 }730 731 // swa tensors may not be allocated if there are no SWA attention layers732 if (inp_attn->self_k_idxs_swa && inp_attn->self_k_idxs_swa->buffer) {733 res &= inp_attn->self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;734 //res &= inp_attn->self_v_idxs_swa->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there735 736 res &= can_reuse_kq_mask(inp_attn->self_kq_mask_swa, attn_ctx->get_swa(), params.ubatch, params.cparams);737 }738 739 res &= inp_rs->s_copy->ne[0] == mctx->get_recr()->get_n_rs();740 741 res &= inp_rs->s_copy_main->ne[0] == params.ubatch.n_seqs;742 res &= inp_rs->s_copy_extra->ne[0] == mctx->get_recr()->get_n_rs() - params.ubatch.n_seqs;743 744 res &= inp_rs->head == mctx->get_recr()->get_head();745 res &= inp_rs->rs_z == mctx->get_recr()->get_rs_z();746 747 return res;748}749 750void llm_graph_input_sampling::set_input(const llama_ubatch * ubatch) {751 // set the inputs only for the active samplers in the current ubatch752 std::unordered_set<llama_seq_id> active_samplers;753 for (uint32_t i = 0; i < ubatch->n_tokens; i++) {754 if (ubatch->output[i]) {755 llama_seq_id seq_id = ubatch->seq_id[i][0];756 active_samplers.insert(seq_id);757 }758 }759 760 for (auto seq_id : active_samplers) {761 if (samplers.find(seq_id) == samplers.end()) {762 continue;763 }764 765 auto & sampler = samplers[seq_id];766 767 if (sampler->iface->backend_set_input) {768 sampler->iface->backend_set_input(sampler);769 }770 }771}772 773bool llm_graph_input_sampling::can_reuse(const llm_graph_params & params) {774 if (samplers.size() != params.samplers.size()) {775 return false;776 }777 778 for (const auto & [seq_id, sampler] : params.samplers) {779 if (samplers[seq_id] != sampler) {780 return false;781 }782 }783 784 return true;785}786 787//788// llm_graph_result789//790 791llm_graph_result::llm_graph_result(int64_t max_nodes) : max_nodes(max_nodes) {792 reset();793 794 const char * LLAMA_GRAPH_RESULT_DEBUG = getenv("LLAMA_GRAPH_RESULT_DEBUG");795 debug = LLAMA_GRAPH_RESULT_DEBUG ? atoi(LLAMA_GRAPH_RESULT_DEBUG) : 0;796}797 798int64_t llm_graph_result::get_max_nodes() const {799 return max_nodes;800}801 802void llm_graph_result::reset() {803 t_inp_tokens = nullptr;804 t_inp_embd = nullptr;805 t_logits = nullptr;806 t_embd = nullptr;807 t_embd_pooled = nullptr;808 t_sampled.clear();809 t_sampled_probs.clear();810 t_sampled_logits.clear();811 t_candidates.clear();812 813 params = {};814 815 inputs.clear();816 817 buf_compute_meta.resize(ggml_tensor_overhead()*max_nodes + ggml_graph_overhead_custom(max_nodes, false));818 819 ggml_init_params params = {820 /*.mem_size =*/ buf_compute_meta.size(),821 /*.mem_buffer =*/ buf_compute_meta.data(),822 /*.no_alloc =*/ true,823 };824 825 ctx_compute.reset(ggml_init(params));826 827 gf = ggml_new_graph_custom(ctx_compute.get(), max_nodes, false);828}829 830void llm_graph_result::set_inputs(const llama_ubatch * ubatch) {831 for (auto & input : inputs) {832 input->set_input(ubatch);833 }834}835 836void llm_graph_result::set_outputs() {837 if (t_logits != nullptr) {838 ggml_set_output(t_logits);839 }840 if (t_embd != nullptr) {841 ggml_set_output(t_embd);842 }843 if (t_embd_pooled != nullptr) {844 ggml_set_output(t_embd_pooled);845 }846 for (auto & [seq_id, t] : t_sampled) {847 if (t != nullptr) {848 ggml_set_output(t);849 }850 }851 for (auto & [seq_id, t] : t_sampled_probs) {852 if (t != nullptr) {853 ggml_set_output(t);854 }855 }856 for (auto & [seq_id, t] : t_sampled_logits) {857 if (t != nullptr) {858 ggml_set_output(t);859 }860 }861 for (auto & [seq_id, t] : t_candidates) {862 if (t != nullptr) {863 ggml_set_output(t);864 }865 }866}867 868bool llm_graph_result::can_reuse(const llm_graph_params & params) {869 if (!this->params.allow_reuse(params)) {870 if (debug > 1) {871 LLAMA_LOG_DEBUG("%s: cannot reuse graph due to incompatible graph parameters\n", __func__);872 }873 874 return false;875 }876 877 if (debug > 1) {878 LLAMA_LOG_DEBUG("%s: checking compatibility of %d inputs:\n", __func__, (int) inputs.size());879 }880 881 bool res = true;882 883 for (auto & input : inputs) {884 const bool cur = input->can_reuse(params);885 886 if (debug > 1) {887 LLAMA_LOG_DEBUG("%s: can_reuse = %d\n", "placeholder", cur);888 }889 890 res = res && cur;891 }892 893 if (debug > 0) {894 LLAMA_LOG_DEBUG("%s: can reuse graph = %d\n", __func__, res);895 }896 897 return res;898}899 900llm_graph_input_i * llm_graph_result::add_input(llm_graph_input_ptr input) {901 inputs.emplace_back(std::move(input));902 return inputs.back().get();903}904 905void llm_graph_result::set_params(const llm_graph_params & params) {906 this->params = params;907}908 909//910// llm_graph_context911//912 913llm_graph_context::llm_graph_context(const llm_graph_params & params) :914 arch (params.arch),915 hparams (params.hparams),916 cparams (params.cparams),917 ubatch (params.ubatch),918 n_embd (hparams.n_embd),919 n_layer (hparams.n_layer),920 n_rot (hparams.n_rot()),921 n_ctx (cparams.n_ctx),922 n_head (hparams.n_head()),923 n_head_kv (hparams.n_head_kv()),924 n_embd_head_k (hparams.n_embd_head_k()),925 n_embd_k_gqa (hparams.n_embd_k_gqa()),926 n_embd_head_v (hparams.n_embd_head_v()),927 n_embd_v_gqa (hparams.n_embd_v_gqa()),928 n_expert (hparams.n_expert),929 n_expert_used (cparams.warmup ? hparams.n_expert : hparams.n_expert_used),930 freq_base (cparams.rope_freq_base),931 freq_scale (cparams.rope_freq_scale),932 ext_factor (cparams.yarn_ext_factor),933 attn_factor (cparams.yarn_attn_factor),934 beta_fast (cparams.yarn_beta_fast),935 beta_slow (cparams.yarn_beta_slow),936 norm_eps (hparams.f_norm_eps),937 norm_rms_eps (hparams.f_norm_rms_eps),938 n_tokens (ubatch.n_tokens),939 n_outputs (params.n_outputs),940 n_ctx_orig (cparams.n_ctx_orig_yarn),941 pooling_type (cparams.pooling_type),942 rope_type (hparams.rope_type),943 sched (params.sched),944 backend_cpu (params.backend_cpu),945 cvec (params.cvec),946 loras (params.loras),947 mctx (params.mctx),948 cross (params.cross),949 samplers (params.samplers),950 cb_func (params.cb),951 res (params.res),952 ctx0 (res->get_ctx()),953 gf (res->get_gf()) {954 res->set_params(params);955 }956 957void llm_graph_context::cb(ggml_tensor * cur, const char * name, int il) const {958 if (cb_func) {959 cb_func(ubatch, cur, name, il);960 }961}962 963ggml_tensor * llm_graph_context::build_cvec(964 ggml_tensor * cur,965 int il) const {966 return cvec->apply_to(ctx0, cur, il);967}968 969ggml_tensor * llm_graph_context::build_lora_mm(970 ggml_tensor * w,971 ggml_tensor * cur,972 ggml_tensor * w_s) const {973 ggml_tensor * res = ggml_mul_mat(ctx0, w, cur);974 975 for (const auto & lora : *loras) {976 llama_adapter_lora_weight * lw = lora.first->get_weight(w);977 if (lw == nullptr) {978 continue;979 }980 981 const float adapter_scale = lora.second;982 const float scale = lw->get_scale(lora.first->alpha, adapter_scale);983 984 ggml_tensor * ab_cur = ggml_mul_mat(985 ctx0, lw->b,986 ggml_mul_mat(ctx0, lw->a, cur)987 );988 989 ab_cur = ggml_scale(ctx0, ab_cur, scale);990 res = ggml_add(ctx0, res, ab_cur);991 }992 993 if (w_s) {994 res = ggml_mul(ctx0, res, w_s);995 }996 997 return res;998}999 1000ggml_tensor * llm_graph_context::build_lora_mm_id(1001 ggml_tensor * w, // ggml_tensor * as1002 ggml_tensor * cur, // ggml_tensor * b1003 ggml_tensor * ids) const {1004 ggml_tensor * res = ggml_mul_mat_id(ctx0, w, cur, ids);1005 for (const auto & lora : *loras) {1006 llama_adapter_lora_weight * lw = lora.first->get_weight(w);1007 if (lw == nullptr) {1008 continue;1009 }1010 1011 const float alpha = lora.first->alpha;1012 const float rank = (float) lw->b->ne[0];1013 const float scale = alpha ? lora.second * alpha / rank : lora.second;1014 1015 ggml_tensor * ab_cur = ggml_mul_mat_id(1016 ctx0, lw->b,1017 ggml_mul_mat_id(ctx0, lw->a, cur, ids),1018 ids1019 );1020 1021 ab_cur = ggml_scale(ctx0, ab_cur, scale);1022 res = ggml_add(ctx0, res, ab_cur);1023 }1024 1025 return res;1026}1027 1028ggml_tensor * llm_graph_context::build_norm(1029 ggml_tensor * cur,1030 ggml_tensor * mw,1031 ggml_tensor * mb,1032 llm_norm_type type,1033 int il) const {1034 switch (type) {1035 case LLM_NORM: cur = ggml_norm (ctx0, cur, hparams.f_norm_eps); break;1036 case LLM_NORM_RMS: cur = ggml_rms_norm(ctx0, cur, hparams.f_norm_rms_eps); break;1037 case LLM_NORM_GROUP:1038 {1039 cur = ggml_reshape_3d(ctx0, cur, cur->ne[0], 1, cur->ne[1]);1040 cur = ggml_group_norm(ctx0, cur, hparams.n_norm_groups, hparams.f_norm_group_eps);1041 cur = ggml_reshape_2d(ctx0, cur, cur->ne[0], cur->ne[2]);1042 } break;1043 }1044 1045 if (mw || mb) {1046 cb(cur, "norm", il);1047 }1048 1049 if (mw) {1050 cur = ggml_mul(ctx0, cur, mw);1051 if (mb) {1052 cb(cur, "norm_w", il);1053 }1054 }1055 1056 if (mb) {1057 cur = ggml_add(ctx0, cur, mb);1058 }1059 1060 return cur;1061}1062 1063 1064llm_graph_qkv llm_graph_context::build_qkv(1065 const llama_layer & layer,1066 ggml_tensor * cur,1067 int64_t n_embd_head,1068 int64_t n_head,1069 int64_t n_head_kv,1070 int il) const {1071 const int64_t n_embd_q = n_embd_head * n_head;1072 const int64_t n_embd_kv = n_embd_head * n_head_kv;1073 1074 ggml_tensor * Qcur, * Kcur, * Vcur;1075 1076 if (layer.wqkv) {1077 // fused QKV path1078 ggml_tensor * qkv = build_lora_mm(layer.wqkv, cur, layer.wqkv_s);1079 cb(qkv, "wqkv", il);1080 if (layer.wqkv_b) {1081 qkv = ggml_add(ctx0, qkv, layer.wqkv_b);1082 cb(qkv, "wqkv_b", il);1083 }1084 if (hparams.f_clamp_kqv > 0.0f) {1085 qkv = ggml_clamp(ctx0, qkv, -hparams.f_clamp_kqv, hparams.f_clamp_kqv);1086 cb(qkv, "wqkv_clamped", il);1087 }1088 Qcur = ggml_view_3d(ctx0, qkv, n_embd_head, n_head, n_tokens,1089 ggml_row_size(qkv->type, n_embd_head), qkv->nb[1], 0);1090 Kcur = ggml_view_3d(ctx0, qkv, n_embd_head, n_head_kv, n_tokens,1091 ggml_row_size(qkv->type, n_embd_head), qkv->nb[1],1092 ggml_row_size(qkv->type, n_embd_q));1093 Vcur = ggml_view_3d(ctx0, qkv, n_embd_head, n_head_kv, n_tokens,1094 ggml_row_size(qkv->type, n_embd_head), qkv->nb[1],1095 ggml_row_size(qkv->type, n_embd_q + n_embd_kv));1096 } else {1097 // separate Q/K/V path1098 Qcur = build_lora_mm(layer.wq, cur, layer.wq_s);1099 cb(Qcur, "Qcur", il);1100 if (layer.wq_b) {1101 Qcur = ggml_add(ctx0, Qcur, layer.wq_b);1102 cb(Qcur, "Qcur", il);1103 }1104 if (hparams.f_clamp_kqv > 0.0f) {1105 Qcur = ggml_clamp(ctx0, Qcur, -hparams.f_clamp_kqv, hparams.f_clamp_kqv);1106 cb(Qcur, "Qcur_clamped", il);1107 }1108 Kcur = build_lora_mm(layer.wk, cur, layer.wk_s);1109 cb(Kcur, "Kcur", il);1110 if (layer.wk_b) {1111 Kcur = ggml_add(ctx0, Kcur, layer.wk_b);1112 cb(Kcur, "Kcur", il);1113 }1114 if (hparams.f_clamp_kqv > 0.0f) {1115 Kcur = ggml_clamp(ctx0, Kcur, -hparams.f_clamp_kqv, hparams.f_clamp_kqv);1116 cb(Kcur, "Kcur_clamped", il);1117 }1118 Vcur = build_lora_mm(layer.wv, cur, layer.wv_s);1119 cb(Vcur, "Vcur", il);1120 if (layer.wv_b) {1121 Vcur = ggml_add(ctx0, Vcur, layer.wv_b);1122 cb(Vcur, "Vcur", il);1123 }1124 if (hparams.f_clamp_kqv > 0.0f) {1125 Vcur = ggml_clamp(ctx0, Vcur, -hparams.f_clamp_kqv, hparams.f_clamp_kqv);1126 cb(Vcur, "Vcur_clamped", il);1127 }1128 Qcur = ggml_reshape_3d(ctx0, Qcur, n_embd_head, n_head, n_tokens);1129 Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens);1130 Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens);1131 }1132 1133 cb(Qcur, "Qcur", il);1134 cb(Kcur, "Kcur", il);1135 cb(Vcur, "Vcur", il);1136 1137 return { Qcur, Kcur, Vcur };1138}1139 1140 1141ggml_tensor * llm_graph_context::build_ffn(1142 ggml_tensor * cur,1143 ggml_tensor * up,1144 ggml_tensor * up_b,1145 ggml_tensor * up_s,1146 ggml_tensor * gate,1147 ggml_tensor * gate_b,1148 ggml_tensor * gate_s,1149 ggml_tensor * down,1150 ggml_tensor * down_b,1151 ggml_tensor * down_s,1152 ggml_tensor * act_scales,1153 llm_ffn_op_type type_op,1154 llm_ffn_gate_type type_gate,1155 int il) const {1156 ggml_tensor * tmp = up ? build_lora_mm(up, cur) : cur;1157 cb(tmp, "ffn_up", il);1158 1159 if (up_b) {1160 tmp = ggml_add(ctx0, tmp, up_b);1161 cb(tmp, "ffn_up_b", il);1162 }1163 1164 if (up_s) {1165 tmp = ggml_mul(ctx0, tmp, up_s);1166 cb(tmp, "ffn_up_s", il);1167 }1168 1169 if (gate) {1170 switch (type_gate) {1171 case LLM_FFN_SEQ:1172 {1173 cur = build_lora_mm(gate, tmp);1174 cb(cur, "ffn_gate", il);1175 } break;1176 case LLM_FFN_PAR:1177 {1178 cur = build_lora_mm(gate, cur);1179 cb(cur, "ffn_gate", il);1180 } break;1181 }1182 1183 if (gate_b) {1184 cur = ggml_add(ctx0, cur, gate_b);1185 cb(cur, "ffn_gate_b", il);1186 }1187 1188 if (gate_s) {1189 cur = ggml_mul(ctx0, cur, gate_s);1190 cb(cur, "ffn_gate_s", il);1191 }1192 1193 } else {1194 cur = tmp;1195 }1196 1197 switch (type_op) {1198 case LLM_FFN_SILU:1199 if (gate && type_gate == LLM_FFN_PAR) {1200 // Step35: HF clamps gate (after SiLU) and up before multiplication