model: M3: Move MSA into a new memory implementation (#26338)

* Move MSA logic from llama-kv-cache into llama-kv-cache-msa

* cont : minor

* cont : ws fix

---------

Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
timkhronos
2026-08-03 15:30:08 +02:00
committed by GitHub
parent 563dec81c1
commit 67d5978bb1
11 changed files with 826 additions and 378 deletions
+57
View 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.