Team Ai
Datasetpublic

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.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes3.1kdownloads
llama-kv-cache-dsv4.cpp2196 linesDownload Raw Back to src
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 

Showing the first 1,200 of 2196 lines. Download the file for the rest.

Brunobkr/llama.cpp_AlgMor24_github · Team Ai