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-kv-cache-dsv4.h"2 3#include "ggml-backend.h"4#include "llama-impl.h"5#include "llama-batch.h"6#include "llama-io.h"7#include "llama-model.h"8 9#include <algorithm>10#include <cassert>11#include <climits>12#include <cstdlib>13#include <cstring>14#include <map>15#include <sstream>16#include <stdexcept>17 18static constexpr uint32_t DSV4_CSA_RATIO = 4;19static constexpr uint32_t DSV4_HCA_RATIO = 128;20 21static constexpr uint32_t DSV4_STATE_MAGIC = 0x34565344; // DSV422static constexpr uint32_t DSV4_STATE_VERSION = 1;23static constexpr uint32_t DSV4_STATE_MODE_FULL = 0;24static constexpr uint32_t DSV4_STATE_MODE_PARTIAL = 1;25static constexpr uint32_t DSV4_K_CACHE_STATE_VER = 2;26static constexpr uint32_t DSV4_COMP_STATE_VER = 1;27 28static uint32_t dsv4_comp_size(uint32_t kv_size, uint32_t ratio) {29 return std::max<uint32_t>(1, (kv_size + ratio - 1)/ratio);30}31 32static void dsv4_clear_tensor_stream(ggml_tensor * tensor, uint32_t stream) {33 GGML_ASSERT(ggml_is_contiguous(tensor));34 GGML_ASSERT(tensor->ne[3] == 1);35 GGML_ASSERT(stream < (uint32_t) tensor->ne[2]);36 37 const size_t stream_size = tensor->nb[2];38 ggml_backend_tensor_memset(tensor, 0, stream*stream_size, stream_size);39}40 41static uint32_t dsv4_state_n_used_k_rows(llama_pos pos_max, uint32_t ratio, uint32_t kv_size) {42 if (pos_max < 0) {43 return 0;44 }45 46 const uint64_t n_rows = ((uint64_t) pos_max + 1)/ratio;47 48 return (uint32_t) std::min<uint64_t>(kv_size, n_rows);49}50 51static int64_t dsv4_stream_offset(uint32_t n_stream, llama_seq_id seq_id, uint32_t size) {52 if (n_stream <= 1) {53 return 0;54 }55 if (seq_id < 0 || (uint32_t) seq_id >= n_stream) {56 throw std::runtime_error("DSV4 sequence id out of stream range");57 }58 59 return (int64_t) seq_id*size;60}61 62static bool dsv4_ubatch_has_coupled(const llama_ubatch & ubatch) {63 for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {64 if (ubatch.n_seq_id[i] > 1) {65 return true;66 }67 }68 69 return false;70}71 72static bool dsv4_token_has_seq(const llama_ubatch & ubatch, uint32_t i, llama_seq_id seq_id) {73 for (int32_t s = 0; s < ubatch.n_seq_id[i]; ++s) {74 if (ubatch.seq_id[i][s] == seq_id) {75 return true;76 }77 }78 79 return false;80}81 82static llama_ubatch dsv4_build_raw_write_ubatch(const llama_ubatch & ubatch) {83 if (!dsv4_ubatch_has_coupled(ubatch)) {84 return ubatch;85 }86 if (ubatch.embd) {87 throw std::runtime_error("DSV4 coupled embedding ubatches are not supported");88 }89 90 std::vector<uint32_t> counts(ubatch.n_seqs_unq, 0);91 uint32_t n_tokens = 0;92 for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) {93 const llama_seq_id seq_id = ubatch.seq_id_unq[s];94 for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {95 if (dsv4_token_has_seq(ubatch, i, seq_id)) {96 ++counts[s];97 ++n_tokens;98 }99 }100 }101 102 if (n_tokens == 0) {103 return ubatch;104 }105 106 const uint32_t n_seq_tokens = counts[0];107 for (uint32_t s = 1; s < counts.size(); ++s) {108 if (counts[s] != n_seq_tokens) {109 throw std::runtime_error("DSV4 coupled raw writes require equal sequence lengths");110 }111 }112 113 auto data = std::make_shared<llama_ubatch::data_t>();114 data->pos.resize((size_t) n_tokens*ubatch.n_pos);115 data->n_seq_id.reserve(n_tokens);116 data->seq_id.reserve(n_tokens);117 data->seq_id_data.reserve(n_tokens);118 data->seq_id_unq.assign(ubatch.seq_id_unq, ubatch.seq_id_unq + ubatch.n_seqs_unq);119 data->seq_idx.assign(LLAMA_MAX_SEQ, -1);120 data->output.assign(n_tokens, 0);121 if (ubatch.token) {122 data->token.reserve(n_tokens);123 }124 125 for (uint32_t s = 0; s < data->seq_id_unq.size(); ++s) {126 data->seq_idx[data->seq_id_unq[s]] = s;127 }128 129 for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) {130 const llama_seq_id seq_id = ubatch.seq_id_unq[s];131 for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {132 if (!dsv4_token_has_seq(ubatch, i, seq_id)) {133 continue;134 }135 136 const uint32_t dst = data->n_seq_id.size();137 if (ubatch.token) {138 data->token.push_back(ubatch.token[i]);139 }140 for (uint32_t p = 0; p < ubatch.n_pos; ++p) {141 data->pos[(size_t) p*n_tokens + dst] = ubatch.pos[(size_t) p*ubatch.n_tokens + i];142 }143 data->n_seq_id.push_back(1);144 data->seq_id_data.push_back(seq_id);145 }146 }147 148 for (uint32_t i = 0; i < n_tokens; ++i) {149 data->seq_id.push_back(&data->seq_id_data[i]);150 }151 152 llama_ubatch res {153 /*.b_equal_seqs =*/ true,154 /*.n_tokens =*/ n_tokens,155 /*.n_seq_tokens =*/ n_seq_tokens,156 /*.n_seqs =*/ ubatch.n_seqs_unq,157 /*.n_seqs_unq =*/ ubatch.n_seqs_unq,158 /*.n_pos =*/ ubatch.n_pos,159 /*.token =*/ data->token.empty() ? nullptr : data->token.data(),160 /*.embd =*/ nullptr,161 /*.pos =*/ data->pos.data(),162 /*.n_seq_id =*/ data->n_seq_id.data(),163 /*.seq_id =*/ data->seq_id.data(),164 /*.seq_id_unq =*/ data->seq_id_unq.data(),165 /*.seq_idx =*/ data->seq_idx.data(),166 /*.output =*/ data->output.data(),167 /*.data =*/ data,168 };169 170 return res;171}172 173static std::vector<llama_ubatch> dsv4_build_raw_write_ubatches(const std::vector<llama_ubatch> & ubatches) {174 std::vector<llama_ubatch> res;175 res.reserve(ubatches.size());176 for (const llama_ubatch & ubatch : ubatches) {177 res.push_back(dsv4_build_raw_write_ubatch(ubatch));178 }179 return res;180}181 182static bool dsv4_batch_has_coupled(const llama_batch & batch) {183 if (!batch.n_seq_id) {184 return false;185 }186 187 for (int32_t i = 0; i < batch.n_tokens; ++i) {188 if (batch.n_seq_id[i] > 1) {189 return true;190 }191 }192 193 return false;194}195 196static int64_t dsv4_comp_graph_n_stream(const llama_ubatch & ubatch, uint32_t n_stream) {197 // Coupled sequence sets must stay in one graph stream because their198 // compressed state is shared. Independent per-seq state can fan out.199 if (n_stream <= 1 || ubatch.n_seqs_unq <= 1 || dsv4_ubatch_has_coupled(ubatch)) {200 return 1;201 }202 203 return ubatch.n_seqs_unq;204}205 206static void dsv4_state_src_stream_range(207 uint32_t n_stream,208 llama_seq_id seq_id,209 uint32_t & s0,210 uint32_t & ns) {211 if (seq_id >= 0 && n_stream > 1) {212 if ((uint32_t) seq_id >= n_stream) {213 throw std::runtime_error("DSV4 state sequence id out of stream range");214 }215 216 s0 = (uint32_t) seq_id;217 ns = 1;218 return;219 }220 221 s0 = 0;222 ns = seq_id >= 0 ? 1 : n_stream;223}224 225static void dsv4_state_dst_stream_range(226 uint32_t n_stream,227 llama_seq_id seq_id,228 uint32_t ns,229 uint32_t & s0) {230 if (seq_id >= 0) {231 if (ns != 1) {232 throw std::runtime_error("DSV4 sequence state stream count mismatch");233 }234 if (n_stream > 1 && (uint32_t) seq_id >= n_stream) {235 throw std::runtime_error("DSV4 state sequence id out of stream range");236 }237 238 s0 = n_stream > 1 ? (uint32_t) seq_id : 0;239 return;240 }241 242 if (ns != n_stream) {243 throw std::runtime_error("DSV4 full state stream count mismatch");244 }245 246 s0 = 0;247}248 249static void dsv4_state_write_tensor_streams(250 llama_io_write_i & io,251 ggml_tensor * tensor,252 uint32_t tensor_rows,253 uint32_t n_rows,254 uint32_t s0,255 uint32_t ns,256 const std::vector<uint32_t> * stream_ids = nullptr) {257 const int32_t type_i = (int32_t) tensor->type;258 const uint64_t ne0 = tensor->ne[0];259 const uint64_t rows = n_rows;260 const uint64_t row_size = ggml_row_size(tensor->type, tensor->ne[0]);261 262 if (n_rows > tensor_rows) {263 throw std::runtime_error("DSV4 state tensor row count exceeds storage");264 }265 266 io.write(&type_i, sizeof(type_i));267 io.write(&ne0, sizeof(ne0));268 io.write(&rows, sizeof(rows));269 io.write(&row_size, sizeof(row_size));270 271 const size_t stream_stride = (size_t) tensor_rows*row_size;272 const size_t size = (size_t) n_rows*row_size;273 if (size == 0) {274 return;275 }276 277 if (stream_ids && stream_ids->size() != ns) {278 throw std::runtime_error("DSV4 state tensor stream map size mismatch");279 }280 281 for (uint32_t s = 0; s < ns; ++s) {282 const uint32_t stream = stream_ids ? (*stream_ids)[s] : s0 + s;283 if ((int64_t) stream >= tensor->ne[2]) {284 throw std::runtime_error("DSV4 state tensor stream out of range");285 }286 const size_t offset = (size_t) stream*stream_stride;287 io.write_tensor(tensor, offset, size);288 }289}290 291static void dsv4_state_read_tensor_streams(292 llama_io_read_i & io,293 ggml_tensor * tensor,294 uint32_t tensor_rows,295 uint32_t n_rows,296 uint32_t s0,297 uint32_t ns) {298 int32_t type_i_ref;299 uint64_t ne0_ref;300 uint64_t rows_ref;301 uint64_t row_size_ref;302 303 io.read(&type_i_ref, sizeof(type_i_ref));304 io.read(&ne0_ref, sizeof(ne0_ref));305 io.read(&rows_ref, sizeof(rows_ref));306 io.read(&row_size_ref, sizeof(row_size_ref));307 308 const int32_t type_i = (int32_t) tensor->type;309 const uint64_t ne0 = tensor->ne[0];310 const uint64_t rows = n_rows;311 const uint64_t row_size = ggml_row_size(tensor->type, tensor->ne[0]);312 313 if (type_i != type_i_ref || ne0 != ne0_ref || rows != rows_ref || row_size != row_size_ref) {314 throw std::runtime_error("DSV4 state tensor metadata mismatch");315 }316 if (n_rows > tensor_rows) {317 throw std::runtime_error("DSV4 state tensor row count exceeds storage");318 }319 320 const size_t stream_stride = (size_t) tensor_rows*row_size;321 const size_t size = (size_t) n_rows*row_size;322 if (size == 0) {323 return;324 }325 326 for (uint32_t s = 0; s < ns; ++s) {327 const size_t offset = (size_t) (s0 + s)*stream_stride;328 io.read_tensor(tensor, offset, size);329 }330}331 332static void dsv4_state_write_k_cache(333 llama_io_write_i & io,334 const llama_kv_cache * kv,335 llama_seq_id seq_id,336 llama_state_seq_flags flags,337 uint32_t n_rows) {338 GGML_UNUSED(flags);339 340 uint32_t s0;341 uint32_t ns;342 dsv4_state_src_stream_range(kv->get_n_stream(), seq_id, s0, ns);343 344 const uint32_t version = DSV4_K_CACHE_STATE_VER;345 const uint32_t kv_size = kv->get_size();346 const auto layer_ids = kv->get_layer_ids();347 const uint32_t n_layer = layer_ids.size();348 349 if (n_rows > kv_size) {350 throw std::runtime_error("DSV4 K-cache state row count exceeds cache size");351 }352 353 io.write(&version, sizeof(version));354 io.write(&n_rows, sizeof(n_rows));355 io.write(&ns, sizeof(ns));356 io.write(&n_layer, sizeof(n_layer));357 358 for (uint32_t il : layer_ids) {359 io.write(&il, sizeof(il));360 dsv4_state_write_tensor_streams(io, kv->get_k_storage(il), kv_size, n_rows, s0, ns);361 }362}363 364static void dsv4_state_read_k_cache(365 llama_io_read_i & io,366 llama_kv_cache * kv,367 llama_seq_id seq_id,368 llama_state_seq_flags flags) {369 GGML_UNUSED(flags);370 371 uint32_t version;372 uint32_t n_rows_ref;373 uint32_t ns;374 uint32_t n_layer_ref;375 376 io.read(&version, sizeof(version));377 io.read(&n_rows_ref, sizeof(n_rows_ref));378 io.read(&ns, sizeof(ns));379 io.read(&n_layer_ref, sizeof(n_layer_ref));380 381 if (version != 1 && version != DSV4_K_CACHE_STATE_VER) {382 throw std::runtime_error("DSV4 K-cache state version mismatch");383 }384 385 const uint32_t kv_size = kv->get_size();386 if (version == 1 && n_rows_ref != kv_size) {387 LLAMA_LOG_INFO("kv size ref %d kv %d\n", n_rows_ref, kv_size);388 throw std::runtime_error("DSV4 K-cache state size mismatch");389 }390 if (n_rows_ref > kv_size) {391 LLAMA_LOG_INFO("kv rows ref %d kv %d\n", n_rows_ref, kv_size);392 throw std::runtime_error("DSV4 K-cache state size mismatch");393 }394 395 uint32_t s0;396 dsv4_state_dst_stream_range(kv->get_n_stream(), seq_id, ns, s0);397 398 const auto layer_ids = kv->get_layer_ids();399 if (n_layer_ref != layer_ids.size()) {400 throw std::runtime_error("DSV4 K-cache layer count mismatch");401 }402 403 for (uint32_t il : layer_ids) {404 uint32_t il_ref;405 io.read(&il_ref, sizeof(il_ref));406 if (il_ref != il) {407 throw std::runtime_error("DSV4 K-cache layer id mismatch");408 }409 410 dsv4_state_read_tensor_streams(io, kv->get_k_storage(il), kv_size, n_rows_ref, s0, ns);411 }412}413 414static std::string dsv4_plan_positions(const std::vector<int32_t> & values) {415 std::ostringstream ss;416 ss << "[";417 for (size_t i = 0; i < values.size(); ++i) {418 if (i > 0) {419 ss << ", ";420 }421 ss << values[i];422 }423 ss << "]";424 return ss.str();425}426 427static llama_kv_cache_dsv4_context::comp_plan dsv4_build_comp_plan(428 const llama_ubatch & ubatch,429 uint32_t ratio,430 bool overlap,431 uint32_t state_size,432 uint32_t kv_size,433 uint32_t n_stream,434 uint32_t n_rs_seq,435 const std::vector<uint32_t> & rs_idx) {436 llama_kv_cache_dsv4_context::comp_plan plan;437 plan.n_visible.resize(ubatch.n_tokens);438 plan.n_stream = dsv4_comp_graph_n_stream(ubatch, n_stream);439 440 // n_stream is the persistent cache/state layout; plan.n_stream is the441 // graph view for this ubatch and can be a subset of those streams.442 if (n_stream <= 1 && ubatch.n_seqs_unq > 1) {443 throw std::runtime_error("DSV4 single compressed stream cannot serve multiple sequences");444 }445 446 const int64_t state_rows = (int64_t) state_size*n_stream;447 448 struct persist_row {449 int32_t dst;450 int32_t src;451 llama_pos pos;452 };453 454 std::vector<persist_row> persist_rows;455 456 // For the overlap compressor, build_overlap_compressed_kv_from_state() consumes457 // state_read_idxs as two contiguous halves: the first ratio*n_blocks entries are458 // the "previous-window" gather indices for every block, followed by the459 // "current-window" indices for every block. Collect them separately here and460 // append cur after prev once the loop has visited all completed blocks461 std::vector<int32_t> overlap_prev_reads;462 std::vector<int32_t> overlap_cur_reads;463 464 std::map<std::pair<llama_seq_id, llama_pos>, int64_t> curr_token_idx_map;465 std::map<llama_seq_id, uint32_t> state_write_counts;466 467 for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {468 for (int32_t s = 0; s < ubatch.n_seq_id[i]; ++s) {469 curr_token_idx_map[std::make_pair(ubatch.seq_id[i][s], ubatch.pos[i])] = i;470 }471 }472 473 const auto state_source_idx = [&](llama_seq_id seq_id, llama_pos pos) -> int32_t {474 if (pos < 0) {475 // The overlap compressor needs a zero/-inf source for the first476 // block's previous half. The graph appends that row after the477 // current-ubatch scratch rows.478 return (int32_t) (state_rows + ubatch.n_tokens);479 }480 481 const auto key = std::make_pair(seq_id, pos);482 if (curr_token_idx_map.find(key) != curr_token_idx_map.end()) {483 return (int32_t) (state_rows + curr_token_idx_map.at(key));484 }485 486 const int64_t stream_off = dsv4_stream_offset(n_stream, seq_id, state_size);487 return (int32_t) (stream_off + pos%state_size);488 };489 490 for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {491 const llama_pos pos = ubatch.pos[i];492 493 if (pos < 0) {494 continue;495 }496 497 plan.state_pos.push_back((int32_t) (pos%ratio));498 499 const int64_t n_visible = (int64_t) (pos + 1)/ratio;500 plan.n_visible[i] = (int32_t) n_visible;501 plan.n_kv = std::max(plan.n_kv, n_visible);502 503 for (int32_t s = 0; s < ubatch.n_seq_id[i]; ++s) {504 const llama_seq_id seq_id = ubatch.seq_id[i][s];505 const int64_t stream_off = dsv4_stream_offset(n_stream, seq_id, state_size);506 const int32_t state_idx = (int32_t) (stream_off + pos%state_size);507 508 const auto it = std::find_if(persist_rows.begin(), persist_rows.end(),509 [state_idx](const persist_row & row) {510 return row.dst == state_idx;511 });512 if (it == persist_rows.end()) {513 persist_rows.push_back({ state_idx, (int32_t) i, pos });514 } else if (pos > it->pos) {515 it->src = (int32_t) i;516 it->pos = pos;517 }518 519 if ((pos + 1) % ratio != 0) {520 continue;521 }522 523 const llama_pos source_start = pos + 1 - ratio;524 const int64_t cache_off = dsv4_stream_offset(n_stream, seq_id, kv_size);525 526 plan.state_write_idxs.push_back(cache_off + pos/ratio);527 plan.state_write_pos.push_back((int32_t) source_start);528 ++state_write_counts[seq_id];529 530 if (overlap) {531 const llama_pos prev_start = source_start - ratio;532 533 for (uint32_t j = 0; j < ratio; ++j) {534 overlap_prev_reads.push_back(state_source_idx(seq_id, prev_start + j));535 }536 for (uint32_t j = 0; j < ratio; ++j) {537 overlap_cur_reads.push_back(state_source_idx(seq_id, source_start + j));538 }539 } else {540 for (uint32_t j = 0; j < ratio; ++j) {541 plan.state_read_idxs.push_back(state_source_idx(seq_id, source_start + j));542 }543 }544 }545 }546 547 if (ratio == DSV4_CSA_RATIO && !plan.state_pos.empty()) {548 assert(kv_size > 0);549 550 // Pad each stream to the reserve plan's block count.551 const auto append_dummy_block = [&](llama_seq_id seq_id, uint32_t i) {552 const int64_t cache_off = dsv4_stream_offset(n_stream, seq_id, kv_size);553 const int32_t source_idx = state_source_idx(seq_id, ubatch.pos[i]);554 555 plan.state_write_idxs.push_back(cache_off + kv_size - 1);556 plan.state_write_pos .push_back(0);557 558 if (overlap) {559 for (uint32_t j = 0; j < ratio; ++j) {560 overlap_prev_reads.push_back(source_idx);561 overlap_cur_reads .push_back(source_idx);562 }563 } else {564 for (uint32_t j = 0; j < ratio; ++j) {565 plan.state_read_idxs.push_back(source_idx);566 }567 }568 };569 570 if (dsv4_ubatch_has_coupled(ubatch)) {571 if (plan.state_write_idxs.empty()) {572 uint32_t i = 0;573 while (i < ubatch.n_tokens && ubatch.pos[i] < 0) {574 ++i;575 }576 assert(i < ubatch.n_tokens);577 append_dummy_block(ubatch.seq_id[i][0], i);578 }579 } else {580 const uint32_t n_blocks = (std::max<uint32_t>(1, ubatch.n_seq_tokens) + ratio - 1)/ratio;581 582 for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) {583 const llama_seq_id seq_id = ubatch.seq_id_unq[s];584 const uint32_t n_writes = state_write_counts[seq_id];585 if (n_writes >= n_blocks) {586 continue;587 }588 if (n_writes + 1 != n_blocks) {589 throw std::runtime_error("DSV4 CSA sequence positions are not contiguous");590 }591 592 uint32_t i = 0;593 while (i < ubatch.n_tokens && (ubatch.pos[i] < 0 || !dsv4_token_has_seq(ubatch, i, seq_id))) {594 ++i;595 }596 assert(i < ubatch.n_tokens);597 append_dummy_block(seq_id, i);598 }599 }600 }601 602 if (overlap) {603 // [ all blocks' prev-window indices | all blocks' cur-window indices ]604 plan.state_read_idxs.reserve(overlap_prev_reads.size() + overlap_cur_reads.size());605 plan.state_read_idxs.insert(plan.state_read_idxs.end(),606 overlap_prev_reads.begin(), overlap_prev_reads.end());607 plan.state_read_idxs.insert(plan.state_read_idxs.end(),608 overlap_cur_reads.begin(), overlap_cur_reads.end());609 }610 611 plan.n_kv = GGML_PAD(plan.n_kv, 256u);612 613 std::sort(persist_rows.begin(), persist_rows.end(),614 [](const persist_row & a, const persist_row & b) {615 return a.dst < b.dst;616 });617 618 for (const persist_row & row : persist_rows) {619 plan.state_persist_src_idxs.push_back(row.src);620 plan.state_persist_dst_idxs.push_back(row.dst);621 }622 623 624 if (n_rs_seq > 0) {625 for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) {626 const llama_seq_id seq_id = ubatch.seq_id_unq[s];627 if (seq_id < 0 || (uint32_t) seq_id >= n_stream) {628 continue;629 }630 631 const int64_t stream_off = dsv4_stream_offset(n_stream, seq_id, state_size);632 const uint32_t rollback = (uint32_t) seq_id < rs_idx.size() ? rs_idx[seq_id] : 0;633 // Keep the restore graph fixed-width when no rollback is pending.634 const int64_t src_plane = rollback > 0 && rollback <= n_rs_seq ? (int64_t) rollback*state_rows : 0;635 for (uint32_t r = 0; r < state_size; ++r) {636 plan.state_restore_src_idxs.push_back((int32_t) (src_plane + stream_off + r));637 plan.state_restore_dst_idxs.push_back((int32_t) (stream_off + r));638 }639 640 std::vector<uint32_t> token_idxs;641 token_idxs.reserve(ubatch.n_tokens);642 for (uint32_t i = 0; i < ubatch.n_tokens; ++i) {643 if (dsv4_token_has_seq(ubatch, i, seq_id)) {644 token_idxs.push_back(i);645 }646 }647 if (token_idxs.empty()) {648 continue;649 }650 651 const uint32_t n_seq_tokens = (uint32_t) token_idxs.size();652 const int64_t scratch_off = (int64_t) state_rows*(1 + n_rs_seq);653 for (uint32_t d = 1; d <= n_rs_seq; ++d) {654 const int64_t dst_plane = (int64_t) d*state_rows;655 656 for (uint32_t r = 0; r < state_size; ++r) {657 int32_t src;658 if (d <= n_seq_tokens) {659 const uint32_t prefix = n_seq_tokens - d;660 src = (int32_t) (stream_off + r);661 662 for (uint32_t j = 0; j < prefix; ++j) {663 const uint32_t i_tok = token_idxs[j];664 if (ubatch.pos[i_tok] >= 0 && (uint32_t) (ubatch.pos[i_tok]%state_size) == r) {665 src = (int32_t) (scratch_off + i_tok);666 }667 }668 } else {669 const int64_t src_plane = (int64_t) (d - n_seq_tokens)*state_rows;670 src = (int32_t) (src_plane + stream_off + r);671 }672 673 plan.state_snapshot_src_idxs.push_back(src);674 plan.state_snapshot_dst_idxs.push_back((int32_t) (dst_plane + stream_off + r));675 }676 }677 }678 }679 680 static const bool debug = []() {681 const char * env = getenv("LLAMA_DSV4_COMPRESS_DEBUG");682 return env && atoi(env) > 0;683 }();684 685 if (debug) {686 LLAMA_LOG_INFO("%s: ratio=%u, n_tokens=%u, state_persist_dst=%s, state_write_pos=%s\n",687 __func__, ratio, ubatch.n_tokens,688 dsv4_plan_positions(plan.state_persist_dst_idxs).c_str(),689 dsv4_plan_positions(plan.state_write_pos).c_str());690 }691 692 return plan;693}694 695static std::vector<llama_kv_cache_dsv4_context::comp_plan> dsv4_build_comp_plans(696 const std::vector<llama_ubatch> & ubatches,697 uint32_t ratio,698 bool overlap,699 uint32_t state_size,700 uint32_t kv_size,701 uint32_t n_stream,702 uint32_t n_rs_seq,703 const std::vector<uint32_t> & rs_idx) {704 std::vector<llama_kv_cache_dsv4_context::comp_plan> plans;705 plans.reserve(ubatches.size());706 707 for (const llama_ubatch & ubatch : ubatches) {708 plans.push_back(dsv4_build_comp_plan(ubatch, ratio, overlap, state_size, kv_size, n_stream, n_rs_seq, rs_idx));709 }710 711 return plans;712}713 714static llama_kv_cache::slot_info_vec_t dsv4_build_comp_sinfos(715 const std::vector<llama_ubatch> & ubatches,716 uint32_t n_stream) {717 llama_kv_cache::slot_info_vec_t sinfos;718 sinfos.reserve(ubatches.size());719 720 for (const llama_ubatch & ubatch : ubatches) {721 if (n_stream <= 1 && ubatch.n_seqs_unq > 1) {722 throw std::runtime_error("DSV4 single compressed stream cannot serve multiple sequences");723 }724 725 const uint32_t ns = (uint32_t) dsv4_comp_graph_n_stream(ubatch, n_stream);726 llama_kv_cache::slot_info sinfo;727 sinfo.s0 = n_stream > 1 ? LLAMA_MAX_SEQ : 0;728 sinfo.s1 = 0;729 sinfo.resize(ns);730 731 for (uint32_t s = 0; s < ns; ++s) {732 const llama_seq_id seq_id = n_stream > 1 ? ubatch.seq_id_unq[s] : 0;733 const uint32_t strm = (uint32_t) dsv4_stream_offset(n_stream, seq_id, 1);734 735 sinfo.s0 = std::min(sinfo.s0, strm);736 sinfo.s1 = std::max(sinfo.s1, strm);737 sinfo.strm[s] = strm;738 sinfo.idxs[s].resize(1, 0);739 }740 741 if (n_stream > 1 && sinfo.s1 - sinfo.s0 + 1 != ns) {742 throw std::runtime_error("DSV4 compressed streams are not contiguous in ubatch");743 }744 745 sinfos.push_back(std::move(sinfo));746 }747 748 return sinfos;749}750 751static llama_kv_cache::slot_info_vec_t dsv4_build_raw_read_sinfos(752 const llama_kv_cache::slot_info_vec_t & sinfos_write,753 const std::vector<llama_ubatch> & ubatches) {754 llama_kv_cache::slot_info_vec_t sinfos;755 sinfos.reserve(ubatches.size());756 757 for (size_t i = 0; i < ubatches.size(); ++i) {758 const llama_ubatch & ubatch = ubatches[i];759 const auto & sinfo_write = sinfos_write[i];760 761 if (!dsv4_ubatch_has_coupled(ubatch)) {762 sinfos.push_back(sinfo_write);763 continue;764 }765 766 const llama_seq_id seq_id = ubatch.seq_id[0][0];767 uint32_t i_stream = 0;768 for (; i_stream < sinfo_write.n_stream(); ++i_stream) {769 if (sinfo_write.strm[i_stream] == seq_id) {770 break;771 }772 }773 if (i_stream == sinfo_write.n_stream()) {774 throw std::runtime_error("DSV4 raw write stream not found for coupled read");775 }776 777 llama_kv_cache::slot_info sinfo;778 sinfo.s0 = sinfo_write.strm[i_stream];779 sinfo.s1 = sinfo_write.strm[i_stream];780 sinfo.resize(1);781 sinfo.strm[0] = sinfo_write.strm[i_stream];782 sinfo.idxs[0] = sinfo_write.idxs[i_stream];783 sinfos.push_back(std::move(sinfo));784 }785 786 return sinfos;787}788 789static llama_kv_cache_dsv4_context::comp_plan dsv4_build_reserve_comp_plan(790 const llama_ubatch & ubatch,791 uint32_t ratio,792 bool overlap,793 uint32_t state_size,794 uint32_t kv_size,795 uint32_t n_stream,796 uint32_t n_rs_seq) {797 llama_kv_cache_dsv4_context::comp_plan plan;798 plan.n_visible.resize(ubatch.n_tokens);799 plan.n_stream = dsv4_comp_graph_n_stream(ubatch, n_stream);800 plan.n_kv = kv_size;801 802 if (ubatch.n_tokens == 0) {803 return plan;804 }805 806 const uint32_t n_seqs = std::max<uint32_t>(1, ubatch.n_seqs);807 const uint32_t n_seq_tokens = std::max<uint32_t>(1, ubatch.n_seq_tokens);808 const uint64_t n_blocks_u64 = (uint64_t) n_seqs*((n_seq_tokens + ratio - 1)/ratio);809 const size_t n_blocks = (size_t) std::max<uint64_t>(1, n_blocks_u64);810 GGML_ASSERT((uint64_t) n_blocks == std::max<uint64_t>(1, n_blocks_u64));811 812 const uint64_t state_rows = (uint64_t) state_size*n_stream;813 const size_t n_persist = (size_t) std::min<uint64_t>(ubatch.n_tokens, state_rows);814 const size_t n_restore = n_rs_seq > 0 ? (size_t) state_size*std::max<uint32_t>(1, ubatch.n_seqs_unq) : 0;815 const size_t n_snapshot = (size_t) n_rs_seq*state_size*std::max<uint32_t>(1, ubatch.n_seqs_unq);816 817 plan.state_pos .resize(ubatch.n_tokens);818 plan.state_persist_src_idxs.resize(n_persist);819 plan.state_persist_dst_idxs.resize(n_persist);820 plan.state_restore_src_idxs.resize(n_restore);821 plan.state_restore_dst_idxs.resize(n_restore);822 plan.state_snapshot_src_idxs.resize(n_snapshot);823 plan.state_snapshot_dst_idxs.resize(n_snapshot);824 plan.state_read_idxs .resize((overlap ? 2u : 1u)*ratio*n_blocks);825 plan.state_write_idxs.resize(n_blocks);826 plan.state_write_pos .resize(n_blocks);827 828 return plan;829}830 831static void dsv4_make_k_only(llama_hparams & hparams) {832 // llama_kv_cache uses hparams.is_mla() to allocate K-only storage.833 hparams.n_embd_head_k_mla_impl = hparams.n_embd_head_k();834 hparams.n_embd_head_v_mla_impl = hparams.n_embd_head_k();835}836 837//838// llama_dsv4_comp_state839//840 841llama_dsv4_comp_state::llama_dsv4_comp_state(842 const llama_model & model,843 bool offload,844 bool unified,845 uint32_t n_seq_max,846 uint32_t ratio,847 uint32_t state_size,848 uint32_t n_embd_state,849 uint32_t n_rs_seq,850 const char * name,851 const llama_memory_i::layer_filter_cb & filter) :852 ratio(ratio),853 state_size(state_size),854 n_embd_state(n_embd_state),855 n_stream(unified ? 1 : n_seq_max),856 n_rs_seq(n_rs_seq) {857 const llama_hparams & hparams = model.hparams;858 859 struct ggml_backend_buft_comparator {860 bool operator()(const ggml_backend_buffer_type_t & lhs, const ggml_backend_buffer_type_t & rhs) const {861 return strcmp(ggml_backend_buft_name(lhs), ggml_backend_buft_name(rhs)) < 0;862 }863 };864 865 std::map<ggml_backend_buffer_type_t, ggml_context_ptr, ggml_backend_buft_comparator> ctx_map;866 867 auto ctx_for_buft = [&](ggml_backend_buffer_type_t buft) -> ggml_context * {868 auto it = ctx_map.find(buft);869 if (it == ctx_map.end()) {870 ggml_init_params params = {871 /*.mem_size =*/ size_t(2u*(1 + n_stream)*hparams.n_layer()*ggml_tensor_overhead()),872 /*.mem_buffer =*/ NULL,873 /*.no_alloc =*/ true,874 };875 876 ggml_context * ctx = ggml_init(params);877 if (!ctx) {878 return nullptr;879 }880 881 ctx_map.emplace(buft, ctx);882 883 return ctx;884 }885 886 return it->second.get();887 };888 889 for (uint32_t il = 0; il < hparams.n_layer(); ++il) {890 if (filter && !filter(il)) {891 continue;892 }893 894 const char * dev_name = "CPU";895 896 ggml_backend_buffer_type_t buft = ggml_backend_cpu_buffer_type();897 898 if (offload) {899 auto * dev = model.dev_layer(il);900 buft = ggml_backend_dev_buffer_type(dev);901 902 dev_name = ggml_backend_dev_name(dev);903 }904 905 LLAMA_LOG_DEBUG("%s: layer %3d: dev = %s\n", __func__, il, dev_name);906 907 ggml_context * ctx = ctx_for_buft(buft);908 if (!ctx) {909 throw std::runtime_error("failed to create ggml context for DSV4 compressor state");910 }911 912 const uint32_t n_planes = n_stream*(1 + n_rs_seq);913 ggml_tensor * kv = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd_state, state_size, n_planes);914 ggml_tensor * score = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd_state, state_size, n_planes);915 916 ggml_format_name(kv, "dsv4_%s_state_kv_l%d", name, il);917 ggml_format_name(score, "dsv4_%s_state_score_l%d", name, il);918 919 std::vector<ggml_tensor *> kv_stream;920 std::vector<ggml_tensor *> score_stream;921 922 for (uint32_t s = 0; s < n_stream; ++s) {923 kv_stream.push_back(ggml_view_2d(ctx, kv, n_embd_state, state_size, kv->nb[1], s*kv->nb[2]));924 score_stream.push_back(ggml_view_2d(ctx, score, n_embd_state, state_size, score->nb[1], s*score->nb[2]));925 }926 927 map_layer_ids[il] = layers.size();928 929 layers.push_back({ il, kv, score, std::move(kv_stream), std::move(score_stream) });930 }931 932 for (auto & [buft, ctx] : ctx_map) {933 ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors_from_buft(ctx.get(), buft);934 if (!buf) {935 throw std::runtime_error("failed to allocate buffer for DSV4 compressor state");936 }937 938 ggml_backend_buffer_clear(buf, 0);939 940 LLAMA_LOG_INFO("%s: %10s DSV4 %s state buffer size = %8.2f MiB\n",941 __func__, ggml_backend_buffer_name(buf), name, ggml_backend_buffer_get_size(buf)/1024.0/1024.0);942 943 ctxs_bufs.emplace_back(std::move(ctx), buf);944 }945 946 LLAMA_LOG_INFO("%s: %s ratio = %u, state = %u x %u, streams = %u, rs_seq = %u, layers = %zu, size = %7.2f MiB\n",947 __func__, name, ratio, state_size, n_embd_state, n_stream, n_rs_seq, layers.size(), total_size()/1024.0/1024.0);948}949 950void llama_dsv4_comp_state::clear(llama_seq_id seq_id, bool data) {951 if (!data) {952 return;953 }954 955 if (seq_id >= 0) {956 GGML_ASSERT((uint32_t) seq_id < n_stream);957 958 for (const auto & layer : layers) {959 for (uint32_t d = 0; d <= n_rs_seq; ++d) {960 const uint32_t stream = d*n_stream + (uint32_t) seq_id;961 dsv4_clear_tensor_stream(layer.kv, stream);962 dsv4_clear_tensor_stream(layer.score, stream);963 }964 }965 return;966 }967 968 for (auto & [_, buf] : ctxs_bufs) {969 ggml_backend_buffer_clear(buf.get(), 0);970 }971}972 973void llama_dsv4_comp_state::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst) {974 GGML_ASSERT(seq_id_src >= 0 && (uint32_t) seq_id_src < n_stream);975 GGML_ASSERT(seq_id_dst >= 0 && (uint32_t) seq_id_dst < n_stream);976 977 if (seq_id_src == seq_id_dst) {978 return;979 }980 981 clear(seq_id_dst, true);982 983 sc_info.ssrc.push_back((uint32_t) seq_id_src);984 sc_info.sdst.push_back((uint32_t) seq_id_dst);985}986 987void llama_dsv4_comp_state::apply_copies(const stream_copy_info & sc_info) const {988 for (size_t i = 0; i < sc_info.ssrc.size(); ++i) {989 const uint32_t ssrc = sc_info.ssrc[i];990 const uint32_t sdst = sc_info.sdst[i];991 992 for (const auto & layer : layers) {993 ggml_backend_tensor_copy(layer.kv_stream[ssrc], layer.kv_stream[sdst]);994 ggml_backend_tensor_copy(layer.score_stream[ssrc], layer.score_stream[sdst]);995 }996 }997}998 999uint32_t llama_dsv4_comp_state::get_ratio() const {1000 return ratio;1001}1002 1003uint32_t llama_dsv4_comp_state::get_state_size() const {1004 return state_size;1005}1006 1007uint32_t llama_dsv4_comp_state::get_n_stream() const {1008 return n_stream;1009}1010 1011uint32_t llama_dsv4_comp_state::get_n_rs_seq() const {1012 return n_rs_seq;1013}1014 1015uint32_t llama_dsv4_comp_state::get_n_rows() const {1016 return state_size*n_stream;1017}1018 1019std::map<ggml_backend_buffer_type_t, size_t> llama_dsv4_comp_state::memory_breakdown() const {1020 std::map<ggml_backend_buffer_type_t, size_t> ret;1021 for (const auto & [_, buf] : ctxs_bufs) {1022 ggml_backend_buffer_type_t buft = ggml_backend_buffer_get_type(buf.get());1023 ret[buft] += ggml_backend_buffer_get_size(buf.get());1024 }1025 return ret;1026}1027 1028void llama_dsv4_comp_state::state_write(1029 llama_io_write_i & io,1030 llama_seq_id seq_id,1031 llama_state_seq_flags flags,1032 const std::vector<uint32_t> & rs_idx) const {1033 GGML_UNUSED(flags);1034 1035 uint32_t s0;1036 uint32_t ns;1037 dsv4_state_src_stream_range(n_stream, seq_id, s0, ns);1038 1039 std::vector<uint32_t> stream_ids(ns);1040 for (uint32_t s = 0; s < ns; ++s) {1041 const uint32_t seq = seq_id >= 0 ? (uint32_t) seq_id : s0 + s;1042 if (seq >= rs_idx.size() || rs_idx[seq] > n_rs_seq) {1043 throw std::runtime_error("DSV4 recurrent state rollback index out of range");1044 }1045 stream_ids[s] = rs_idx[seq]*n_stream + s0 + s;1046 }1047 1048 const uint32_t version = DSV4_COMP_STATE_VER;1049 const uint32_t n_layer = layers.size();1050 1051 io.write(&version, sizeof(version));1052 io.write(&ratio, sizeof(ratio));1053 io.write(&state_size, sizeof(state_size));1054 io.write(&n_embd_state, sizeof(n_embd_state));1055 io.write(&ns, sizeof(ns));1056 io.write(&n_layer, sizeof(n_layer));1057 1058 for (const auto & layer : layers) {1059 io.write(&layer.il, sizeof(layer.il));1060 1061 dsv4_state_write_tensor_streams(io, layer.kv, state_size, state_size, s0, ns, &stream_ids);1062 dsv4_state_write_tensor_streams(io, layer.score, state_size, state_size, s0, ns, &stream_ids);1063 }1064}1065 1066void llama_dsv4_comp_state::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {1067 GGML_UNUSED(flags);1068 1069 uint32_t version;1070 uint32_t ratio_ref;1071 uint32_t state_size_ref;1072 uint32_t n_embd_state_ref;1073 uint32_t ns;1074 uint32_t n_layer_ref;1075 1076 io.read(&version, sizeof(version));1077 io.read(&ratio_ref, sizeof(ratio_ref));1078 io.read(&state_size_ref, sizeof(state_size_ref));1079 io.read(&n_embd_state_ref, sizeof(n_embd_state_ref));1080 io.read(&ns, sizeof(ns));1081 io.read(&n_layer_ref, sizeof(n_layer_ref));1082 1083 if (version != DSV4_COMP_STATE_VER) {1084 throw std::runtime_error("DSV4 compressor state version mismatch");1085 }1086 if (ratio_ref != ratio || state_size_ref != state_size || n_embd_state_ref != n_embd_state) {1087 throw std::runtime_error("DSV4 compressor state metadata mismatch");1088 }1089 if (n_layer_ref != layers.size()) {1090 throw std::runtime_error("DSV4 compressor state layer count mismatch");1091 }1092 1093 uint32_t s0;1094 dsv4_state_dst_stream_range(n_stream, seq_id, ns, s0);1095 1096 for (const auto & layer : layers) {1097 uint32_t il_ref;1098 io.read(&il_ref, sizeof(il_ref));1099 if (il_ref != layer.il) {1100 throw std::runtime_error("DSV4 compressor state layer id mismatch");1101 }1102 1103 dsv4_state_read_tensor_streams(io, layer.kv, state_size, state_size, s0, ns);1104 dsv4_state_read_tensor_streams(io, layer.score, state_size, state_size, s0, ns);1105 }1106}1107 1108ggml_tensor * llama_dsv4_comp_state::get_kv_all(ggml_context * ctx, int32_t il) const {1109 const int32_t ids = map_layer_ids.at(il);1110 ggml_tensor * state = layers[ids].kv;1111 1112 return ggml_view_2d(ctx, state, state->ne[0], get_n_rows()*(1 + n_rs_seq), state->nb[1], 0);1113}1114 1115ggml_tensor * llama_dsv4_comp_state::get_score_all(ggml_context * ctx, int32_t il) const {1116 const int32_t ids = map_layer_ids.at(il);1117 ggml_tensor * state = layers[ids].score;1118 1119 return ggml_view_2d(ctx, state, state->ne[0], get_n_rows()*(1 + n_rs_seq), state->nb[1], 0);1120}1121 1122ggml_tensor * llama_dsv4_comp_state::get_kv(ggml_context * ctx, int32_t il) const {1123 ggml_tensor * state = get_kv_all(ctx, il);1124 const size_t row_size = ggml_row_size(state->type, state->ne[0]);1125 1126 return ggml_view_2d(ctx, state, state->ne[0], get_n_rows(), state->nb[1], 0*row_size);1127}1128 1129ggml_tensor * llama_dsv4_comp_state::get_score(ggml_context * ctx, int32_t il) const {1130 ggml_tensor * state = get_score_all(ctx, il);1131 const size_t row_size = ggml_row_size(state->type, state->ne[0]);1132 1133 return ggml_view_2d(ctx, state, state->ne[0], get_n_rows(), state->nb[1], 0*row_size);1134}1135 1136ggml_tensor * llama_dsv4_comp_state::cpy_kv(ggml_context * ctx, ggml_tensor * cur, ggml_tensor * idxs, int32_t il) const {1137 return ggml_set_rows(ctx, get_kv_all(ctx, il), cur, idxs);1138}1139 1140ggml_tensor * llama_dsv4_comp_state::cpy_score(ggml_context * ctx, ggml_tensor * cur, ggml_tensor * idxs, int32_t il) const {1141 return ggml_set_rows(ctx, get_score_all(ctx, il), cur, idxs);1142}1143 1144size_t llama_dsv4_comp_state::total_size() const {1145 size_t size = 0;1146 1147 for (const auto & [_, buf] : ctxs_bufs) {1148 size += ggml_backend_buffer_get_size(buf.get());1149 }1150 1151 return size;1152}1153 1154//1155// llama_kv_cache_dsv41156//1157 1158llama_kv_cache_dsv4::llama_kv_cache_dsv4(1159 const llama_model & model,1160 ggml_type type_k,1161 ggml_type type_v,1162 bool v_trans,1163 bool offload,1164 bool swa_full,1165 bool unified,1166 uint32_t kv_size,1167 uint32_t n_seq_max,1168 uint32_t n_ubatch,1169 uint32_t n_pad,1170 uint32_t n_rs_seq,1171 const layer_filter_cb & filter,1172 const layer_reuse_cb & reuse) :1173 hparams_raw(model.hparams),1174 hparams_csa(model.hparams),1175 hparams_hca(model.hparams),1176 hparams_lid(model.hparams),1177 n_seq_max(n_seq_max),1178 n_rs_seq(n_rs_seq),1179 rs_idx(n_seq_max, 0) {1180 1181 const layer_filter_cb filter_raw = [&](int32_t il) {1182 if (filter && !filter(il)) {1183 return false;1184 }1185 1186 return true;1187 };1188 1189 GGML_UNUSED(unified);1190 1191 // Keep DSV4 KV/state streams per sequence even when public KV mode is unified.1192 const bool unified_raw = false;1193 1194 hparams_raw.n_layer_nextn = 0;1195 hparams_csa.n_layer_nextn = 0;1196 hparams_hca.n_layer_nextn = 0;1197 hparams_lid.n_layer_nextn = 0;1198 1199 LLAMA_LOG_INFO("%s: creating DSV4 raw KV cache\n", __func__);1200 