CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 4d agoView on Hugging Face
0likes1.1kdownloads
llama-kv-cache-iswa.cpp364 linesDownload Raw Back to src
1#include "llama-kv-cache-iswa.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 10//11// llama_kv_cache_iswa12//13 14llama_kv_cache_iswa::llama_kv_cache_iswa(15        const llama_model & model,16                ggml_type   type_k,17                ggml_type   type_v,18                     bool   v_trans,19                     bool   offload,20                     bool   swa_full,21                     bool   unified,22                 uint32_t   kv_size,23                 uint32_t   n_seq_max,24                 uint32_t   n_ubatch,25                 uint32_t   n_pad,26           llama_memory_t   mem_other,27    const layer_filter_cb & filter,28    const  layer_reuse_cb & reuse,29    const  layer_share_cb & share) :30    llama_kv_cache_iswa(model, model.hparams, type_k, type_v, v_trans, offload, swa_full, unified,31            kv_size, n_seq_max, n_ubatch, n_pad, mem_other, filter, reuse, share) {32}33 34llama_kv_cache_iswa::llama_kv_cache_iswa(35        const llama_model & model,36        const llama_hparams & hparams,37                ggml_type   type_k,38                ggml_type   type_v,39                     bool   v_trans,40                     bool   offload,41                     bool   swa_full,42                     bool   unified,43                 uint32_t   kv_size,44                 uint32_t   n_seq_max,45                 uint32_t   n_ubatch,46                 uint32_t   n_pad,47           llama_memory_t   mem_other,48    const layer_filter_cb & filter,49    const  layer_reuse_cb & reuse,50    const  layer_share_cb & share) : unified(unified) {51 52    // chain filters53    const layer_filter_cb filter_base = [&](int32_t il) {54        if (filter && !filter(il)) {55            return false;56        }57 58        return !model.hparams.is_swa(il);59    };60 61    const layer_filter_cb filter_swa  = [&](int32_t il) {62        if (filter && !filter(il)) {63            return false;64        }65 66        return  model.hparams.is_swa(il);67    };68 69    const uint32_t size_base = kv_size;70 71    // note: the SWA cache is always padded to 256 for performance72    //       https://github.com/ggml-org/llama.cpp/issues/1703773    uint32_t size_swa = GGML_PAD(std::min(size_base, hparams.n_swa*(unified ? n_seq_max : 1) + n_ubatch), 256);74 75    // when using full-size SWA cache, we set the SWA cache size to be equal to the base cache size76    if (swa_full) {77        LLAMA_LOG_WARN("%s: using full-size SWA cache (ref: %s)\n",78                __func__, "https://github.com/ggml-org/llama.cpp/pull/13194#issuecomment-2868343055");79 80        size_swa = size_base;81    }82 83    LLAMA_LOG_INFO("%s: creating non-SWA KV cache, size = %u cells\n", __func__, size_base);84 85    llama_memory_t mem_other_base = nullptr;86    if (mem_other) {87        mem_other_base = static_cast<llama_kv_cache_iswa *>(mem_other)->get_base();88    }89 90    llama_memory_t mem_other_swa = nullptr;91    if (mem_other) {92        mem_other_swa = static_cast<llama_kv_cache_iswa *>(mem_other)->get_swa();93    }94 95    kv_base = std::make_unique<llama_kv_cache>(96            model, hparams, type_k, type_v,97            v_trans, offload, unified, size_base, n_seq_max, n_pad,98            0, LLAMA_SWA_TYPE_NONE, mem_other_base, filter_base, reuse, share);99 100    LLAMA_LOG_INFO("%s: creating     SWA KV cache, size = %u cells\n", __func__, size_swa);101 102    kv_swa = std::make_unique<llama_kv_cache>(103            model, hparams, type_k, type_v,104            v_trans, offload, unified, size_swa, n_seq_max, n_pad,105            hparams.n_swa, hparams.swa_type, mem_other_swa, filter_swa, reuse, share);106}107 108void llama_kv_cache_iswa::clear(bool data) {109    kv_base->clear(data);110    kv_swa ->clear(data);111}112 113bool llama_kv_cache_iswa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {114    bool res = true;115 116    res = res & kv_base->seq_rm(seq_id, p0, p1);117    res = res & kv_swa ->seq_rm(seq_id, p0, p1);118 119    return res;120}121 122void llama_kv_cache_iswa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {123    kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1);124    kv_swa ->seq_cp(seq_id_src, seq_id_dst, p0, p1);125}126 127void llama_kv_cache_iswa::seq_keep(llama_seq_id seq_id) {128    kv_base->seq_keep(seq_id);129    kv_swa ->seq_keep(seq_id);130}131 132void llama_kv_cache_iswa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {133    kv_base->seq_add(seq_id, p0, p1, shift);134    kv_swa ->seq_add(seq_id, p0, p1, shift);135}136 137void llama_kv_cache_iswa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {138    kv_base->seq_div(seq_id, p0, p1, d);139    kv_swa ->seq_div(seq_id, p0, p1, d);140}141 142llama_pos llama_kv_cache_iswa::seq_pos_min(llama_seq_id seq_id) const {143    // the base cache is a superset of the SWA cache, so we can just check the SWA cache144    return kv_swa->seq_pos_min(seq_id);145}146 147llama_pos llama_kv_cache_iswa::seq_pos_max(llama_seq_id seq_id) const {148    return kv_swa->seq_pos_max(seq_id);149}150 151std::map<ggml_backend_buffer_type_t, size_t> llama_kv_cache_iswa::memory_breakdown() const {152    std::map<ggml_backend_buffer_type_t, size_t> mb = kv_base->memory_breakdown();153    for (const auto & buft_size : kv_swa->memory_breakdown()) {154        mb[buft_size.first] += buft_size.second;155    }156    return mb;157}158 159llama_memory_context_ptr llama_kv_cache_iswa::init_batch(llama_batch_allocr & balloc, uint32_t n_ubatch, bool embd_all) {160    GGML_UNUSED(embd_all);161 162    // first try simple split163    do {164        if (!unified) {165            // requires equal splits, so we skip the simple split166            break;167        }168 169        balloc.split_reset();170 171        std::vector<llama_ubatch> ubatches;172        while (true) {173            auto ubatch = balloc.split_simple(n_ubatch);174 175            if (ubatch.n_tokens == 0) {176                break;177            }178 179            ubatches.push_back(std::move(ubatch)); // NOLINT180        }181 182        if (balloc.get_n_used() < balloc.get_n_tokens()) {183            // failed to find a suitable split184            break;185        }186 187        auto sinfos_base = kv_base->prepare(ubatches);188        if (sinfos_base.empty()) {189            break;190        }191 192        auto sinfos_swa = kv_swa->prepare(ubatches);193        if (sinfos_swa.empty()) {194            break;195        }196 197        assert(sinfos_base.size() == sinfos_swa.size());198 199        return std::make_unique<llama_kv_cache_iswa_context>(200                this, std::move(sinfos_base), std::move(sinfos_swa), std::move(ubatches));201    } while (false);202 203    // if it fails, try equal split204    do {205        balloc.split_reset();206 207        std::vector<llama_ubatch> ubatches;208        while (true) {209            auto ubatch = balloc.split_equal(n_ubatch, !unified, 0);210 211            if (ubatch.n_tokens == 0) {212                break;213            }214 215            ubatches.push_back(std::move(ubatch)); // NOLINT216        }217 218        if (balloc.get_n_used() < balloc.get_n_tokens()) {219            // failed to find a suitable split220            break;221        }222 223        auto sinfos_base = kv_base->prepare(ubatches);224        if (sinfos_base.empty()) {225            break;226        }227 228        auto sinfos_swa = kv_swa->prepare(ubatches);229        if (sinfos_swa.empty()) {230            break;231        }232 233        assert(sinfos_base.size() == sinfos_swa.size());234 235        return std::make_unique<llama_kv_cache_iswa_context>(236                this, std::move(sinfos_base), std::move(sinfos_swa), std::move(ubatches));237    } while (false);238 239    // TODO: if we fail again, we should attempt different splitting strategies240    //       but to do that properly, we first have to refactor the batches to be more flexible241 242    return std::make_unique<llama_kv_cache_iswa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);243}244 245llama_memory_context_ptr llama_kv_cache_iswa::init_full() {246    return std::make_unique<llama_kv_cache_iswa_context>(this);247}248 249llama_memory_context_ptr llama_kv_cache_iswa::init_update(llama_context * lctx, bool optimize) {250    return std::make_unique<llama_kv_cache_iswa_context>(this, lctx, optimize);251}252 253bool llama_kv_cache_iswa::get_can_shift() const {254    return kv_base->get_can_shift() &&255           kv_swa->get_can_shift() &&256           kv_base->get_size() == kv_swa->get_size();257}258 259void llama_kv_cache_iswa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {260    if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) {261        kv_base->state_write(io, seq_id, flags);262    }263 264    kv_swa->state_write(io, seq_id, flags);265}266 267void llama_kv_cache_iswa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {268    if ((flags & LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) == 0) {269        kv_base->state_read(io, seq_id, flags);270    }271 272    kv_swa->state_read(io, seq_id, flags);273}274 275llama_kv_cache * llama_kv_cache_iswa::get_base() const {276    return kv_base.get();277}278 279llama_kv_cache * llama_kv_cache_iswa::get_swa() const {280    return kv_swa.get();281}282 283//284// llama_kv_cache_iswa_context285//286 287llama_kv_cache_iswa_context::llama_kv_cache_iswa_context(llama_memory_status status) : status(status) {}288 289llama_kv_cache_iswa_context::llama_kv_cache_iswa_context(290        llama_kv_cache_iswa * kv) :291    ctx_base(kv->get_base()->init_full()),292    ctx_swa (kv->get_swa ()->init_full()),293    status(llama_memory_status_combine(ctx_base->get_status(), ctx_swa->get_status())) {294}295 296llama_kv_cache_iswa_context::llama_kv_cache_iswa_context(297        llama_kv_cache_iswa * kv,298        llama_context * lctx,299        bool optimize) :300    ctx_base(kv->get_base()->init_update(lctx, optimize)),301    ctx_swa (kv->get_swa ()->init_update(lctx, optimize)),302    status(llama_memory_status_combine(ctx_base->get_status(), ctx_swa->get_status())) {303}304 305llama_kv_cache_iswa_context::llama_kv_cache_iswa_context(306        llama_kv_cache_iswa * kv,307        slot_info_vec_t sinfos_base,308        slot_info_vec_t sinfos_swa,309        std::vector<llama_ubatch> ubatches) :310    ubatches(std::move(ubatches)),311    // note: here we copy the ubatches. not sure if this is ideal312    ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)),313    ctx_swa (new llama_kv_cache_context(kv->get_swa (), std::move(sinfos_swa),  this->ubatches)),314    status(llama_memory_status_combine(ctx_base->get_status(), ctx_swa->get_status())) {315}316 317llama_kv_cache_iswa_context:: ~llama_kv_cache_iswa_context() = default;318 319bool llama_kv_cache_iswa_context::next() {320    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);321 322    ctx_base->next();323    ctx_swa ->next();324 325    if (++i_next >= ubatches.size()) {326        return false;327    }328 329    return true;330}331 332bool llama_kv_cache_iswa_context::apply() {333    assert(!llama_memory_status_is_fail(status));334 335    bool res = true;336 337    res = res & ctx_base->apply();338    res = res & ctx_swa ->apply();339 340    return res;341}342 343llama_memory_status llama_kv_cache_iswa_context::get_status() const {344    return status;345}346 347const llama_ubatch & llama_kv_cache_iswa_context::get_ubatch() const {348    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);349 350    return ubatches[i_next];351}352 353const llama_kv_cache_context * llama_kv_cache_iswa_context::get_base() const {354    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);355 356    return static_cast<const llama_kv_cache_context *>(ctx_base.get());357}358 359const llama_kv_cache_context * llama_kv_cache_iswa_context::get_swa()  const {360    assert(status == LLAMA_MEMORY_STATUS_SUCCESS);361 362    return static_cast<const llama_kv_cache_context *>(ctx_swa.get());363}364