]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
model: M3: Move MSA into a new memory implementation (#26338)
authortimkhronos <redacted>
Mon, 3 Aug 2026 13:30:08 +0000 (15:30 +0200)
committerGitHub <redacted>
Mon, 3 Aug 2026 13:30:08 +0000 (16:30 +0300)
* Move MSA logic from llama-kv-cache into llama-kv-cache-msa

* cont : minor

* cont : ws fix

---------

Co-authored-by: Georgi Gerganov <redacted>
src/CMakeLists.txt
src/llama-graph.cpp
src/llama-graph.h
src/llama-hparams.cpp
src/llama-hparams.h
src/llama-kv-cache-msa.cpp [new file with mode: 0644]
src/llama-kv-cache-msa.h [new file with mode: 0644]
src/llama-kv-cache.cpp
src/llama-kv-cache.h
src/llama-model.cpp
src/models/minimax-m3.cpp

index 320784c3a8cc8f83e1d01be64b7c6208ccd97916..24f05cc91673217726b919229e1626b7f74a7bcb 100644 (file)
@@ -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
index 1a35692300c1d28f64b488d892252311fd7707fb..fdab7b8dde679b068b6b024e2c9049373bb92c92 100644 (file)
@@ -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<const llama_kv_cache_msa_context *>(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<const llama_kv_cache_msa_context *>(mctx);
+
+    auto inp = std::make_unique<llm_graph_input_attn_kv_msa>(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.
index 160e294135523839a535da36c6948798343bd881..ff216302dbb11f7b4aa404dab890e335e97ec0d5 100644 (file)
@@ -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,
index 50af97f358c339369b637587aed6923709c866c1..846d4c69a6265b1cb7663605befe757c6b2f75bd 100644 (file)
@@ -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
index fc770bf003e612c45236bd246d12629de09fd209..6e8336c987481de11f985f9ca89f5446f4ff077d 100644 (file)
@@ -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 (file)
index 0000000..55ef286
--- /dev/null
@@ -0,0 +1,395 @@
+#include "llama-kv-cache-msa.h"
+
+#include "llama-impl.h"
+#include "llama-batch.h"
+#include "llama-model.h"
+
+#include <algorithm>
+#include <cassert>
+#include <cmath>
+
+// 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<llama_kv_cache>(
+            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<llama_kv_cache>(
+            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<ggml_backend_buffer_type_t, size_t> llama_kv_cache_msa::memory_breakdown() const {
+    std::map<ggml_backend_buffer_type_t, size_t> 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<llama_ubatch> 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<llama_kv_cache_msa_context>(
+                this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches));
+    } while (false);
+
+    return std::make_unique<llama_kv_cache_msa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);
+}
+
+llama_memory_context_ptr llama_kv_cache_msa::init_full() {
+    return std::make_unique<llama_kv_cache_msa_context>(this);
+}
+
+llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) {
+    return std::make_unique<llama_kv_cache_msa_context>(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<llama_ubatch> 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<const llama_kv_cache_context *>(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<const llama_kv_cache_context *>(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<int32_t> 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 (file)
index 0000000..f09b6d3
--- /dev/null
@@ -0,0 +1,153 @@
+#pragma once
+
+#include "llama-kv-cache.h"
+
+#include <vector>
+
+// 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<ggml_backend_buffer_type_t, size_t> 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<llama_kv_cache> kv_base;
+    std::unique_ptr<llama_kv_cache> 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<llama_ubatch> 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<llama_ubatch> ubatches;
+
+    const llama_memory_context_ptr ctx_base;
+    const llama_memory_context_ptr ctx_idx;
+
+    const llama_memory_status status;
+};
index 44cb1668dacf74d899a1658a9a51ba825a281510..8678a326d9eec69e9a8d696568242df58c0a39bc 100644 (file)
@@ -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<ggml_tensor *> 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;
-
-        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);
-        }
+        const size_t memory_size_k = size_k_bytes();
+        const size_t memory_size_v = size_v_bytes();
 
-        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<llama_pos>::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);
 }
index d5a92f4405b572ebc5357516b8a1420d97bd5ca5..6cb6dbd2f9843aa95319660654b7bc3eb5293740 100644 (file)
@@ -164,6 +164,8 @@ public:
     std::vector<uint32_t> 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<ggml_tensor *> k_stream;
         std::vector<ggml_tensor *> v_stream;
-        std::vector<ggml_tensor *> 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
index 938d98798cdbb4e942b05b5b498e311e13f7e0c8..8fff1a432667dfcfccc7b66447d5f8df0662d68b 100644 (file)
@@ -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:
             {
index 0773ad5435c98f483a4eaa67001264a73b86ed98..8bd6a4298e0d7bfb2d8553fc27b6caa0c6a36214 100644 (file)
@@ -1,5 +1,5 @@
 #include "models.h"
-#include "llama-kv-cache.h"
+#include "llama-kv-cache-msa.h"
 #include <cmath>
 #include <vector>
 #include <cstdint>
@@ -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<llm_graph_context> llama_model_minimax_m3::build_arch_graph(cons
     return std::make_unique<graph>(*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<float> 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<float> 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<const llama_kv_cache_context *>(params.mctx);
+        const auto * mctx_new = static_cast<const llama_kv_cache_msa_context *>(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<const llama_kv_cache_msa_context *>(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<llm_graph_input_msa>(mctx_msa, blk, mm.msa_p.local);
+
+        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);
 
-        auto loc = std::make_unique<llm_graph_input_msa_local>(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));
+            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<const llama_kv_cache_msa_context *>(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 * 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, tj, (float) HKV),
+                            ggml_scale(ctx0, cs, (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 * 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];