CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 2d agoView on Hugging Face
0likes1.1kdownloads
llama-kv-cache-msa.cpp396 linesDownload Raw Back to src
1#include "llama-kv-cache-msa.h"2 3#include "llama-impl.h"4#include "llama-batch.h"5#include "llama-model.h"6 7#include <algorithm>8#include <cassert>9#include <cmath>10 11// llama_kv_cache_msa12 13llama_kv_cache_msa::llama_kv_cache_msa(14        const llama_model & model,15                ggml_type   type_k,16                ggml_type   type_v,17                     bool   v_trans,18                     bool   offload,19                     bool   unified,20                 uint32_t   kv_size,21                 uint32_t   n_seq_max,22                 uint32_t   n_pad,23                 uint32_t   n_swa,24           llama_swa_type   swa_type,25    const layer_filter_cb & filter,26    const layer_filter_cb & filter_idx,27    const  layer_reuse_cb & reuse) :28    hparams_idx(model.hparams),29    n_stream(unified ? 1 : n_seq_max), n_seq_max(n_seq_max), n_pad(n_pad),30    n_swa(n_swa), swa_type(swa_type) {31 32    LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size);33 34    kv_base = std::make_unique<llama_kv_cache>(35            model, model.hparams, type_k, type_v,36            v_trans, offload, unified, kv_size, n_seq_max, n_pad,37            n_swa, swa_type, nullptr, filter, reuse, nullptr);38 39    // the MSA indexer uses a single key head per layer40    std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1);41    hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size;42    // the rope parameters are kept identical to the main cache43 44    LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size);45 46    kv_idx = std::make_unique<llama_kv_cache>(47            model, hparams_idx, type_k, type_v,48            v_trans, offload, unified, kv_size, n_seq_max, n_pad,49            n_swa, swa_type, nullptr, filter_idx, reuse, nullptr);50}51 52void llama_kv_cache_msa::clear(bool data) {53    kv_base->clear(data);54    kv_idx ->clear(data);55}56 57bool llama_kv_cache_msa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {58    bool res = true;59 60    res = res & kv_base->seq_rm(seq_id, p0, p1);61    res = res & kv_idx ->seq_rm(seq_id, p0, p1);62 63    return res;64}65 66void llama_kv_cache_msa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {67    kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1);68    kv_idx ->seq_cp(seq_id_src, seq_id_dst, p0, p1);69}70 71void llama_kv_cache_msa::seq_keep(llama_seq_id seq_id) {72    kv_base->seq_keep(seq_id);73    kv_idx ->seq_keep(seq_id);74}75 76void llama_kv_cache_msa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {77    kv_base->seq_add(seq_id, p0, p1, shift);78    kv_idx ->seq_add(seq_id, p0, p1, shift);79}80 81void llama_kv_cache_msa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {82    kv_base->seq_div(seq_id, p0, p1, d);83    kv_idx ->seq_div(seq_id, p0, p1, d);84}85 86llama_pos llama_kv_cache_msa::seq_pos_min(llama_seq_id seq_id) const {87    return kv_base->seq_pos_min(seq_id);88}89 90llama_pos llama_kv_cache_msa::seq_pos_max(llama_seq_id seq_id) const {91    return kv_base->seq_pos_max(seq_id);92}93 94std::map<ggml_backend_buffer_type_t, size_t> llama_kv_cache_msa::memory_breakdown() const {95    std::map<ggml_backend_buffer_type_t, size_t> mb = kv_base->memory_breakdown();96    for (const auto & buft_size : kv_idx->memory_breakdown()) {97        mb[buft_size.first] += buft_size.second;98    }99    return mb;100}101 102llama_memory_context_ptr llama_kv_cache_msa::init_batch(103            llama_batch_allocr & balloc,104            uint32_t n_ubatch,105            bool embd_all) {106    GGML_UNUSED(embd_all);107 108    do {109        balloc.split_reset();110 111        std::vector<llama_ubatch> ubatches;112        while (true) {113            auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0);114 115            if (ubatch.n_tokens == 0) {116                break;117            }118 119            ubatches.push_back(std::move(ubatch));120        }121 122        if (balloc.get_n_used() < balloc.get_n_tokens()) {123            // failed to find a suitable split124            break;125        }126 127        auto sinfos_base = kv_base->prepare(ubatches);128        if (sinfos_base.empty()) {129            break;130        }131 132        auto sinfos_idx = kv_idx->prepare(ubatches);133        if (sinfos_idx.empty()) {134            break;135        }136 137        assert(sinfos_base.size() == sinfos_idx.size());138 139        return std::make_unique<llama_kv_cache_msa_context>(140                this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches));141    } while (false);142 143    return std::make_unique<llama_kv_cache_msa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);144}145 146llama_memory_context_ptr llama_kv_cache_msa::init_full() {147    return std::make_unique<llama_kv_cache_msa_context>(this);148}149 150llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) {151    return std::make_unique<llama_kv_cache_msa_context>(this, lctx, optimize);152}153 154bool llama_kv_cache_msa::get_can_shift() const {155    return kv_base->get_can_shift() &&156           kv_idx ->get_can_shift() &&157           kv_base->get_size() == kv_idx->get_size();158}159 160void llama_kv_cache_msa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {161    kv_base->state_write(io, seq_id, flags);162    kv_idx ->state_write(io, seq_id, flags);163}164 165void llama_kv_cache_msa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {166    kv_base->state_read(io, seq_id, flags);167    kv_idx ->state_read(io, seq_id, flags);168}169 170llama_kv_cache * llama_kv_cache_msa::get_base() const {171    return kv_base.get();172}173 174llama_kv_cache * llama_kv_cache_msa::get_idx() const {175    return kv_idx.get();176}177 178// llama_kv_cache_msa_context179 180llama_kv_cache_msa_context::llama_kv_cache_msa_context(llama_memory_status status) :181    kv(nullptr), status(status) {}182 183llama_kv_cache_msa_context::llama_kv_cache_msa_context(184        llama_kv_cache_msa * kv) :185    kv(kv),186    ctx_base(kv->get_base()->init_full()),187    ctx_idx (kv->get_idx ()->init_full()),188    status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {189}190 191llama_kv_cache_msa_context::llama_kv_cache_msa_context(192        llama_kv_cache_msa * kv,193        llama_context * lctx,194        bool optimize) :195    kv(kv),196    ctx_base(kv->get_base()->init_update(lctx, optimize)),197    ctx_idx (kv->get_idx ()->init_update(lctx, optimize)),198    status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {199}200 201llama_kv_cache_msa_context::llama_kv_cache_msa_context(202        llama_kv_cache_msa * kv,203        slot_info_vec_t sinfos_base,204        slot_info_vec_t sinfos_idx,205        std::vector<llama_ubatch> ubatches) :206    kv(kv),207    ubatches(std::move(ubatches)),208    // here we copy the ubatches. not sure if this is ideal209    ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)),210    ctx_idx (new llama_kv_cache_context(kv->get_idx (), std::move(sinfos_idx),  this->ubatches)),211    status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {212}213 214llama_kv_cache_msa_context::~llama_kv_cache_msa_context() = default;215 216bool llama_kv_cache_msa_context::next() {217    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);218 219    ctx_base->next();220    ctx_idx ->next();221 222    if (++i_next >= ubatches.size()) {223        return false;224    }225 226    return true;227}228 229bool llama_kv_cache_msa_context::apply() {230    assert(!llama_memory_status_is_fail(status));231 232    bool res = true;233 234    res = res & ctx_base->apply();235    res = res & ctx_idx ->apply();236 237    return res;238}239 240llama_memory_status llama_kv_cache_msa_context::get_status() const {241    return status;242}243 244const llama_ubatch & llama_kv_cache_msa_context::get_ubatch() const {245    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);246 247    return ubatches[i_next];248}249 250const llama_kv_cache_context * llama_kv_cache_msa_context::get_base() const {251    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);252 253    return static_cast<const llama_kv_cache_context *>(ctx_base.get());254}255 256const llama_kv_cache_context * llama_kv_cache_msa_context::get_idx() const {257    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);258 259    return static_cast<const llama_kv_cache_context *>(ctx_idx.get());260}261 262uint32_t llama_kv_cache_msa_context::get_n_pos() const {263    // pad the value so that the graph remains constant across batches and can be reused264    const uint32_t n_pad_cur = std::max(kv->get_n_pad(), 256u);265 266    llama_pos pos_max = -1;267 268    for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) kv->get_n_seq_max(); ++seq_id) {269        pos_max = std::max(pos_max, kv->seq_pos_max(seq_id));270    }271 272    return std::max(n_pad_cur, GGML_PAD((uint32_t) (pos_max + 1), n_pad_cur));273}274 275void llama_kv_cache_msa_context::set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const {276    GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));277    GGML_ASSERT(dst->type == GGML_TYPE_I32);278    GGML_ASSERT(div > 0);279 280    const int64_t n_tokens    = ubatch->n_tokens;281    const int64_t n_kv        = dst->ne[0];282    const int64_t n_stream_ub = dst->ne[1];283 284    GGML_ASSERT(n_tokens % n_stream_ub == 0);285    const int64_t n_tps = n_tokens/n_stream_ub;286 287    int32_t * data = (int32_t *) dst->data;288 289    for (int64_t s = 0; s < n_stream_ub; ++s) {290        const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];291 292        const auto & cells = kv->get_base()->get_cells(seq_id);293 294        for (int64_t j = 0; j < n_kv; ++j) {295            // the value for empty or other-sequence cells is irrelevant as consumers mask them296            data[s*n_kv + j] =297                cells.is_empty(j) || !cells.seq_has(j, seq_id)298                    ? 0299                    : (int32_t) (cells.pos_get(j)/div);300        }301    }302}303 304void llama_kv_cache_msa_context::set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const {305    GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));306    GGML_ASSERT(dst->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_F32);307 308    const int64_t n_tokens    = ubatch->n_tokens;309    const int64_t n_pos       = dst->ne[0];310    const int64_t n_stream_ub = dst->ne[1];311 312    GGML_ASSERT(n_tokens % n_stream_ub == 0);313    const int64_t n_tps = n_tokens/n_stream_ub;314 315    for (int64_t s = 0; s < n_stream_ub; ++s) {316        const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];317 318        const auto & cells = kv->get_base()->get_cells(seq_id);319 320        std::vector<int32_t> map(n_pos, 0);321 322        for (uint32_t j = 0; j < cells.size(); ++j) {323            if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {324                continue;325            }326 327            const llama_pos p0 = cells.pos_get(j);328 329            if (p0 < 0 || p0 >= n_pos) {330                continue;331            }332 333            map[p0] = (int32_t) j;334        }335 336        if (dst->type == GGML_TYPE_I32) {337            int32_t * data = (int32_t *) dst->data + s*n_pos;338            std::copy(map.begin(), map.end(), data);339        } else {340            float * data = (float *) dst->data + s*n_pos;341            for (int64_t p = 0; p < n_pos; ++p) {342                data[p] = (float) map[p];343            }344        }345    }346}347 348void llama_kv_cache_msa_context::set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const {349    GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));350    GGML_ASSERT(dst->type == GGML_TYPE_F32);351 352    const int64_t n_tokens = ubatch->n_tokens;353    const int64_t n_pos    = dst->ne[0];354 355    GGML_ASSERT(dst->ne[1] == n_tokens);356 357    const uint32_t       n_swa    = kv->get_n_swa();358    const llama_swa_type swa_type = kv->get_swa_type();359 360    float * data = (float *) dst->data;361 362    std::fill(data, data + n_pos*n_tokens, -INFINITY);363 364    for (int64_t i = 0; i < n_tokens; ++i) {365        const llama_seq_id seq_id = ubatch->seq_id[i][0];366 367        const auto & cells = kv->get_base()->get_cells(seq_id);368 369        const llama_pos p1 = ubatch->pos[i];370 371        for (uint32_t j = 0; j < cells.size(); ++j) {372            if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {373                continue;374            }375 376            const llama_pos p0 = cells.pos_get(j);377 378            if (p0 < 0 || p0 >= n_pos) {379                continue;380            }381 382            // causal mask383            if (p0 > p1) {384                continue;385            }386 387            // apply SWA if any388            if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {389                continue;390            }391 392            data[i*n_pos + p0] = 0.0f;393        }394    }395}396