Felipe97/llama-cpp-compiled
01.1k
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 