Team Ai
Datasetpublic

echodict/llama.cpp

version https://git-lfs.github.com/spec/v1 oid sha256:cfc44b7ba25614df70e6b65e3341cae0310163bd32fd31a6b928a542df433faf size 30786

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes610downloads
llama-memory-recurrent.cpp1164 linesDownload Raw Back to src
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