echodict/llama.cpp
version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786
0610
1#include "llama-memory-recurrent.h"2 3#include "ggml-backend.h"4#include "llama-impl.h"5#include "llama-io.h"6#include "llama-batch.h"7#include "llama-model.h"8 9#include <algorithm>10#include <cassert>11#include <cstring>12#include <limits>13#include <map>14#include <stdexcept>15 16//17// llama_memory_recurrent18//19 20llama_memory_recurrent::llama_memory_recurrent(21 const llama_model & model,22 ggml_type type_r,23 ggml_type type_s,24 bool offload,25 uint32_t mem_size,26 uint32_t n_seq_max,27 const layer_filter_cb & filter) : hparams(model.hparams), n_seq_max(n_seq_max) {28 const int32_t n_layer = hparams.n_layer;29 30 head = 0;31 size = mem_size;32 used = 0;33 34 cells.clear();35 cells.resize(mem_size);36 37 // define a comparator for the buft -> ctx map to ensure that the order is well-defined:38 struct ggml_backend_buft_comparator {39 bool operator()(const ggml_backend_buffer_type_t & lhs, const ggml_backend_buffer_type_t & rhs) const {40 return strcmp(ggml_backend_buft_name(lhs), ggml_backend_buft_name(rhs)) < 0;41 }42 };43 std::map<ggml_backend_buffer_type_t, ggml_context_ptr, ggml_backend_buft_comparator> ctx_map;44 45 // create a context for each buffer type46 auto ctx_for_buft = [&](ggml_backend_buffer_type_t buft) -> ggml_context * {47 auto it = ctx_map.find(buft);48 if (it == ctx_map.end()) {49 ggml_init_params params = {50 /*.mem_size =*/ size_t(2u*n_layer*ggml_tensor_overhead()),51 /*.mem_buffer =*/ NULL,52 /*.no_alloc =*/ true,53 };54 55 ggml_context * ctx = ggml_init(params);56 if (!ctx) {57 return nullptr;58 }59 60 ctx_map.emplace(buft, ctx);61 62 return ctx;63 }64 65 return it->second.get();66 };67 68 r_l.resize(n_layer);69 s_l.resize(n_layer);70 71 for (int i = 0; i < n_layer; i++) {72 if (filter && !filter(i)) {73 LLAMA_LOG_DEBUG("%s: layer %3d: skipped\n", __func__, i);74 continue;75 }76 77 const char * dev_name = "CPU";78 79 ggml_backend_buffer_type_t buft = ggml_backend_cpu_buffer_type();80 81 if (offload) {82 auto * dev = model.dev_layer(i);83 buft = ggml_backend_dev_buffer_type(dev);84 85 dev_name = ggml_backend_dev_name(dev);86 }87 88 LLAMA_LOG_DEBUG("%s, layer %3d: dev = %s\n", __func__, i, dev_name);89 90 ggml_context * ctx = ctx_for_buft(buft);91 if (!ctx) {92 throw std::runtime_error("failed to create ggml context for rs cache");93 }94 95 ggml_tensor * r = ggml_new_tensor_2d(ctx, type_r, hparams.n_embd_r(), mem_size);96 ggml_tensor * s = ggml_new_tensor_2d(ctx, type_s, hparams.n_embd_s(), mem_size);97 ggml_format_name(r, "cache_r_l%d", i);98 ggml_format_name(s, "cache_s_l%d", i);99 r_l[i] = r;100 s_l[i] = s;101 }102 103 // allocate tensors and initialize the buffers to avoid NaNs in the padding104 for (auto & [buft, ctx] : ctx_map) {105 ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors_from_buft(ctx.get(), buft);106 if (!buf) {107 throw std::runtime_error("failed to allocate buffer for rs cache");108 }109 ggml_backend_buffer_clear(buf, 0);110 LLAMA_LOG_INFO("%s: %10s RS buffer size = %8.2f MiB\n", __func__, ggml_backend_buffer_name(buf), ggml_backend_buffer_get_size(buf)/1024.0/1024.0);111 ctxs_bufs.emplace_back(std::move(ctx), buf);112 }113 114 {115 const size_t memory_size_r = size_r_bytes();116 const size_t memory_size_s = size_s_bytes();117 118 LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u seqs), R (%s): %7.2f MiB, S (%s): %7.2f MiB\n", __func__,119 (float)(memory_size_r + memory_size_s) / (1024.0f * 1024.0f), mem_size, n_layer, n_seq_max,120 ggml_type_name(type_r), (float)memory_size_r / (1024.0f * 1024.0f),121 ggml_type_name(type_s), (float)memory_size_s / (1024.0f * 1024.0f));122 }123}124 125void llama_memory_recurrent::clear(bool data) {126 for (int32_t i = 0; i < (int32_t) size; ++i) {127 cells[i].pos = -1;128 cells[i].seq_id.clear();129 cells[i].src = -1;130 cells[i].tail = -1;131 }132 133 head = 0;134 used = 0;135 136 if (data) {137 for (auto & [_, buf] : ctxs_bufs) {138 ggml_backend_buffer_clear(buf.get(), 0);139 }140 }141}142 143bool llama_memory_recurrent::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {144 //printf("[DEBUG] calling llama_memory_recurrent::seq_rm` with `seq_id=%d, p0=%d, p1=%d`\n", seq_id, p0, p1);145 uint32_t new_head = size;146 147 if (p0 < 0) {148 p0 = 0;149 }150 151 if (p1 < 0) {152 p1 = std::numeric_limits<llama_pos>::max();153 }154 155 // models like Mamba or RWKV can't have a state partially erased at the end156 // of the sequence because their state isn't preserved for previous tokens157 if (seq_id >= (int64_t) size) {158 // could be fatal159 return false;160 }161 if (0 <= seq_id) {162 int32_t & tail_id = cells[seq_id].tail;163 if (tail_id >= 0) {164 const auto & cell = cells[tail_id];165 // partial intersection is invalid if it includes the final pos166 if (0 < p0 && p0 <= cell.pos && p1 > cell.pos) {167 //printf("[DEBUG] inside `llama_memory_recurrent::seq_rm`: partial intersection is invalid, so returning false, p0 = %d, cell.pos = %d, p1 = %d\n", p0, cell.pos, p1);168 return false;169 }170 // invalidate tails which will be cleared171 if (p0 <= cell.pos && cell.pos < p1) {172 tail_id = -1;173 }174 }175 } else {176 // seq_id is negative, then the range should include everything or nothing177 if (p0 != p1 && (p0 != 0 || p1 != std::numeric_limits<llama_pos>::max())) {178 //printf("[DEBUG] inside `llama_memory_recurrent::seq_rm`: `seq_id` is negative, so returning false\n");179 return false;180 }181 }182 183 for (uint32_t i = 0; i < size; ++i) {184 if (cells[i].pos >= p0 && cells[i].pos < p1) {185 if (seq_id < 0) {186 cells[i].seq_id.clear();187 } else if (cells[i].has_seq_id(seq_id)) {188 cells[i].seq_id.erase(seq_id);189 } else {190 continue;191 }192 if (cells[i].is_empty()) {193 // keep count of the number of used cells194 if (cells[i].pos >= 0) {195 used--;196 }197 cells[i].pos = -1;198 cells[i].src = -1;199 if (new_head == size) {200 new_head = i;201 }202 }203 }204 }205 206 // If we freed up a slot, set head to it so searching can start there.207 if (new_head != size && new_head < head) {208 head = new_head;209 }210 211 return true;212}213 214void llama_memory_recurrent::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {215 if (seq_id_src == seq_id_dst) {216 return;217 }218 219 if (p0 < 0) {220 p0 = 0;221 }222 223 if (p1 < 0) {224 p1 = std::numeric_limits<llama_pos>::max();225 }226 227 if ((uint32_t) seq_id_dst < size && (uint32_t) seq_id_src < size) {228 auto & tail_src = cells[seq_id_src];229 auto & tail_dst = cells[seq_id_dst];230 if (tail_dst.tail >= 0) {231 // clear destination seq_id if it wasn't empty232 auto & cell_dst = cells[tail_dst.tail];233 234 cell_dst.seq_id.erase(seq_id_dst);235 tail_dst.tail = -1;236 if (cell_dst.seq_id.empty()) {237 cell_dst.pos = -1;238 cell_dst.src = -1;239 used -= 1;240 }241 }242 if (tail_src.tail >= 0) {243 auto & cell_src = cells[tail_src.tail];244 245 cell_src.seq_id.insert(seq_id_dst);246 tail_dst.tail = tail_src.tail;247 }248 }249}250 251void llama_memory_recurrent::seq_keep(llama_seq_id seq_id) {252 uint32_t new_head = size;253 254 for (uint32_t i = 0; i < size; ++i) {255 if ((llama_seq_id) i != seq_id) {256 cells[i].tail = -1;257 }258 259 if (!cells[i].has_seq_id(seq_id)) {260 if (cells[i].pos >= 0) {261 used--;262 }263 264 cells[i].pos = -1;265 cells[i].src = -1;266 cells[i].seq_id.clear();267 268 if (new_head == size){269 new_head = i;270 }271 } else {272 cells[i].seq_id.clear();273 cells[i].seq_id.insert(seq_id);274 }275 }276 277 // If we freed up a slot, set head to it so searching can start there.278 if (new_head != size && new_head < head) {279 head = new_head;280 }281}282 283void llama_memory_recurrent::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {284 if (shift == 0) {285 return;286 }287 288 if (p0 < 0) {289 p0 = 0;290 }291 292 if (p1 < 0) {293 p1 = std::numeric_limits<llama_pos>::max();294 }295 296 // If there is no range then return early to avoid looping over the297 if (p0 == p1) {298 return;299 }300 301 // for Mamba-like or RWKV models, only the pos needs to be shifted302 if (0 <= seq_id && seq_id < (int64_t) size) {303 const int32_t tail_id = cells[seq_id].tail;304 if (tail_id >= 0) {305 auto & cell = cells[tail_id];306 if (cell.has_seq_id(seq_id) && p0 <= cell.pos && cell.pos < p1) {307 cell.pos += shift;308 }309 }310 }311}312 313void llama_memory_recurrent::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {314 if (d == 1) {315 return;316 }317 318 if (p0 < 0) {319 p0 = 0;320 }321 322 if (p1 < 0) {323 p1 = std::numeric_limits<llama_pos>::max();324 }325 326 // If there is no range then return early to avoid looping over the cache.327 if (p0 == p1) {328 return;329 }330 331 // for Mamba-like or RWKV models, only the pos needs to be changed332 if (0 <= seq_id && seq_id < (int64_t) size) {333 const int32_t tail_id = cells[seq_id].tail;334 if (tail_id >= 0) {335 auto & cell = cells[tail_id];336 if (cell.has_seq_id(seq_id) && p0 <= cell.pos && cell.pos < p1) {337 cell.pos /= d;338 }339 }340 }341}342 343llama_pos llama_memory_recurrent::seq_pos_min(llama_seq_id seq_id) const {344 llama_pos result = std::numeric_limits<llama_pos>::max();345 346 for (uint32_t i = 0; i < size; ++i) {347 if (cells[i].has_seq_id(seq_id)) {348 result = std::min(result, cells[i].pos);349 }350 }351 352 if (result == std::numeric_limits<llama_pos>::max()) {353 result = -1;354 }355 356 return result;357}358 359llama_pos llama_memory_recurrent::seq_pos_max(llama_seq_id seq_id) const {360 llama_pos result = -1;361 362 for (uint32_t i = 0; i < size; ++i) {363 if (cells[i].has_seq_id(seq_id)) {364 result = std::max(result, cells[i].pos);365 }366 }367 368 return result;369}370 371std::map<ggml_backend_buffer_type_t, size_t> llama_memory_recurrent::memory_breakdown() const {372 std::map<ggml_backend_buffer_type_t, size_t> ret;373 for (const auto & [_, buf] : ctxs_bufs) {374 ret[ggml_backend_buffer_get_type(buf.get())] += ggml_backend_buffer_get_size(buf.get());375 }376 return ret;377}378 379llama_memory_context_ptr llama_memory_recurrent::init_batch(llama_batch_allocr & balloc, uint32_t n_ubatch, bool embd_all) {380 do {381 balloc.split_reset();382 383 std::vector<llama_ubatch> ubatches;384 while (true) {385 llama_ubatch ubatch;386 387 if (embd_all) {388 // if all tokens are output, split by sequence389 ubatch = balloc.split_seq(n_ubatch);390 } else {391 // TODO: non-sequential equal split can be done if using unified KV cache392 // for simplicity, we always use sequential equal split for now393 ubatch = balloc.split_equal(n_ubatch, true);394 }395 396 if (ubatch.n_tokens == 0) {397 break;398 }399 400 ubatches.push_back(std::move(ubatch)); // NOLINT401 }402 403 if (balloc.get_n_used() < balloc.get_n_tokens()) {404 // failed to find a suitable split405 break;406 }407 408 if (!prepare(ubatches)) {409 break;410 }411 412 return std::make_unique<llama_memory_recurrent_context>(this, std::move(ubatches));413 } while (false);414 415 return std::make_unique<llama_memory_recurrent_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);416}417 418llama_memory_context_ptr llama_memory_recurrent::init_full() {419 return std::make_unique<llama_memory_recurrent_context>(this);420}421 422llama_memory_context_ptr llama_memory_recurrent::init_update(llama_context * lctx, bool optimize) {423 GGML_UNUSED(lctx);424 GGML_UNUSED(optimize);425 426 return std::make_unique<llama_memory_recurrent_context>(LLAMA_MEMORY_STATUS_NO_UPDATE);427}428 429bool llama_memory_recurrent::prepare(const std::vector<llama_ubatch> & ubatches) {430 // simply remember the full state because it is very small for this type of cache431 // TODO: optimize432 auto org_cells = cells;433 auto org_used = used;434 auto org_head = head;435 436 bool success = true;437 438 for (const auto & ubatch : ubatches) {439 if (!find_slot(ubatch)) {440 success = false;441 break;442 }443 }444 445 // restore the original state446 cells = std::move(org_cells);447 used = org_used;448 head = org_head;449 450 return success;451}452 453bool llama_memory_recurrent::find_slot(const llama_ubatch & ubatch) {454 const uint32_t n_seq_tokens = ubatch.n_seq_tokens;455 const uint32_t n_seqs = ubatch.n_seqs;456 457 // if we have enough unused cells before the current head ->458 // better to start searching from the beginning of the cache, hoping to fill it459 if (head > used + 2*n_seqs) {460 head = 0;461 }462 463 // For recurrent state architectures (like Mamba or RWKV),464 // each cache cell can store the state for a whole sequence.465 // A slot should be always be contiguous.466 467 // can only process batches with an equal number of new tokens in each sequence468 GGML_ASSERT(ubatch.equal_seqs());469 470 int32_t min = size - 1;471 int32_t max = 0;472 473 // everything should fit if all seq_ids are smaller than the max474 for (uint32_t s = 0; s < n_seqs; ++s) {475 const uint32_t i = s*n_seq_tokens; // first token of sequence set s476 const uint32_t n_seq_id = ubatch.n_seq_id[i];477 478 for (uint32_t j = 0; j < n_seq_id; ++j) {479 const llama_seq_id seq_id = ubatch.seq_id[i][j];480 481 if (seq_id < 0 || (uint32_t) seq_id >= size) {482 // too big seq_id483 // TODO: would it be possible to resize the cache instead?484 LLAMA_LOG_ERROR("%s: seq_id=%d >= n_seq_max=%u Try using a bigger --parallel value\n", __func__, seq_id, n_seq_max);485 return false;486 }487 if (j > 0) {488 auto & seq = cells[seq_id];489 if (seq.tail >= 0) {490 auto & cell = cells[seq.tail];491 // clear cells from seq_ids that become shared492 // (should not normally happen, but let's handle it anyway)493 cell.seq_id.erase(seq_id);494 seq.tail = -1;495 if (cell.seq_id.empty()) {496 cell.pos = -1;497 cell.src = -1;498 used -= 1;499 }500 }501 }502 }503 }504 505#ifndef NDEBUG506 {507 std::vector<int32_t> tails_verif;508 tails_verif.assign(size, -1);509 for (uint32_t i = 0; i < size; ++i) {510 auto & cell = cells[i];511 for (llama_seq_id seq_id : cell.seq_id) {512 if (tails_verif[seq_id] != -1) {513 LLAMA_LOG_ERROR("%s: duplicate tail for seq_id %d in cell %d and %d\n", __func__, seq_id, i, tails_verif[seq_id]);514 }515 tails_verif[seq_id] = i;516 }517 }518 for (uint32_t i = 0; i < size; ++i) {519 if (tails_verif[i] != cells[i].tail) {520 LLAMA_LOG_ERROR("%s: wrong tail for seq_id %d, (%d instead of %d)\n", __func__, i, cells[i].tail, tails_verif[i]);521 }522 }523 }524#endif525 526 // find next empty cell527 uint32_t next_empty_cell = head;528 529 for (uint32_t i = 0; i < size; ++i) {530 if (next_empty_cell >= size) { next_empty_cell -= size; }531 auto & cell = cells[next_empty_cell];532 if (cell.is_empty()) { break; }533 next_empty_cell += 1;534 }535 536 // find usable cell range537 for (uint32_t s = 0; s < n_seqs; ++s) {538 const uint32_t i = s*n_seq_tokens;539 const llama_seq_id seq_id = ubatch.seq_id[i][0];540 auto & seq_meta = cells[seq_id];541 bool has_cell = false;542 if (seq_meta.tail >= 0) {543 auto & cell = cells[seq_meta.tail];544 GGML_ASSERT(cell.has_seq_id(seq_id));545 // does this seq_id "own" the cell?546 if (cell.seq_id.size() == 1) { has_cell = true; }547 }548 if (!has_cell) {549 auto & empty_cell = cells[next_empty_cell];550 GGML_ASSERT(empty_cell.is_empty());551 // copy old tail into the empty cell552 if (seq_meta.tail >= 0) {553 auto & orig_cell = cells[seq_meta.tail];554 empty_cell.pos = orig_cell.pos;555 empty_cell.src = orig_cell.src;556 orig_cell.seq_id.erase(seq_id);557 empty_cell.seq_id.insert(seq_id); // will be overwritten558 GGML_ASSERT(!orig_cell.is_empty()); // has at least one remaining seq_id559 }560 seq_meta.tail = next_empty_cell;561 // find next empty cell562 if (s + 1 < n_seqs) {563 for (uint32_t j = 0; j < size; ++j) {564 next_empty_cell += 1;565 if (next_empty_cell >= size) { next_empty_cell -= size; }566 auto & cell = cells[next_empty_cell];567 if (cell.is_empty()) { break; }568 }569 }570 }571 if (min > seq_meta.tail) { min = seq_meta.tail; }572 if (max < seq_meta.tail) { max = seq_meta.tail; }573 }574 575 // gather and re-order576 for (uint32_t s = 0; s < n_seqs; ++s) {577 const uint32_t i = s*n_seq_tokens;578 const int32_t dst_id = s + min;579 const int32_t src_id = cells[ubatch.seq_id[i][0]].tail;580 if (dst_id != src_id) {581 auto & dst_cell = cells[dst_id];582 auto & src_cell = cells[src_id];583 584 std::swap(dst_cell.pos, src_cell.pos);585 std::swap(dst_cell.src, src_cell.src);586 std::swap(dst_cell.seq_id, src_cell.seq_id);587 588 // swap tails589 for (uint32_t j = 0; j < size; ++j) {590 int32_t & tail = cells[j].tail;591 if (tail == src_id) {592 tail = dst_id;593 } else if (tail == dst_id) {594 tail = src_id;595 }596 }597 }598 }599 600 // update the pos of the used seqs601 for (uint32_t s = 0; s < n_seqs; ++s) {602 const uint32_t i = s*n_seq_tokens;603 const llama_pos last_pos = ubatch.pos[i + n_seq_tokens - 1];604 const int32_t cell_id = s + min;605 auto & cell = cells[cell_id];606 607 if (cell.pos >= 0 && last_pos != cell.pos + (llama_pos) n_seq_tokens) {608 // What should happen when the pos backtracks or skips a value?609 // Clearing the state mid-batch would require special-casing which isn't done.610 LLAMA_LOG_WARN("%s: non-consecutive token position %d after %d for sequence %d with %u new tokens\n",611 __func__, last_pos, cell.pos, ubatch.seq_id[i][0], n_seq_tokens);612 }613 cell.pos = last_pos;614 cell.seq_id.clear();615 for (int32_t j = 0; j < ubatch.n_seq_id[i]; ++j) {616 const llama_seq_id seq_id = ubatch.seq_id[i][j];617 cell.seq_id.insert(seq_id);618 cells[seq_id].tail = cell_id;619 }620 }621 622 // Find first cell without src refs, to use as the zero-ed state623 {624 // TODO: bake-in src refcounts in the cell metadata625 std::vector<int32_t> refcounts(size, 0);626 for (size_t i = 0; i < size; ++i) {627 const int32_t src = cells[i].src;628 if (src >= 0) {629 refcounts[src] += 1;630 }631 }632 633 rs_z = -1;634 for (int i = min; i <= max; ++i) {635 if (refcounts[i] == 0) {636 rs_z = i;637 break;638 }639 }640 641 for (int i = min; i <= max; ++i) {642 if (cells[i].src < 0) {643 GGML_ASSERT(rs_z >= 0);644 cells[i].src0 = rs_z;645 } else {646 // Stage the source ids for all used cells to allow correct seq_* behavior647 // and still make these values available when setting the inputs648 cells[i].src0 = cells[i].src;649 }650 cells[i].src = i; // avoid moving or clearing twice651 }652 }653 654 // allow getting the range of used cells, from head to head + n655 head = min;656 n = max - min + 1;657 used = std::count_if(cells.begin(), cells.end(),658 [](const mem_cell & cell){ return !cell.is_empty(); });659 660 // sanity check661 return n >= n_seqs;662}663 664bool llama_memory_recurrent::get_can_shift() const {665 // shifting the pos is trivial for recurrent models666 return true;667}668 669size_t llama_memory_recurrent::total_size() const {670 size_t size = 0;671 for (const auto & [_, buf] : ctxs_bufs) {672 size += ggml_backend_buffer_get_size(buf.get());673 }674 675 return size;676}677 678size_t llama_memory_recurrent::size_r_bytes() const {679 size_t size_r_bytes = 0;680 681 for (const auto & r : r_l) {682 if (r != nullptr) {683 size_r_bytes += ggml_nbytes(r);684 }685 }686 687 return size_r_bytes;688}689 690size_t llama_memory_recurrent::size_s_bytes() const {691 size_t size_s_bytes = 0;692 693 for (const auto & s : s_l) {694 if (s != nullptr) {695 size_s_bytes += ggml_nbytes(s);696 }697 }698 699 return size_s_bytes;700}701 702void llama_memory_recurrent::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {703 GGML_UNUSED(flags);704 705 std::vector<std::pair<uint32_t, uint32_t>> cell_ranges; // ranges, from inclusive, to exclusive706 uint32_t cell_count = 0;707 708 // Count the number of cells with the specified seq_id709 // Find all the ranges of cells with this seq id (or all, when -1)710 uint32_t cell_range_begin = size;711 for (uint32_t i = 0; i < size; ++i) {712 const auto & cell = cells[i];713 if ((seq_id == -1 && !cell.is_empty()) || cell.has_seq_id(seq_id)) {714 ++cell_count;715 if (cell_range_begin == size) {716 cell_range_begin = i;717 }718 } else {719 if (cell_range_begin != size) {720 cell_ranges.emplace_back(cell_range_begin, i);721 cell_range_begin = size;722 }723 }724 }725 if (cell_range_begin != size) {726 cell_ranges.emplace_back(cell_range_begin, size);727 }728 729 // DEBUG CHECK: Sum of cell counts in ranges should equal the total cell count730 uint32_t cell_count_check = 0;731 for (const auto & range : cell_ranges) {732 cell_count_check += range.second - range.first;733 }734 GGML_ASSERT(cell_count == cell_count_check);735 736 io.write(&cell_count, sizeof(cell_count));737 738 state_write_meta(io, cell_ranges, seq_id);739 state_write_data(io, cell_ranges);740}741 742void llama_memory_recurrent::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {743 GGML_UNUSED(flags);744 745 uint32_t cell_count;746 io.read_to(&cell_count, sizeof(cell_count));747 748 bool res = true;749 750 res = res && state_read_meta(io, cell_count, seq_id);751 res = res && state_read_data(io, cell_count);752 753 if (!res) {754 if (seq_id == -1) {755 clear(true);756 } else {757 seq_rm(seq_id, -1, -1);758 }759 throw std::runtime_error("failed to restore kv cache");760 }761}762 763void llama_memory_recurrent::state_write_meta(llama_io_write_i & io, const std::vector<std::pair<uint32_t, uint32_t>> & cell_ranges, llama_seq_id seq_id) const {764 for (const auto & range : cell_ranges) {765 for (uint32_t i = range.first; i < range.second; ++i) {766 const auto & cell = cells[i];767 const llama_pos pos = cell.pos;768 const uint32_t n_seq_id = seq_id == -1 ? cell.seq_id.size() : 0;769 770 io.write(&pos, sizeof(pos));771 io.write(&n_seq_id, sizeof(n_seq_id));772 773 if (n_seq_id) {774 for (auto seq_id : cell.seq_id) {775 io.write(&seq_id, sizeof(seq_id));776 }777 }778 }779 }780}781 782void llama_memory_recurrent::state_write_data(llama_io_write_i & io, const std::vector<std::pair<uint32_t, uint32_t>> & cell_ranges) const {783 const uint32_t s_trans = 0;784 const uint32_t n_layer = hparams.n_layer;785 786 io.write(&s_trans, sizeof(s_trans));787 io.write(&n_layer, sizeof(n_layer));788 789 // Iterate and write all the R tensors first, each row is a cell790 // Get whole range at a time791 for (uint32_t il = 0; il < n_layer; ++il) {792 // skip null layers (read_data will handle this by checking "r_l" and "s_l" for null)793 if (r_l[il] == nullptr) continue;794 795 // Write R tensor type796 const int32_t r_type_i = (int32_t)r_l[il]->type;797 io.write(&r_type_i, sizeof(r_type_i));798 799 // Write row size of R tensor800 const uint64_t r_size_row = ggml_row_size(r_l[il]->type, hparams.n_embd_r());801 io.write(&r_size_row, sizeof(r_size_row));802 803 // Write each range of cells of r_size_row length804 for (const auto & range : cell_ranges) {805 const size_t range_size = range.second - range.first;806 const size_t buf_size = range_size * r_size_row;807 io.write_tensor(r_l[il], range.first * r_size_row, buf_size);808 }809 }810 811 if (!s_trans) {812 for (uint32_t il = 0; il < n_layer; ++il) {813 // skip null layers (read_data will handle this by checking "r_l" and "s_l" for null)814 if (s_l[il] == nullptr) continue;815 816 // Write S tensor type817 const int32_t s_type_i = (int32_t)s_l[il]->type;818 io.write(&s_type_i, sizeof(s_type_i));819 820 // Write row size of S tensor821 const uint64_t s_size_row = ggml_row_size(s_l[il]->type, hparams.n_embd_s());822 io.write(&s_size_row, sizeof(s_size_row));823 824 // Write each range of S tensor rows825 for (const auto & range : cell_ranges) {826 const size_t range_size = range.second - range.first;827 const size_t buf_size = range_size * s_size_row;828 io.write_tensor(s_l[il], range.first * s_size_row, buf_size);829 }830 }831 } else {832 // When S tensor is transposed, we also need the element size and get the element ranges from each row833 const uint32_t mem_size = size;834 for (uint32_t il = 0; il < n_layer; ++il) {835 // skip null layers (read_data will handle this by checking "r_l" and "s_l" for null)836 if (s_l[il] == nullptr) continue;837 838 const uint32_t n_embd_s = hparams.n_embd_s();839 840 // Write S tensor type841 const int32_t s_type_i = (int32_t)s_l[il]->type;842 io.write(&s_type_i, sizeof(s_type_i));843 844 // Write element size845 const uint32_t s_size_el = ggml_type_size(s_l[il]->type);846 io.write(&s_size_el, sizeof(s_size_el));847 848 // Write GQA embedding size849 io.write(&n_embd_s, sizeof(n_embd_s));850 851 // For each row, we get the element values of each cell852 for (uint32_t j = 0; j < n_embd_s; ++j) {853 // Write each range of cells of s_size_el length854 for (const auto & range : cell_ranges) {855 const size_t range_size = range.second - range.first;856 const size_t src_offset = (range.first + j * mem_size) * s_size_el;857 const size_t buf_size = range_size * s_size_el;858 io.write_tensor(s_l[il], src_offset, buf_size);859 }860 }861 }862 }863}864 865bool llama_memory_recurrent::state_read_meta(llama_io_read_i & io, uint32_t cell_count, llama_seq_id dest_seq_id) {866 if (dest_seq_id != -1) {867 // single sequence868 seq_rm(dest_seq_id, -1, -1);869 870 if (cell_count == 0) {871 return true;872 }873 874 llama_batch_allocr balloc(hparams.n_pos_per_embd());875 876 llama_ubatch ubatch = balloc.ubatch_reserve(cell_count, 1);877 878 for (uint32_t i = 0; i < cell_count; ++i) {879 llama_pos pos;880 uint32_t n_seq_id;881 882 io.read_to(&pos, sizeof(pos));883 io.read_to(&n_seq_id, sizeof(n_seq_id));884 885 if (n_seq_id != 0) {886 LLAMA_LOG_ERROR("%s: invalid seq_id-agnostic kv cell\n", __func__);887 return false;888 }889 890 ubatch.pos[i] = pos;891 }892 ubatch.n_seq_id[0] = 1;893 ubatch.seq_id[0] = &dest_seq_id;894 895 if (!find_slot(ubatch)) {896 LLAMA_LOG_ERROR("%s: failed to find available cells in kv cache\n", __func__);897 return false;898 }899 900 // DEBUG CHECK: kv.head should be our first cell, kv.head + cell_count - 1 should be our last cell (verify seq_id and pos values)901 // Assume that this is one contiguous block of cells902 GGML_ASSERT(head + cell_count <= size);903 GGML_ASSERT(cells[head].pos == ubatch.pos[0]);904 GGML_ASSERT(cells[head + cell_count - 1].pos == ubatch.pos[cell_count - 1]);905 GGML_ASSERT(cells[head].has_seq_id(dest_seq_id));906 GGML_ASSERT(cells[head + cell_count - 1].has_seq_id(dest_seq_id));907 } else {908 // whole KV cache restore909 910 if (cell_count > size) {911 LLAMA_LOG_ERROR("%s: not enough cells in kv cache\n", __func__);912 return false;913 }914 915 clear(true);916 917 for (uint32_t i = 0; i < cell_count; ++i) {918 auto & cell = cells[i];919 920 llama_pos pos;921 uint32_t n_seq_id;922 923 io.read_to(&pos, sizeof(pos));924 io.read_to(&n_seq_id, sizeof(n_seq_id));925 926 cell.pos = pos;927 928 for (uint32_t j = 0; j < n_seq_id; ++j) {929 llama_seq_id seq_id;930 io.read_to(&seq_id, sizeof(seq_id));931 932 if (seq_id < 0 || (uint32_t) seq_id >= this->n_seq_max) {933 LLAMA_LOG_ERROR("%s: invalid seq_id, %d is out of range [0, %u)\n", __func__, seq_id, this->n_seq_max);934 return false;935 }936 937 cell.seq_id.insert(seq_id);938 939 int32_t & tail = cells[seq_id].tail;940 if (tail != -1) {941 LLAMA_LOG_ERROR("%s: duplicate tail for seq_id %d in cell %d and %d\n", __func__, seq_id, i, tail);942 return false;943 }944 tail = i;945 }946 }947 948 head = 0;949 used = cell_count;950 }951 952 for (uint32_t i = 0; i < cell_count; ++i) {953 uint32_t cell_id = head + i;954 // make sure the recurrent states will keep their restored state955 cells[cell_id].src = cell_id;956 }957 958 return true;959}960 961bool llama_memory_recurrent::state_read_data(llama_io_read_i & io, uint32_t cell_count) {962 uint32_t s_trans;963 uint32_t n_layer;964 io.read_to(&s_trans, sizeof(s_trans));965 io.read_to(&n_layer, sizeof(n_layer));966 967 if (n_layer != hparams.n_layer) {968 LLAMA_LOG_ERROR("%s: mismatched layer count (%u instead of %u)\n", __func__, n_layer, hparams.n_layer);969 return false;970 }971 if (cell_count > size) {972 LLAMA_LOG_ERROR("%s: not enough cells in kv cache to restore state (%u > %u)\n", __func__, cell_count, size);973 return false;974 }975 if (false != (bool) s_trans) {976 LLAMA_LOG_ERROR("%s: incompatible s transposition\n", __func__);977 return false;978 }979 980 // For each layer, read the keys for each cell, one row is one cell, read as one contiguous block981 for (uint32_t il = 0; il < n_layer; ++il) {982 // skip null layers983 if (r_l[il] == nullptr) continue;984 985 // Read type of key986 int32_t r_type_i_ref;987 io.read_to(&r_type_i_ref, sizeof(r_type_i_ref));988 const int32_t r_type_i = (int32_t) r_l[il]->type;989 if (r_type_i != r_type_i_ref) {990 LLAMA_LOG_ERROR("%s: mismatched r type (%d != %d, layer %d)\n", __func__, r_type_i, r_type_i_ref, il);991 return false;992 }993 994 // Read row size of key995 uint64_t r_size_row_ref;996 io.read_to(&r_size_row_ref, sizeof(r_size_row_ref));997 const size_t r_size_row = ggml_row_size(r_l[il]->type, hparams.n_embd_r());998 if (r_size_row != r_size_row_ref) {999 LLAMA_LOG_ERROR("%s: mismatched r row size (%zu != %zu, layer %d)\n", __func__, r_size_row, (size_t) r_size_row_ref, il);1000 return false;1001 }1002 1003 if (cell_count) {1004 // Read and set the keys for the whole cell range1005 ggml_backend_tensor_set(r_l[il], io.read(cell_count * r_size_row), head * r_size_row, cell_count * r_size_row);1006 }1007 }1008 1009 if (!s_trans) {1010 for (uint32_t il = 0; il < n_layer; ++il) {1011 // skip null layers1012 if (s_l[il] == nullptr) continue;1013 1014 // Read type of value1015 int32_t s_type_i_ref;1016 io.read_to(&s_type_i_ref, sizeof(s_type_i_ref));1017 const int32_t s_type_i = (int32_t)s_l[il]->type;1018 1019 if (s_type_i != s_type_i_ref) {1020 LLAMA_LOG_ERROR("%s: mismatched s type (%d != %d, layer %d)\n", __func__, s_type_i, s_type_i_ref, il);1021 return false;1022 }1023 1024 // Read row size of value1025 uint64_t s_size_row_ref;1026 io.read_to(&s_size_row_ref, sizeof(s_size_row_ref));1027 const size_t s_size_row = ggml_row_size(s_l[il]->type, hparams.n_embd_s());1028 if (s_size_row != s_size_row_ref) {1029 LLAMA_LOG_ERROR("%s: mismatched s row size (%zu != %zu, layer %d)\n", __func__, s_size_row, (size_t) s_size_row_ref, il);1030 return false;1031 }1032 1033 if (cell_count) {1034 // Read and set the values for the whole cell range1035 ggml_backend_tensor_set(s_l[il], io.read(cell_count * s_size_row), head * s_size_row, cell_count * s_size_row);1036 }1037 }1038 } else {1039 // For each layer, read the values for each cell (transposed)1040 for (uint32_t il = 0; il < n_layer; ++il) {1041 // skip null layers1042 if (s_l[il] == nullptr) continue;1043 1044 const uint32_t n_embd_s = hparams.n_embd_s();1045 1046 // Read type of value1047 int32_t s_type_i_ref;1048 io.read_to(&s_type_i_ref, sizeof(s_type_i_ref));1049 const int32_t s_type_i = (int32_t)s_l[il]->type;1050 if (s_type_i != s_type_i_ref) {1051 LLAMA_LOG_ERROR("%s: mismatched s type (%d != %d, layer %d)\n", __func__, s_type_i, s_type_i_ref, il);1052 return false;1053 }1054 1055 // Read element size of value1056 uint32_t s_size_el_ref;1057 io.read_to(&s_size_el_ref, sizeof(s_size_el_ref));1058 const size_t s_size_el = ggml_type_size(s_l[il]->type);1059 if (s_size_el != s_size_el_ref) {1060 LLAMA_LOG_ERROR("%s: mismatched s element size (%zu != %zu, layer %d)\n", __func__, s_size_el, (size_t) s_size_el_ref, il);1061 return false;1062 }1063 1064 // Read state embedding size1065 uint32_t n_embd_s_ref;1066 io.read_to(&n_embd_s_ref, sizeof(n_embd_s_ref));1067 if (n_embd_s != n_embd_s_ref) {1068 LLAMA_LOG_ERROR("%s: mismatched s embedding size (%u != %u, layer %d)\n", __func__, n_embd_s, n_embd_s_ref, il);1069 return false;1070 }1071 1072 if (cell_count) {1073 // For each row in the transposed matrix, read the values for the whole cell range1074 for (uint32_t j = 0; j < n_embd_s; ++j) {1075 const size_t dst_offset = (head + j * size) * s_size_el;1076 ggml_backend_tensor_set(s_l[il], io.read(cell_count * s_size_el), dst_offset, cell_count * s_size_el);1077 }1078 }1079 }1080 }1081 1082 return true;1083}1084 1085//1086// llama_memory_recurrent_context1087//1088 1089llama_memory_recurrent_context::llama_memory_recurrent_context(llama_memory_status status) : status(status) {}1090 1091llama_memory_recurrent_context::llama_memory_recurrent_context(1092 llama_memory_recurrent * mem) : status(LLAMA_MEMORY_STATUS_SUCCESS), mem(mem), is_full(true) {1093}1094 1095llama_memory_recurrent_context::llama_memory_recurrent_context(1096 llama_memory_recurrent * mem,1097 std::vector<llama_ubatch> ubatches) : status(LLAMA_MEMORY_STATUS_SUCCESS), mem(mem), ubatches(std::move(ubatches)) {}1098 1099llama_memory_recurrent_context::~llama_memory_recurrent_context() = default;1100 1101bool llama_memory_recurrent_context::next() {1102 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);1103 1104 if (++i_next >= ubatches.size()) {1105 return false;1106 }1107 1108 return true;1109}1110 1111bool llama_memory_recurrent_context::apply() {1112 assert(!llama_memory_status_is_fail(status));1113 1114 // no ubatches -> this is an update1115 if (ubatches.empty()) {1116 // recurrent cache never performs updates1117 assert(status == LLAMA_MEMORY_STATUS_NO_UPDATE);1118 1119 return true;1120 }1121 1122 mem->find_slot(ubatches[i_next]);1123 1124 return true;1125}1126 1127llama_memory_status llama_memory_recurrent_context::get_status() const {1128 return status;1129}1130 1131const llama_ubatch & llama_memory_recurrent_context::get_ubatch() const {1132 assert(status == LLAMA_MEMORY_STATUS_SUCCESS);1133 1134 return ubatches[i_next];1135}1136 1137uint32_t llama_memory_recurrent_context::get_n_rs() const {1138 return is_full ? mem->size : mem->n;1139}1140 1141uint32_t llama_memory_recurrent_context::get_head() const {1142 return is_full ? 0 : mem->head;1143}1144 1145int32_t llama_memory_recurrent_context::get_rs_z() const {1146 return is_full ? 0 : mem->rs_z;1147}1148 1149uint32_t llama_memory_recurrent_context::get_size() const {1150 return mem->size;1151}1152 1153ggml_tensor * llama_memory_recurrent_context::get_r_l(int32_t il) const {1154 return mem->r_l[il];1155}1156 1157ggml_tensor * llama_memory_recurrent_context::get_s_l(int32_t il) const {1158 return mem->s_l[il];1159}1160 1161int32_t llama_memory_recurrent_context::s_copy(int i) const {1162 return mem->cells[i + mem->head].src0;1163}1164 