diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt index 320784c3a8..24f05cc916 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -25,6 +25,7 @@ add_library(llama llama-kv-cache.cpp llama-kv-cache-iswa.cpp llama-kv-cache-dsa.cpp + llama-kv-cache-msa.cpp llama-kv-cache-dsv4.cpp llama-memory.cpp llama-memory-hybrid.cpp diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 1a35692300..fdab7b8dde 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -8,6 +8,7 @@ #include "llama-kv-cache.h" #include "llama-kv-cache-iswa.h" #include "llama-kv-cache-dsa.h" +#include "llama-kv-cache-msa.h" #include "llama-kv-cache-dsv4.h" #include "llama-memory-hybrid.h" #include "llama-memory-hybrid-iswa.h" @@ -518,6 +519,36 @@ bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) { return res; } +llm_graph_input_attn_kv_msa::llm_graph_input_attn_kv_msa( + const llama_hparams & hparams, + const llama_cparams & cparams, + const llama_kv_cache_msa_context * mctx) : + llm_graph_input_attn_kv(hparams, cparams, mctx->get_base()), + mctx_msa(mctx) { +} + +void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) { + llm_graph_input_attn_kv::set_input(ubatch); + + mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch); +} + +bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) { + mctx_msa = static_cast(params.mctx); + + // the parent class operates on the base cache context + this->mctx = mctx_msa->get_base(); + + bool res = true; + + res &= self_k_idxs ->ne[0] == params.ubatch.n_tokens; + res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens; + + res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams); + + return res; +} + void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) { mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch); @@ -3187,6 +3218,32 @@ llm_graph_input_attn_k_dsa * llm_graph_context::build_attn_inp_k_dsa() const { return (llm_graph_input_attn_k_dsa *) res->add_input(std::move(inp)); } +llm_graph_input_attn_kv_msa * llm_graph_context::build_attn_inp_kv_msa() const { + const auto * mctx_cur = static_cast(mctx); + + auto inp = std::make_unique(hparams, cparams, mctx_cur); + + const auto * mctx_base = mctx_cur->get_base(); + const auto * mctx_idx = mctx_cur->get_idx(); + + { + GGML_ASSERT(hparams.swa_type == LLAMA_SWA_TYPE_NONE && "Use llama_kv_cache_iswa for SWA"); + + inp->self_k_idxs = mctx_base->build_input_k_idxs(ctx0, ubatch); + inp->self_v_idxs = mctx_base->build_input_v_idxs(ctx0, ubatch); + + inp->self_kq_mask = build_attn_inp_kq_mask(ctx0, mctx_base, ubatch, cparams); + inp->self_kq_mask_cnv = inp->self_kq_mask; + } + + inp->self_k_rot = mctx_base->build_input_k_rot(ctx0); + inp->self_v_rot = mctx_base->build_input_v_rot(ctx0); + + inp->self_k_idxs_idx = mctx_idx->build_input_k_idxs(ctx0, ubatch); + + return (llm_graph_input_attn_kv_msa *) res->add_input(std::move(inp)); +} + // TODO: maybe separate the inner implementation into a separate function // like with the non-sliding window equivalent // once sliding-window hybrid caches are a thing. diff --git a/src/llama-graph.h b/src/llama-graph.h index 160e294135..ff216302db 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -23,6 +23,7 @@ struct llama_memory_context_i; class llama_kv_cache_context; class llama_kv_cache_dsa_context; +class llama_kv_cache_msa_context; class llama_kv_cache_dsv4_raw_context; class llama_kv_cache_dsv4_context; class llama_kv_cache_iswa_context; @@ -425,6 +426,26 @@ public: const llama_kv_cache_dsa_context * mctx; }; +// standard K/V attention input against the base cache, plus destination indices for the indexer key cache +class llm_graph_input_attn_kv_msa : public llm_graph_input_attn_kv { +public: + llm_graph_input_attn_kv_msa( + const llama_hparams & hparams, + const llama_cparams & cparams, + const llama_kv_cache_msa_context * mctx); + ~llm_graph_input_attn_kv_msa() = default; + + void set_input(const llama_ubatch * ubatch) override; + + bool can_reuse(const llm_graph_params & params) override; + + ggml_tensor * get_k_idxs_idx() const { return self_k_idxs_idx; } + + ggml_tensor * self_k_idxs_idx = nullptr; // I64 [n_batch] + + const llama_kv_cache_msa_context * mctx_msa; +}; + class llm_graph_input_attn_kv_iswa : public llm_graph_input_i { public: llm_graph_input_attn_kv_iswa( @@ -1169,6 +1190,8 @@ struct llm_graph_context { llm_graph_input_attn_k_dsa * build_attn_inp_k_dsa() const; + llm_graph_input_attn_kv_msa * build_attn_inp_kv_msa() const; + ggml_tensor * build_attn( llm_graph_input_attn_k_dsa * inp, ggml_tensor * wo, diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index 50af97f358..846d4c69a6 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -180,16 +180,6 @@ uint32_t llama_hparams::n_embd_v_gqa_max() const { return val; } -uint32_t llama_hparams::n_embd_k_idx(uint32_t il) const { - if (!indexer_kv || indexer_head_size == 0) { - return 0; // arch without a MSA indexer - } - if (il < n_layer_dense_lead) { - return 0; // leading dense layers carry no indexer - } - return indexer_head_size; // 128 -} - uint32_t llama_hparams::n_embd_r() const { if (wkv_head_size != 0) { // for RWKV models diff --git a/src/llama-hparams.h b/src/llama-hparams.h index fc770bf003..6e8336c987 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -230,8 +230,6 @@ struct llama_hparams { // MSA uint32_t indexer_block_size = 0; uint32_t indexer_local_blocks = 0; - // MSA stores its indexer keys in the main KV cache (k_idx tensors); - bool indexer_kv = false; // Indexer is "full" (1) or "shared" (0) // Shared indexers reuse top-k from previous full layer @@ -356,9 +354,6 @@ struct llama_hparams { uint32_t n_embd_k_gqa_max() const; uint32_t n_embd_v_gqa_max() const; - // dimension of the single-head MSA indexer key stream - uint32_t n_embd_k_idx(uint32_t il = 0) const; - // dimension of the rolling state embeddings // corresponds to Mamba's conv_states size or RWKV's token_shift states size uint32_t n_embd_r() const; diff --git a/src/llama-kv-cache-msa.cpp b/src/llama-kv-cache-msa.cpp new file mode 100644 index 0000000000..55ef286caf --- /dev/null +++ b/src/llama-kv-cache-msa.cpp @@ -0,0 +1,395 @@ +#include "llama-kv-cache-msa.h" + +#include "llama-impl.h" +#include "llama-batch.h" +#include "llama-model.h" + +#include +#include +#include + +// llama_kv_cache_msa + +llama_kv_cache_msa::llama_kv_cache_msa( + const llama_model & model, + ggml_type type_k, + ggml_type type_v, + bool v_trans, + bool offload, + bool unified, + uint32_t kv_size, + uint32_t n_seq_max, + uint32_t n_pad, + uint32_t n_swa, + llama_swa_type swa_type, + const layer_filter_cb & filter, + const layer_filter_cb & filter_idx, + const layer_reuse_cb & reuse) : + hparams_idx(model.hparams), + n_stream(unified ? 1 : n_seq_max), n_seq_max(n_seq_max), n_pad(n_pad), + n_swa(n_swa), swa_type(swa_type) { + + LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size); + + kv_base = std::make_unique( + model, model.hparams, type_k, type_v, + v_trans, offload, unified, kv_size, n_seq_max, n_pad, + n_swa, swa_type, nullptr, filter, reuse, nullptr); + + // the MSA indexer uses a single key head per layer + std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1); + hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size; + // the rope parameters are kept identical to the main cache + + LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size); + + kv_idx = std::make_unique( + model, hparams_idx, type_k, type_v, + v_trans, offload, unified, kv_size, n_seq_max, n_pad, + n_swa, swa_type, nullptr, filter_idx, reuse, nullptr); +} + +void llama_kv_cache_msa::clear(bool data) { + kv_base->clear(data); + kv_idx ->clear(data); +} + +bool llama_kv_cache_msa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { + bool res = true; + + res = res & kv_base->seq_rm(seq_id, p0, p1); + res = res & kv_idx ->seq_rm(seq_id, p0, p1); + + return res; +} + +void llama_kv_cache_msa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) { + kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1); + kv_idx ->seq_cp(seq_id_src, seq_id_dst, p0, p1); +} + +void llama_kv_cache_msa::seq_keep(llama_seq_id seq_id) { + kv_base->seq_keep(seq_id); + kv_idx ->seq_keep(seq_id); +} + +void llama_kv_cache_msa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) { + kv_base->seq_add(seq_id, p0, p1, shift); + kv_idx ->seq_add(seq_id, p0, p1, shift); +} + +void llama_kv_cache_msa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) { + kv_base->seq_div(seq_id, p0, p1, d); + kv_idx ->seq_div(seq_id, p0, p1, d); +} + +llama_pos llama_kv_cache_msa::seq_pos_min(llama_seq_id seq_id) const { + return kv_base->seq_pos_min(seq_id); +} + +llama_pos llama_kv_cache_msa::seq_pos_max(llama_seq_id seq_id) const { + return kv_base->seq_pos_max(seq_id); +} + +std::map llama_kv_cache_msa::memory_breakdown() const { + std::map mb = kv_base->memory_breakdown(); + for (const auto & buft_size : kv_idx->memory_breakdown()) { + mb[buft_size.first] += buft_size.second; + } + return mb; +} + +llama_memory_context_ptr llama_kv_cache_msa::init_batch( + llama_batch_allocr & balloc, + uint32_t n_ubatch, + bool embd_all) { + GGML_UNUSED(embd_all); + + do { + balloc.split_reset(); + + std::vector ubatches; + while (true) { + auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0); + + if (ubatch.n_tokens == 0) { + break; + } + + ubatches.push_back(std::move(ubatch)); + } + + if (balloc.get_n_used() < balloc.get_n_tokens()) { + // failed to find a suitable split + break; + } + + auto sinfos_base = kv_base->prepare(ubatches); + if (sinfos_base.empty()) { + break; + } + + auto sinfos_idx = kv_idx->prepare(ubatches); + if (sinfos_idx.empty()) { + break; + } + + assert(sinfos_base.size() == sinfos_idx.size()); + + return std::make_unique( + this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches)); + } while (false); + + return std::make_unique(LLAMA_MEMORY_STATUS_FAILED_PREPARE); +} + +llama_memory_context_ptr llama_kv_cache_msa::init_full() { + return std::make_unique(this); +} + +llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) { + return std::make_unique(this, lctx, optimize); +} + +bool llama_kv_cache_msa::get_can_shift() const { + return kv_base->get_can_shift() && + kv_idx ->get_can_shift() && + kv_base->get_size() == kv_idx->get_size(); +} + +void llama_kv_cache_msa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { + kv_base->state_write(io, seq_id, flags); + kv_idx ->state_write(io, seq_id, flags); +} + +void llama_kv_cache_msa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { + kv_base->state_read(io, seq_id, flags); + kv_idx ->state_read(io, seq_id, flags); +} + +llama_kv_cache * llama_kv_cache_msa::get_base() const { + return kv_base.get(); +} + +llama_kv_cache * llama_kv_cache_msa::get_idx() const { + return kv_idx.get(); +} + +// llama_kv_cache_msa_context + +llama_kv_cache_msa_context::llama_kv_cache_msa_context(llama_memory_status status) : + kv(nullptr), status(status) {} + +llama_kv_cache_msa_context::llama_kv_cache_msa_context( + llama_kv_cache_msa * kv) : + kv(kv), + ctx_base(kv->get_base()->init_full()), + ctx_idx (kv->get_idx ()->init_full()), + status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) { +} + +llama_kv_cache_msa_context::llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + llama_context * lctx, + bool optimize) : + kv(kv), + ctx_base(kv->get_base()->init_update(lctx, optimize)), + ctx_idx (kv->get_idx ()->init_update(lctx, optimize)), + status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) { +} + +llama_kv_cache_msa_context::llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + slot_info_vec_t sinfos_base, + slot_info_vec_t sinfos_idx, + std::vector ubatches) : + kv(kv), + ubatches(std::move(ubatches)), + // here we copy the ubatches. not sure if this is ideal + ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)), + ctx_idx (new llama_kv_cache_context(kv->get_idx (), std::move(sinfos_idx), this->ubatches)), + status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) { +} + +llama_kv_cache_msa_context::~llama_kv_cache_msa_context() = default; + +bool llama_kv_cache_msa_context::next() { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + ctx_base->next(); + ctx_idx ->next(); + + if (++i_next >= ubatches.size()) { + return false; + } + + return true; +} + +bool llama_kv_cache_msa_context::apply() { + assert(!llama_memory_status_is_fail(status)); + + bool res = true; + + res = res & ctx_base->apply(); + res = res & ctx_idx ->apply(); + + return res; +} + +llama_memory_status llama_kv_cache_msa_context::get_status() const { + return status; +} + +const llama_ubatch & llama_kv_cache_msa_context::get_ubatch() const { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + return ubatches[i_next]; +} + +const llama_kv_cache_context * llama_kv_cache_msa_context::get_base() const { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + return static_cast(ctx_base.get()); +} + +const llama_kv_cache_context * llama_kv_cache_msa_context::get_idx() const { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + return static_cast(ctx_idx.get()); +} + +uint32_t llama_kv_cache_msa_context::get_n_pos() const { + // pad the value so that the graph remains constant across batches and can be reused + const uint32_t n_pad_cur = std::max(kv->get_n_pad(), 256u); + + llama_pos pos_max = -1; + + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) kv->get_n_seq_max(); ++seq_id) { + pos_max = std::max(pos_max, kv->seq_pos_max(seq_id)); + } + + return std::max(n_pad_cur, GGML_PAD((uint32_t) (pos_max + 1), n_pad_cur)); +} + +void llama_kv_cache_msa_context::set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(dst->type == GGML_TYPE_I32); + GGML_ASSERT(div > 0); + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_kv = dst->ne[0]; + const int64_t n_stream_ub = dst->ne[1]; + + GGML_ASSERT(n_tokens % n_stream_ub == 0); + const int64_t n_tps = n_tokens/n_stream_ub; + + int32_t * data = (int32_t *) dst->data; + + for (int64_t s = 0; s < n_stream_ub; ++s) { + const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0]; + + const auto & cells = kv->get_base()->get_cells(seq_id); + + for (int64_t j = 0; j < n_kv; ++j) { + // the value for empty or other-sequence cells is irrelevant as consumers mask them + data[s*n_kv + j] = + cells.is_empty(j) || !cells.seq_has(j, seq_id) + ? 0 + : (int32_t) (cells.pos_get(j)/div); + } + } +} + +void llama_kv_cache_msa_context::set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(dst->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_F32); + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_pos = dst->ne[0]; + const int64_t n_stream_ub = dst->ne[1]; + + GGML_ASSERT(n_tokens % n_stream_ub == 0); + const int64_t n_tps = n_tokens/n_stream_ub; + + for (int64_t s = 0; s < n_stream_ub; ++s) { + const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0]; + + const auto & cells = kv->get_base()->get_cells(seq_id); + + std::vector map(n_pos, 0); + + for (uint32_t j = 0; j < cells.size(); ++j) { + if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) { + continue; + } + + const llama_pos p0 = cells.pos_get(j); + + if (p0 < 0 || p0 >= n_pos) { + continue; + } + + map[p0] = (int32_t) j; + } + + if (dst->type == GGML_TYPE_I32) { + int32_t * data = (int32_t *) dst->data + s*n_pos; + std::copy(map.begin(), map.end(), data); + } else { + float * data = (float *) dst->data + s*n_pos; + for (int64_t p = 0; p < n_pos; ++p) { + data[p] = (float) map[p]; + } + } + } +} + +void llama_kv_cache_msa_context::set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_pos = dst->ne[0]; + + GGML_ASSERT(dst->ne[1] == n_tokens); + + const uint32_t n_swa = kv->get_n_swa(); + const llama_swa_type swa_type = kv->get_swa_type(); + + float * data = (float *) dst->data; + + std::fill(data, data + n_pos*n_tokens, -INFINITY); + + for (int64_t i = 0; i < n_tokens; ++i) { + const llama_seq_id seq_id = ubatch->seq_id[i][0]; + + const auto & cells = kv->get_base()->get_cells(seq_id); + + const llama_pos p1 = ubatch->pos[i]; + + for (uint32_t j = 0; j < cells.size(); ++j) { + if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) { + continue; + } + + const llama_pos p0 = cells.pos_get(j); + + if (p0 < 0 || p0 >= n_pos) { + continue; + } + + // causal mask + if (p0 > p1) { + continue; + } + + // apply SWA if any + if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) { + continue; + } + + data[i*n_pos + p0] = 0.0f; + } + } +} diff --git a/src/llama-kv-cache-msa.h b/src/llama-kv-cache-msa.h new file mode 100644 index 0000000000..f09b6d32b0 --- /dev/null +++ b/src/llama-kv-cache-msa.h @@ -0,0 +1,153 @@ +#pragma once + +#include "llama-kv-cache.h" + +#include + +// llama_kv_cache_msa + +// uses two instances of llama_kv_cache, one for K/V tensors, and one for the MSA indexer tensors +// both receive identical sequence operations and identical ubatches, so their cell layouts stay in synced. +// the context also exposes per-ubatch pos - cell translation maps populated from llama_kv_cells via +// llama_kv_cache::get_cells(), which the model graph uses to run MSA block selection in position space + +class llama_kv_cache_msa : public llama_memory_i { +public: + llama_kv_cache_msa( + const llama_model & model, + ggml_type type_k, + ggml_type type_v, + bool v_trans, + bool offload, + bool unified, + uint32_t kv_size, + uint32_t n_seq_max, + uint32_t n_pad, + uint32_t n_swa, + llama_swa_type swa_type, + const layer_filter_cb & filter, + const layer_filter_cb & filter_idx, + const layer_reuse_cb & reuse); + + ~llama_kv_cache_msa() = default; + + // llama_memory_i + + llama_memory_context_ptr init_batch( + llama_batch_allocr & balloc, + uint32_t n_ubatch, + bool embd_all) override; + + llama_memory_context_ptr init_full() override; + + llama_memory_context_ptr init_update(llama_context * lctx, bool optimize) override; + + bool get_can_shift() const override; + + void clear(bool data) override; + + bool seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) override; + void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) override; + void seq_keep(llama_seq_id seq_id) override; + void seq_add (llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) override; + void seq_div (llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) override; + + llama_pos seq_pos_min(llama_seq_id seq_id) const override; + llama_pos seq_pos_max(llama_seq_id seq_id) const override; + + std::map memory_breakdown() const override; + + // state write/load + + void state_write(llama_io_write_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) const override; + void state_read (llama_io_read_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) override; + + // llama_kv_cache_msa specific API + + llama_kv_cache * get_base() const; + llama_kv_cache * get_idx () const; + + uint32_t get_n_pad() const { return n_pad; } + uint32_t get_n_seq_max() const { return n_seq_max; } + uint32_t get_n_swa() const { return n_swa; } + llama_swa_type get_swa_type() const { return swa_type; } + +private: + // keep the indexer KV cache hparams instance here as llama_kv_cache stores only a reference + llama_hparams hparams_idx; + + const uint32_t n_stream = 1; + const uint32_t n_seq_max = 1; + const uint32_t n_pad = 1; + + const uint32_t n_swa = 0; + const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE; + + std::unique_ptr kv_base; + std::unique_ptr kv_idx; +}; + +class llama_kv_cache_msa_context : public llama_memory_context_i { +public: + using slot_info_vec_t = llama_kv_cache::slot_info_vec_t; + + // used for errors + llama_kv_cache_msa_context(llama_memory_status status); + + // used to create a full-cache context + llama_kv_cache_msa_context( + llama_kv_cache_msa * kv); + + // used to create an update context + llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + llama_context * lctx, + bool optimize); + + // used to create a batch processing context from a batch + llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + slot_info_vec_t sinfos_base, + slot_info_vec_t sinfos_idx, + std::vector ubatches); + + virtual ~llama_kv_cache_msa_context(); + + // llama_memory_context_i + + bool next() override; + bool apply() override; + + llama_memory_status get_status() const override; + const llama_ubatch & get_ubatch() const override; + + // llama_kv_cache_msa_context specific API + + const llama_kv_cache_context * get_base() const; + const llama_kv_cache_context * get_idx () const; + + // max position currently present in the cache plus one, padded MSA blocks are defined over token positions + // so the block-selection tensors are sized by this value rather than by the number of cells + uint32_t get_n_pos() const; + + // position <-> cell translation maps, populated from the base cache cells + // the model graph relates cache contents to token positions only through these per ubatch inputs + // value for empty or other-sequence cells is 0 so consumers must mask them + void set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const; + // positions without a cell map to cell 0, consumers must mask them assumes one sequence per stream + void set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const; + void set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const; + +private: + llama_kv_cache_msa * kv; + + // the index of the next ubatch to process + size_t i_next = 0; + + std::vector ubatches; + + const llama_memory_context_ptr ctx_base; + const llama_memory_context_ptr ctx_idx; + + const llama_memory_status status; +}; diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 44cb1668da..8678a326d9 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -112,7 +112,7 @@ llama_kv_cache::llama_kv_cache( auto it = ctx_map.find(buft); if (it == ctx_map.end()) { ggml_init_params params = { - /*.mem_size =*/ size_t(3u*(1 + n_stream)*n_layer*ggml_tensor_overhead()), //Reserve tensor metadata for up to 3 tensors per layer (K, V, and optional K_idx), plus one view per tensor per stream. + /*.mem_size =*/ size_t(2u*(1 + n_stream)*n_layer*ggml_tensor_overhead()), /*.mem_buffer =*/ NULL, /*.no_alloc =*/ true, }; @@ -242,25 +242,9 @@ llama_kv_cache::llama_kv_cache( v_stream.push_back(has_v ? ggml_view_2d(ctx, v, n_embd_v_gqa, kv_size, v->nb[1], s*v->nb[2]) : nullptr); } - const uint32_t n_embd_k_idx = hparams.n_embd_k_idx(il); - ggml_tensor * k_idx = n_embd_k_idx > 0 - ? ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd_k_idx, kv_size, n_stream) - : nullptr; - if (k_idx) { - ggml_format_name(k_idx, "cache_k_idx_l%d", il); - msa_strict_slots = (n_stream == n_seq_max); - } - - std::vector k_idx_stream; - for (uint32_t s = 0; s < n_stream; ++s) { - k_idx_stream.push_back(k_idx - ? ggml_view_2d(ctx, k_idx, n_embd_k_idx, kv_size, k_idx->nb[1], s*k_idx->nb[2]) - : nullptr); - } - map_layer_ids[il] = layers.size(); - layers.push_back({ il, k, v, k_idx, k_stream, v_stream, k_idx_stream }); + layers.push_back({ il, k, v, k_stream, v_stream, }); } if (reuse) { @@ -309,24 +293,13 @@ llama_kv_cache::llama_kv_cache( } { - const size_t memory_size_k = size_k_bytes(); - const size_t memory_size_v = size_v_bytes(); - const size_t memory_size_k_idx = size_k_idx_bytes(); - const size_t memory_size_total = memory_size_k + memory_size_v + memory_size_k_idx; + const size_t memory_size_k = size_k_bytes(); + const size_t memory_size_v = size_v_bytes(); - constexpr float mib = 1024.0f * 1024.0f; - - const std::string k_log = format(", K (%s): %7.2f MiB", ggml_type_name(type_k), (float) memory_size_k / mib); - const std::string v_log = format(", V (%s): %7.2f MiB", ggml_type_name(type_v), (float) memory_size_v / mib); - - std::string k_idx_log; - if (memory_size_k_idx > 0) { - k_idx_log = format(", K_idx (%s): %7.2f MiB", ggml_type_name(GGML_TYPE_F32), (float) memory_size_k_idx / mib); - } - - LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs)%s%s%s\n", __func__, - (float) memory_size_total / mib, kv_size, (int) layers.size(), n_seq_max, n_stream, - k_log.c_str(), v_log.c_str(), k_idx_log.c_str()); + LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs), K (%s): %7.2f MiB, V (%s): %7.2f MiB\n", __func__, + (float)(memory_size_k + memory_size_v) / (1024.0f * 1024.0f), kv_size, (int) layers.size(), n_seq_max, n_stream, + ggml_type_name(type_k), (float)memory_size_k / (1024.0f * 1024.0f), + ggml_type_name(type_v), (float)memory_size_v / (1024.0f * 1024.0f)); } // TODO: refactor [TAG_KV_CACHE_SHARE_CELLS] @@ -419,39 +392,6 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { p1 = std::numeric_limits::max(); } - // empty range - nothing to remove - if (p0 >= p1) { - return true; - } - - // MSA anchors block selection to absolute cache slots (slot == position). Tail trim and full removal preserve this invariant, but removing a prefix - // or middle range would free slots while later cells survive, desynchronizing the indexer cache. Reject such removals before modifying the cache. - if (msa_strict_slots) { - for (llama_seq_id sid = 0; sid < (llama_seq_id) seq_to_stream.size(); ++sid) { - if (seq_id >= 0 && sid != seq_id) { - continue; - } - - const auto & cells = v_cells[seq_to_stream[sid]]; - - const llama_pos pmin = cells.seq_pos_min(sid); - const llama_pos pmax = cells.seq_pos_max(sid); - - if (pmin < 0) { - continue; // empty sequence - } - - const bool overlaps = p0 <= pmax && p1 > pmin; // the range removes something - const bool leaves_tail = p1 <= pmax; // cells beyond the range survive - - if (overlaps && leaves_tail) { - LLAMA_LOG_WARN("%s: MSA: partial (non-suffix) removal [%d, %d) for seq %d is not supported " - "(block selection is anchored to cache slots) - rejected\n", __func__, p0, p1, sid); - return false; - } - } - } - if (seq_id >= 0) { auto & cells = v_cells[seq_to_stream[seq_id]]; auto & head = v_heads[seq_to_stream[seq_id]]; @@ -906,10 +846,6 @@ bool llama_kv_cache::update(llama_context * lctx, bool do_shift, const stream_co if (layer.v_stream[ssrc]) { ggml_backend_tensor_copy(layer.v_stream[ssrc], layer.v_stream[sdst]); } - if (layer.k_idx_stream[ssrc]) { - GGML_ASSERT(layer.k_idx_stream[sdst]); - ggml_backend_tensor_copy(layer.k_idx_stream[ssrc], layer.k_idx_stream[sdst]); - } } } } @@ -1058,44 +994,6 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch, const auto & cells = v_cells[seq_to_stream[seq_id]]; - if (n_tokens > cells.size()) { - LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size()); - return { }; - } - - // MSA block selection assumes slot == logical position (append-only streams). - if (msa_strict_slots) { - for (uint32_t ii = 0; ii < n_tokens; ++ii) { - const llama_pos pos = ubatch.pos[s*n_tokens + ii]; - - if (pos < 0 || (uint64_t) pos >= cells.size()) { - LLAMA_LOG_WARN("%s: MSA: position %d is outside the cache range [0, %u)\n", - __func__, pos, cells.size()); - return { }; - } - - const uint32_t idx = (uint32_t) pos; - - if (!cells.is_empty(idx)) { - LLAMA_LOG_WARN("%s: MSA: required slot %u is already occupied (stream %u)\n", - __func__, idx, seq_to_stream[seq_id]); - return { }; - } - - // strictly increasing positions, rules out duplicates and, for contiguous requests, is tightened to exact adjacency - if (!res.idxs[s].empty() && (cont ? idx != res.idxs[s].back() + 1 - : idx <= res.idxs[s].back())) { - LLAMA_LOG_WARN("%s: MSA: token positions are not %s within the ubatch\n", - __func__, cont ? "contiguous" : "strictly increasing"); - return { }; - } - - res.idxs[s].push_back(idx); - } - - continue; - } - uint32_t head_cur = v_heads[seq_to_stream[seq_id]]; // if we have enough unused cells before the current head -> @@ -1104,6 +1002,11 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch, head_cur = 0; } + if (n_tokens > cells.size()) { + LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size()); + return { }; + } + uint32_t n_tested = 0; // for continuous slots, we test that all tokens in the ubatch fit, starting from the current head @@ -1210,15 +1113,6 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & const auto idx = sinfo.idxs[s][ii]; - if (msa_strict_slots && (llama_pos) idx != ubatch.pos[i]) { - LLAMA_LOG_ERROR("%s: MSA slot/position invariant violated: " - "writing pos %d into cell %u (stream %u). The indexer cache " - "would desync and block selection would silently corrupt. " - "This is a bug, please report it with reproduction steps.\n", - __func__, ubatch.pos[i], idx, sinfo.strm[s]); - GGML_ABORT("MSA: slot != pos"); - } - if (!cells.is_empty(idx)) { assert(cells.seq_count(idx) == 1); @@ -1262,8 +1156,7 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & LLAMA_LOG_DEBUG("%s: purging positions [%d, %d] of sequence %d from KV cache\n", __func__, cells.seq_pos_min(s), seq_pos_max_rm[s], s); - // under MSA strict slots this path should be unreachable, since strict MSA placement never selects occupied cells - GGML_ASSERT(seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1)); + seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1); } } @@ -1283,12 +1176,6 @@ bool llama_kv_cache::get_can_shift() const { if (hparams.n_pos_per_embd() > 1) { return false; } - // shifting would leave k_idx stale - for (const auto & layer : layers) { - if (layer.k_idx) { - return false; - } - } return true; } @@ -1337,6 +1224,12 @@ ggml_tensor * llama_kv_cache::get_k_storage(int32_t il) const { return layers[ikv].k; } +const llama_kv_cells & llama_kv_cache::get_cells(llama_seq_id seq_id) const { + GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < seq_to_stream.size()); + + return v_cells[seq_to_stream[seq_id]]; +} + uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const { uint32_t result = 0; @@ -1405,23 +1298,6 @@ ggml_tensor * llama_kv_cache::get_v(ggml_context * ctx, int32_t il, uint32_t n_k ggml_row_size(v->type, kv_size*n_embd_v_gqa)*sinfo.s0); } -ggml_tensor * llama_kv_cache::get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const { - const int32_t ikv = map_layer_ids.at(il); - auto * k_idx = layers[ikv].k_idx; - GGML_ASSERT(k_idx); - - const uint64_t kv_size = get_size(); - const int64_t n_idx = k_idx->ne[0]; // 128 - const uint32_t ns = sinfo.s1 - sinfo.s0 + 1; - - return ggml_view_4d(ctx, k_idx, - n_idx, 1, n_kv, ns, - ggml_row_size(k_idx->type, n_idx), // nb1 (single head) - ggml_row_size(k_idx->type, n_idx), // nb2 (per cell) - ggml_row_size(k_idx->type, n_idx*kv_size), // nb3 (per stream) - ggml_row_size(k_idx->type, n_idx*kv_size)*sinfo.s0); -} - ggml_tensor * llama_kv_cache::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const { GGML_UNUSED(sinfo); @@ -1523,28 +1399,6 @@ ggml_tensor * llama_kv_cache::build_input_k_idxs(ggml_context * ctx, const llama return k_idxs; } -ggml_tensor * llama_kv_cache::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const { - GGML_UNUSED(sinfo); - const int32_t ikv = map_layer_ids.at(il); - ggml_tensor * k_idx = layers[ikv].k_idx; - GGML_ASSERT(k_idx && "cpy_k_idx on a layer with no indexer cache"); - - const int64_t n_embd_head = k_idx_cur->ne[0]; // 128 - const int64_t n_head = k_idx_cur->ne[1]; // 1 - const int64_t n_tokens = k_idx_cur->ne[2]; - const int64_t n_embd_gqa = n_embd_head*n_head; // 128 - - GGML_ASSERT(ggml_row_size(k_idx_cur->type, n_embd_head) == k_idx_cur->nb[1]); - k_idx_cur = ggml_view_2d(ctx, k_idx_cur, n_embd_gqa, n_tokens, k_idx_cur->nb[2], 0); - - const int64_t n_stream = k_idx->ne[2]; - if (n_stream > 1) { - const int64_t kv_size = get_size(); - k_idx = ggml_reshape_2d(ctx, k_idx, n_embd_gqa, kv_size*n_stream); - } - return ggml_set_rows(ctx, k_idx, k_idx_cur, k_idxs); // same k_idxs as the K store -} - ggml_tensor * llama_kv_cache::build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const { const uint32_t n_tokens = ubatch.n_tokens; @@ -1979,18 +1833,6 @@ size_t llama_kv_cache::size_v_bytes() const { return size_v_bytes; } -size_t llama_kv_cache::size_k_idx_bytes() const { - size_t size_k_idx_bytes = 0; - - for (const auto & layer : layers) { - if (layer.k_idx) { - size_k_idx_bytes += ggml_nbytes(layer.k_idx); - } - } - - return size_k_idx_bytes; -} - ggml_tensor * llama_kv_cache::build_rope_shift( const llama_cparams & cparams, ggml_context * ctx, @@ -2303,36 +2145,6 @@ void llama_kv_cache::state_write_data(llama_io_write_i & io, const cell_ranges_t } } - if (size_k_idx_bytes() > 0) { - const uint32_t has_k_idx_u32 = 1; - io.write(&has_k_idx_u32, sizeof(has_k_idx_u32)); - - for (const auto & layer : layers) { - const uint32_t layer_has_k_idx = layer.k_idx ? 1 : 0; - io.write(&layer_has_k_idx, sizeof(layer_has_k_idx)); - - if (!layer_has_k_idx) { - continue; - } - - GGML_ASSERT(layer.k_idx_stream[cr.strm]); - - const int32_t k_idx_type_i = (int32_t) layer.k_idx->type; - io.write(&k_idx_type_i, sizeof(k_idx_type_i)); - - const uint64_t k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]); - io.write(&k_idx_size_row, sizeof(k_idx_size_row)); - - for (const auto & range : cr.data) { - const size_t range_size = range.second - range.first; - const size_t buf_size = range_size * k_idx_size_row; - const size_t offset = range.first * k_idx_size_row; - - io.write_tensor(layer.k_idx_stream[cr.strm], offset, buf_size); - } - } - } - if (!v_trans) { for (const auto & layer : layers) { const uint32_t il = layer.il; @@ -2581,68 +2393,6 @@ bool llama_kv_cache::state_read_data(llama_io_read_i & io, uint32_t strm, uint32 } } - if (size_k_idx_bytes() > 0) { - uint32_t has_k_idx_u32 = 0; - io.read(&has_k_idx_u32, sizeof(has_k_idx_u32)); - - if (has_k_idx_u32 != 1) { - LLAMA_LOG_ERROR("%s: missing k_idx data in KV cache state\n", __func__); - return false; - } - - for (const auto & layer : layers) { - uint32_t layer_has_k_idx = 0; - io.read(&layer_has_k_idx, sizeof(layer_has_k_idx)); - - const uint32_t expected_layer_has_k_idx = layer.k_idx ? 1 : 0; - - if (layer_has_k_idx != expected_layer_has_k_idx) { - LLAMA_LOG_ERROR( - "%s: mismatched k_idx state for layer: got %u, expected %u\n", - __func__, layer_has_k_idx, expected_layer_has_k_idx); - return false; - } - - if (!layer_has_k_idx) { - continue; - } - - GGML_ASSERT(layer.k_idx_stream[strm]); - - int32_t k_idx_type_i = -1; - io.read(&k_idx_type_i, sizeof(k_idx_type_i)); - - if (k_idx_type_i != (int32_t) layer.k_idx->type) { - LLAMA_LOG_ERROR( - "%s: mismatched k_idx type: got %d, expected %d\n", - __func__, k_idx_type_i, (int32_t) layer.k_idx->type); - return false; - } - - uint64_t k_idx_size_row = 0; - io.read(&k_idx_size_row, sizeof(k_idx_size_row)); - - const uint64_t expected_k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]); - - if (k_idx_size_row != expected_k_idx_size_row) { - LLAMA_LOG_ERROR( - "%s: mismatched k_idx row size: got %zu, expected %zu\n", - __func__, (size_t) k_idx_size_row, (size_t) expected_k_idx_size_row); - return false; - } - - if (cell_count) { - if (sinfo.is_contiguous()) { - io.read_tensor(layer.k_idx_stream[strm], sinfo.head() * k_idx_size_row, cell_count * k_idx_size_row); - } else { - for (uint32_t i = 0; i < cell_count; ++i) { - io.read_tensor(layer.k_idx_stream[strm], sinfo.idxs[0][i] * k_idx_size_row, k_idx_size_row); - } - } - } - } - } - if (!this->v_trans) { for (const auto & layer : layers) { const uint32_t il = layer.il; @@ -2844,10 +2594,6 @@ ggml_tensor * llama_kv_cache_context::get_v(ggml_context * ctx, int32_t il) cons return kv->get_v(ctx, il, n_kv, sinfos[i_cur]); } -ggml_tensor * llama_kv_cache_context::get_k_idx(ggml_context * ctx, int32_t il) const { - return kv->get_k_idx(ctx, il, n_kv, sinfos[i_cur]); -} - ggml_tensor * llama_kv_cache_context::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const { return kv->cpy_k(ctx, k_cur, k_idxs, il, sinfos[i_cur]); } @@ -2856,10 +2602,6 @@ ggml_tensor * llama_kv_cache_context::cpy_v(ggml_context * ctx, ggml_tensor * v_ return kv->cpy_v(ctx, v_cur, v_idxs, il, sinfos[i_cur]); } -ggml_tensor * llama_kv_cache_context::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const { - return kv->cpy_k_idx(ctx, k_idx_cur, k_idxs, il, sinfos[i_cur]); -} - ggml_tensor * llama_kv_cache_context::build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const { return kv->build_input_k_idxs(ctx, ubatch); } diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index d5a92f4405..6cb6dbd2f9 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -164,6 +164,8 @@ public: std::vector get_layer_ids() const; ggml_tensor * get_k_storage(int32_t il) const; + const llama_kv_cells & get_cells(llama_seq_id seq_id) const; + // // graph_build API // @@ -173,12 +175,10 @@ public: // get views of the current state of the cache ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; - ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; // store k_cur and v_cur in the cache based on the provided head location ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const; ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const; - ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const; // // preparation API @@ -230,11 +230,9 @@ private: ggml_tensor * k; ggml_tensor * v; - ggml_tensor * k_idx; // MSA single-head indexer keys, F32 std::vector k_stream; std::vector v_stream; - std::vector k_idx_stream; }; bool v_trans = true; // the value tensor is transposed @@ -263,9 +261,6 @@ private: // env: LLAMA_KV_CACHE_DEBUG int debug = 0; - // set when a k_idx (indexer) cache exists and the stream layout supports MSA (single seq, or one stream per seq) - bool msa_strict_slots = false; - // this is the SWA type of the cache - not to be confused with the model SWA type const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE; @@ -298,7 +293,6 @@ private: size_t size_k_bytes() const; size_t size_v_bytes() const; - size_t size_k_idx_bytes() const; ggml_tensor * build_rope_shift( const llama_cparams & cparams, @@ -378,7 +372,6 @@ public: // get views of the current state of the cache ggml_tensor * get_k(ggml_context * ctx, int32_t il) const; ggml_tensor * get_v(ggml_context * ctx, int32_t il) const; - ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il) const; // store k_cur and v_cur in the cache based on the provided head location // note: the heads in k_cur and v_cur should be laid out contiguously in memory @@ -388,7 +381,6 @@ public: // - v_idxs [n_tokens] or [n_tokens*n_embd_v_gqa] depending if V cache is transposed ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const; ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il) const; - ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const; // create destination indices for each head of the current batch for where it would be written in the KV cache // the indices address the global KV cache (not per stream) - this is not relevant for the user of this API, but diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 938d98798c..8fff1a4326 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -11,6 +11,7 @@ #include "llama-kv-cache.h" #include "llama-kv-cache-iswa.h" #include "llama-kv-cache-dsa.h" +#include "llama-kv-cache-msa.h" #include "llama-kv-cache-dsv4.h" #include "llama-memory-hybrid.h" #include "llama-memory-hybrid-iswa.h" @@ -2071,6 +2072,28 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, { res = nullptr; } break; + case LLM_ARCH_MINIMAX_M3: + { + // sparse (MSA) layers carry an indexer key cache, but leading dense layers do not + llama_kv_cache::layer_filter_cb filter_idx = + [&](int32_t il) { return (uint32_t) il >= hparams.n_layer_dense_lead; }; + + res = new llama_kv_cache_msa( + *this, + params.type_k, + params.type_v, + !cparams.flash_attn, + cparams.offload_kqv, + cparams.kv_unified, + cparams.n_ctx_seq, + cparams.n_seq_max, + 1, + hparams.n_swa, + hparams.swa_type, + nullptr, + filter_idx, + nullptr); + } break; case LLM_ARCH_GLM_DSA: case LLM_ARCH_DEEPSEEK32: { diff --git a/src/models/minimax-m3.cpp b/src/models/minimax-m3.cpp index 0773ad5435..8bd6a4298e 100644 --- a/src/models/minimax-m3.cpp +++ b/src/models/minimax-m3.cpp @@ -1,5 +1,5 @@ #include "models.h" -#include "llama-kv-cache.h" +#include "llama-kv-cache-msa.h" #include #include #include @@ -7,7 +7,8 @@ // MiniMax-M3: MiniMax-M2 style GQA (per-head QK-norm, partial rotary) with // DeepSeek-V3 leading-dense + routed/shared experts (sigmoid gating, routed scaling), // swigluoai activation, and MiniMax Sparse Attention (MSA). MTP is not in released model weights. -// Notes: Blocks are anchored to absolute KV cache slots. +// MSA blocks are defined over token positions. The graph translates between position space (block +// selection) and cell space (K/V/indexer storage) via per-ubatch pos<->cell maps populated from llama_kv_cells void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); @@ -23,7 +24,6 @@ void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, hparams.indexer_block_size); ml.get_key(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks); msa_p = { (int) hparams.indexer_block_size, (int) hparams.indexer_top_k, (int) hparams.indexer_local_blocks }; - hparams.indexer_kv = true; switch (hparams.n_layer()) { case 60: type = LLM_TYPE_428B_A23B; break; @@ -86,43 +86,83 @@ std::unique_ptr llama_model_minimax_m3::build_arch_graph(cons return std::make_unique(*this, params); } -// per-query local-force bias for MSA selection -// local window always wins a slot -class llm_graph_input_msa_local : public llm_graph_input_i { +class llm_graph_input_msa : public llm_graph_input_i { public: - llm_graph_input_msa_local(int blk, int local, int64_t nblk) : blk(blk), local(local), nblk(nblk) {} + llm_graph_input_msa(const llama_kv_cache_msa_context * mctx, int blk, int local) : + mctx(mctx), blk(blk), local(local) {} void set_input(const llama_ubatch * ubatch) override { - if (!bias || !ubatch->pos) { - return; - } - const int64_t n_tokens = ubatch->n_tokens; - std::vector data((size_t) nblk * n_tokens, 0.0f); - for (int64_t i = 0; i < n_tokens; ++i) { - const int64_t L = ubatch->pos[i] / blk; - for (int l = 0; l < local && L - l >= 0; ++l) { - if (L - l < nblk) { - data[(size_t) i * nblk + (L - l)] = 1e30f; + if (pos_slot_i) { mctx->set_input_pos_slot(pos_slot_i, ubatch); } + if (pos_slot_f) { mctx->set_input_pos_slot(pos_slot_f, ubatch); } + if (cell_blk) { mctx->set_input_cell_pos(cell_blk, ubatch, blk); } + if (pos_mask) { mctx->set_input_pos_mask(pos_mask, ubatch); } + + // local-force bias over position blocks + if (bias && ubatch->pos) { + const int64_t n_tokens = ubatch->n_tokens; + const int64_t nblk = bias->ne[0]; + std::vector data((size_t) nblk * n_tokens, 0.0f); + for (int64_t i = 0; i < n_tokens; ++i) { + const int64_t L = ubatch->pos[i] / blk; + for (int l = 0; l < local && L - l >= 0; ++l) { + if (L - l < nblk) { + data[(size_t) i * nblk + (L - l)] = 1e30f; + } } } + ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float)); } - ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float)); } - // valid as long as the bias tensor dims still match the new ubatch/cache window + // valid as long as the tensor dims still match the new ubatch/cache window and the + // ubatch is in the same regime (decode graphs have pos_slot_f, batch graphs cell_blk) bool can_reuse(const llm_graph_params & params) override { - const auto * mctx = static_cast(params.mctx); + const auto * mctx_new = static_cast(params.mctx); + + this->mctx = mctx_new; + + const int64_t n_ps = GGML_PAD((int64_t) mctx_new->get_n_pos(), blk); + const int64_t ns = params.cparams.kv_unified ? 1 : params.ubatch.n_seqs_unq; + + const bool decode = params.ubatch.n_tokens == ns; // one token per stream bool res = true; - res &= bias->ne[1] == params.ubatch.n_tokens; - res &= bias->ne[0] * blk == (int64_t) mctx->get_n_kv(); + + res &= bias->ne[0] * blk == n_ps; + res &= bias->ne[1] == params.ubatch.n_tokens; + + res &= pos_mask->ne[0] == n_ps; + res &= pos_mask->ne[1] == params.ubatch.n_tokens; + + res &= pos_slot_i->ne[0] == n_ps; + res &= pos_slot_i->ne[1] == ns; + + res &= decode == (pos_slot_f != nullptr); + res &= decode == (cell_blk == nullptr); + + if (pos_slot_f) { + res &= pos_slot_f->ne[0] == n_ps; + res &= pos_slot_f->ne[1] == ns; + } + + if (cell_blk) { + res &= cell_blk->ne[0] == (int64_t) mctx_new->get_base()->get_n_kv(); + res &= cell_blk->ne[1] == ns; + } + return res; } - ggml_tensor * bias = nullptr; - int blk; - int local; - int64_t nblk; + ggml_tensor * bias = nullptr; // F32 [nblk, n_tokens] local-force bias (position blocks) + ggml_tensor * pos_mask = nullptr; // F32 [n_ps, n_tokens] 0/-inf visibility, by position + ggml_tensor * pos_slot_i = nullptr; // I32 [n_ps, ns] pos -> cell (get_rows index) + ggml_tensor * pos_slot_f = nullptr; // F32 [n_ps, ns] pos -> cell (gatherable values, decode) + ggml_tensor * cell_blk = nullptr; // I32 [n_kv, ns] cell -> position block (batch) + + const llama_kv_cache_msa_context * mctx; + + int blk; + int local; }; // One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3]) @@ -173,7 +213,7 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ inpL = build_inp_embd(model.tok_embd); ggml_tensor * inp_pos = build_inp_pos(); - auto inp_attn = build_attn_inp_kv(); + auto inp_attn = build_attn_inp_kv_msa(); // MSA calls ggml_flash_attn_ext directly and assumes the non-transposed V layout that // llama.cpp only provides when flash attention is enabled. Block selection is anchored @@ -199,34 +239,51 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ } // hoisted per-graph MSA state (shared by every sparse layer) - llm_graph_input_msa_local * msa_loc = nullptr; + llm_graph_input_msa * msa = nullptr; ggml_tensor * msa_kqm = nullptr; - ggml_tensor * msa_mf = nullptr; - int64_t n_kv = 0, nblk = 0, ns = 1, n_tps = 0; + ggml_tensor * msa_mf = nullptr; // F32 copy of the FA mask for the final mask add + int64_t n_kv = 0, n_ps = 0, nblk = 0, ns = 1, n_tps = 0; bool msa_decode = false; // gather (1 token per stream) vs mask const int blk = mm.msa_p.blk; const int64_t Hd = hparams.indexer_n_head; // one indexer head per GQA group if (msa_enabled) { + const auto * mctx_msa = static_cast(mctx); + msa_kqm = inp_attn->get_kq_mask(); n_kv = msa_kqm->ne[0]; n_tps = msa_kqm->ne[1]; // tokens per stream ns = msa_kqm->ne[3]; // streams in this ubatch GGML_ASSERT(msa_kqm->type == GGML_TYPE_F16 && "MSA requires the FA (f16) mask"); GGML_ASSERT(n_tps*ns == n_tokens); - GGML_ASSERT(n_kv % blk == 0 && - "MSA: KV/mask n_kv must be a multiple of indexer.block_size (128); " - "the flash-attention KV padding must be a multiple of the block size. " - "A non-multiple would silently drop the partial tail block."); - nblk = n_kv / blk; + + // the position axis covers every position currently in the cache and is padded to whole blocks + n_ps = GGML_PAD((int64_t) mctx_msa->get_n_pos(), blk); + nblk = n_ps / blk; msa_decode = n_tps == 1; - msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32); + auto inp = std::make_unique(mctx_msa, blk, mm.msa_p.local); - auto loc = std::make_unique(blk, mm.msa_p.local, nblk); - loc->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens - ggml_set_input(loc->bias); - msa_loc = (llm_graph_input_msa_local *) res->add_input(std::move(loc)); + inp->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens + ggml_set_input(inp->bias); + + inp->pos_mask = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, n_tokens); + ggml_set_input(inp->pos_mask); + + inp->pos_slot_i = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_ps, ns); + ggml_set_input(inp->pos_slot_i); + + if (msa_decode) { + inp->pos_slot_f = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, ns); + ggml_set_input(inp->pos_slot_f); + } else { + inp->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, ns); + ggml_set_input(inp->cell_blk); + + msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32); + } + + msa = (llm_graph_input_msa *) res->add_input(std::move(inp)); } ggml_tensor * inp_out_ids = build_inp_out_ids(); @@ -283,9 +340,11 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ ik = ggml_rope_ext(ctx0, ik, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow); - const auto * mctx_cur = inp_attn->mctx; - ggml_build_forward_expand(gf, mctx_cur->cpy_k_idx(ctx0, ik, inp_attn->get_k_idxs(), il)); - ggml_tensor * ik_kv = mctx_cur->get_k_idx(ctx0, il); + const auto * mctx_msa_l = static_cast(mctx); + const auto * mctx_cur = mctx_msa_l->get_base(); + const auto * mctx_idx = mctx_msa_l->get_idx(); + ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, ik, inp_attn->get_k_idxs_idx(), il)); + ggml_tensor * ik_kv = mctx_idx->get_k(ctx0, il); if (inp_attn->self_k_rot) { Qcur = llama_mul_mat_hadamard(ctx0, Qcur, inp_attn->self_k_rot); @@ -316,42 +375,52 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ if (msa_decode) { // decode: batched over streams top-k + gather, one grouped FA - // scores: per-stream batched matmul over the stream dim (ne[3]). - // the cache views are not contiguous across streams (stride = kv_size, not n_kv) - ggml_tensor * ikv4 = ggml_view_4d(ctx0, ik_kv, n_idx_dim, n_kv, 1, ns, - ik_kv->nb[2], ik_kv->nb[3], ik_kv->nb[3], 0); + // gather the indexer keys through the pos -> cell map + ggml_tensor * ik3 = ggml_view_3d(ctx0, ik_kv, n_idx_dim, n_kv, ns, + ik_kv->nb[2], ik_kv->nb[3], 0); + ggml_tensor * ikp = ggml_get_rows(ctx0, ik3, msa->pos_slot_i); // [n_idx_dim, n_ps, ns] ggml_tensor * iq4 = ggml_reshape_4d(ctx0, iq, n_idx_dim, Hd, 1, ns); - ggml_tensor * sc = ggml_mul_mat(ctx0, ikv4, iq4); + ggml_tensor * sc = ggml_mul_mat(ctx0, + ggml_reshape_4d(ctx0, ikp, n_idx_dim, n_ps, 1, ns), iq4); ggml_mul_mat_set_prec(sc, GGML_PREC_F32); - sc = ggml_add_inplace(ctx0, sc, msa_mf); + // unmapped positions come out -inf, so they can never rank into the top-k + sc = ggml_add_inplace(ctx0, sc, + ggml_reshape_4d(ctx0, msa->pos_mask, n_ps, 1, 1, ns)); ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0); cb(bs, "msa_bs", il); ggml_tensor * bsf = ggml_add(ctx0, bs, - ggml_reshape_4d(ctx0, msa_loc->bias, nblk, 1, 1, ns)); - ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); + ggml_reshape_4d(ctx0, msa->bias, nblk, 1, 1, ns)); + ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); // position blocks - // token idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (for the mask gather) - // row idx: tr[t,k,h,s] = tj*HKV + h (for the per-stream K/V gather) + // pos idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (positions - mask gather) + // cell idx: cs[t,k,h,s] = pos_slot[tj] (pos -> cell translation) + // row idx: tr[t,k,h,s] = cs*HKV + h (per-stream K/V gather) ggml_tensor * a = ggml_scale(ctx0, ggml_cast(ctx0, idx, GGML_TYPE_F32), (float) blk); a = ggml_reshape_4d(ctx0, a, 1, K, Hd, ns); ggml_tensor * tj = ggml_add(ctx0, ggml_repeat_4d(ctx0, a, blk, K, Hd, ns), ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) blk, 1.0f), blk, 1, 1)); - ggml_tensor * tr = ggml_add(ctx0, - ggml_scale(ctx0, tj, (float) HKV), - ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd)); ggml_tensor * tokj = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tj, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32); + + ggml_tensor * cs = ggml_get_rows(ctx0, + ggml_reshape_3d(ctx0, msa->pos_slot_f, 1, n_ps, ns), tokj); // [1, blk*K*Hd, ns] + cs = ggml_reshape_4d(ctx0, cs, blk, K, Hd, ns); + + ggml_tensor * tr = ggml_add(ctx0, + ggml_scale(ctx0, cs, (float) HKV), + ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd)); + ggml_tensor * tokr = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tr, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32); ggml_tensor * k3 = ggml_view_3d(ctx0, k, D, HKV*n_kv, ns, k->nb[1], k->nb[3], 0); ggml_tensor * v3 = ggml_view_3d(ctx0, v, D, HKV*n_kv, ns, v->nb[1], v->nb[3], 0); - ggml_tensor * m3 = ggml_reshape_3d(ctx0, msa_kqm, 1, n_kv, ns); + ggml_tensor * mp = ggml_reshape_3d(ctx0, msa->pos_mask, 1, n_ps, ns); ggml_tensor * kg = ggml_get_rows(ctx0, k3, tokr); ggml_tensor * vg = ggml_get_rows(ctx0, v3, tokr); - ggml_tensor * mg = ggml_get_rows(ctx0, m3, tokj); + ggml_tensor * mg = ggml_get_rows(ctx0, mp, tokj); // fold (group, stream) onto the FA channel dim const ggml_type kt = ggml_is_quantized(k->type) ? GGML_TYPE_F16 : k->type; @@ -372,12 +441,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ iq->nb[1], iq->nb[2], st*n_tps*iq->nb[2]); ggml_tensor * ik_s = ggml_view_2d(ctx0, ik_kv, n_idx_dim, n_kv, ik_kv->nb[2], st*ik_kv->nb[3]); - ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, 1, n_tps, - msa_mf->nb[1], msa_mf->nb[1], st*msa_mf->nb[3]); - ggml_tensor * km_s = ggml_view_3d(ctx0, msa_kqm, n_kv, n_tps, 1, - msa_kqm->nb[1], msa_kqm->nb[3], st*msa_kqm->nb[3]); - ggml_tensor * bias_s = ggml_view_3d(ctx0, msa_loc->bias, nblk, 1, n_tps, - msa_loc->bias->nb[1], msa_loc->bias->nb[1], st*n_tps*msa_loc->bias->nb[1]); + ggml_tensor * psl_s = ggml_view_1d(ctx0, msa->pos_slot_i, n_ps, + st*msa->pos_slot_i->nb[1]); + ggml_tensor * pm_s = ggml_view_3d(ctx0, msa->pos_mask, n_ps, 1, n_tps, + msa->pos_mask->nb[1], msa->pos_mask->nb[1], st*n_tps*msa->pos_mask->nb[1]); + ggml_tensor * cb_s = ggml_view_1d(ctx0, msa->cell_blk, n_kv, + st*msa->cell_blk->nb[1]); + ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, n_tps, 1, + msa_mf->nb[1], msa_mf->nb[3], st*msa_mf->nb[3]); + ggml_tensor * bias_s = ggml_view_3d(ctx0, msa->bias, nblk, 1, n_tps, + msa->bias->nb[1], msa->bias->nb[1], st*n_tps*msa->bias->nb[1]); ggml_tensor * q_s = ggml_view_3d(ctx0, Qcur, D, n_head, n_tps, Qcur->nb[1], Qcur->nb[2], st*n_tps*Qcur->nb[2]); ggml_tensor * k_s = ggml_view_4d(ctx0, k, D, HKV, n_kv, 1, @@ -385,14 +458,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ ggml_tensor * v_s = ggml_view_4d(ctx0, v, D, HKV, n_kv, 1, v->nb[1], v->nb[2], v->nb[3], st*v->nb[3]); - // block scores: bs = maxpool_blk(idx_q * idx_k^T + causal mask) + // block scores: the indexer keys are gathered through the pos -> cell map first // scores are unscaled, only the top-k ordering matters - ggml_tensor * sc = ggml_mul_mat(ctx0, ik_s, + ggml_tensor * ikp = ggml_get_rows(ctx0, ik_s, psl_s); // [n_idx_dim, n_ps] + ggml_tensor * sc = ggml_mul_mat(ctx0, ikp, ggml_reshape_2d(ctx0, iq_s, n_idx_dim, Hd*n_tps)); // indexer scores run in F32 ggml_mul_mat_set_prec(sc, GGML_PREC_F32); - sc = ggml_reshape_3d(ctx0, sc, n_kv, Hd, n_tps); - sc = ggml_add_inplace(ctx0, sc, mf_s); + sc = ggml_reshape_3d(ctx0, sc, n_ps, Hd, n_tps); + // unmapped positions (holes, padding, empty cells) come out -inf + sc = ggml_add_inplace(ctx0, sc, pm_s); ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0); cb(bs, "msa_bs", il); @@ -416,14 +491,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ bm = ggml_cont(ctx0, ggml_permute(ctx0, bm, 0, 2, 1, 3)); // [nblk, n_tps, Hd] cb(bm, "msa_block_mask", il); - // expand block -> token granularity (j = bk*blk + t), - // then combine with the causal mask in place - ggml_tensor * bmx = ggml_repeat_4d(ctx0, - ggml_reshape_3d(ctx0, bm, 1, nblk, n_tps*Hd), - blk, nblk, n_tps*Hd, 1); + // expand block -> cell granularity through the cell -> position block + // map, then combine with the causal mask. empty cells are masked by the causal mask. + ggml_tensor * bm2 = ggml_cont(ctx0, ggml_transpose(ctx0, + ggml_reshape_2d(ctx0, bm, nblk, n_tps*Hd))); // [n_tps*Hd, nblk] + ggml_tensor * bmc = ggml_get_rows(ctx0, bm2, cb_s); // [n_tps*Hd, n_kv] F32 + ggml_tensor * bmx = ggml_cont(ctx0, ggml_transpose(ctx0, bmc)); bmx = ggml_reshape_3d(ctx0, bmx, n_kv, n_tps, Hd); - ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, km_s); - mask4 = ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd); + ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, mf_s); + mask4 = ggml_cast(ctx0, + ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd), GGML_TYPE_F16); cb(mask4, "msa_mask4", il); // cache views with groups on ne[3];