]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
model: Add MiniMax-M3 (MSA: MiniMax Sparse Attention) support (#24908)
authortimkhronos <redacted>
Sun, 26 Jul 2026 17:43:45 +0000 (19:43 +0200)
committerGitHub <redacted>
Sun, 26 Jul 2026 17:43:45 +0000 (19:43 +0200)
* Add preliminary MiniMax-M3 support

Text-only port that re-uses existing components: MiniMax-M2 style GQA with
per-head QK-norm and partial rotary, DeepSeek-V3 style leading-dense and
routed/shared experts, and swigluoai activation. Sparse attention is not
yet supported (dense fallback); vision tower and MTP heads are dropped.

* MiniMax-M3 vision tower (mmproj + clip graph)

* Delete m3_vision_ref.py

* Update clip.cpp

* MSA

* Update constants.py

* Update minimax.py

* Cache creation. Working withotu flash attention

* Added flash attention for sparse layers

* Decomposed slow cpu OP into GPU + CPU ops. Massive speedup over long ctx

* Rewrote indexer op to be cuda native. Modified flash attention to match per group block picking

* Implement sparse attention calc out of stock ops.

* Fix a cache allocation and cont issue

* Fixed -fa auto crash, flagged debug spots

* Delete vocab.json

* Delete model.safetensors.index.json

* Delete generation_config.json

* Delete Minimax directory

* Handled multi stream case to fall back on Dense Attention

* Development scaffolding cleanup. No functional change to the decode or
4-way paths. Full debug harness remains at <8136a9c68ed7a5eb009aa67bba3fda8062f4648f> for reproducing the
selection-parity validation.

* Remove redundant comment from minimax-m3.cpp

* Changed 3 Gelu Ops for vision into Gelu_erf ops

* Assert that n_kv is multiple of 128

* Rename MSA index tensors to indexer convention

Note: All GGUFs generated before this change will need to be regenerated.

* Fix incorrect Assert

* Review driven changes (#3)

* Remove comment from conversion minimax.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Remove whitespaces from constants.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Tighten comment in minimax.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* inherit MiniMax-M3 from MiniMax-M2

* drop dead text_config fallbacks

* Add indexer writer methods

* Reuse LLM_FFN_SWIGLU_OAI_MOE

* Remove duplicate  indexer setters, add only block_size/local_blocks, follow value naming convention

* Fix conversion error /gguf_writer.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Update gguf-py/gguf/gguf_writer.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Update gguf-py/gguf/tensor_mapping.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Update conversion/minimax.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Update conversion/minimax.py

Co-authored-by: Sigbjørn Skjæret <redacted>
* Remove whitespace in src/llama-kv-cache.cpp

Co-authored-by: Sigbjørn Skjæret <redacted>
* Remove Whitespace in Update src/llama-model.h

Co-authored-by: Sigbjørn Skjæret <redacted>
* Remove whitespace in src/llama-hparams.h

Co-authored-by: Sigbjørn Skjæret <redacted>
* remove multimodal code upon maintainer request. Will be made as a separate PR

* Whitespace clean in tensor_mapping.py

* Log cache size on launch, block ctx shift, support prompt caching

Log indexer cache size on launch

Disallow ctx shift

Support prompt caching

* Update minimax-m3.cpp

* Optimize implementation, add multi stream support.

Fully rewrote minimax-m3.cpp for speed and buffer size gains:

Unified the 4-way + decode, 1 FA call per layer instead of 4, with the groups mapped onto ne[3]

Custom CPU op now emits block-level mask, expanded on GPU, which causes CPU to GPU transfer to shrinks at prefill

Decode: ~25 nodes/layer vs ~50, no per-group concats/conts

Unified selection semantics, so both regimes rank bs + local bias (position-anchored local force), which means prefill/decode can no longer disagree on selection

can_reuse on the MSA bias input. Graph reuse at decode restored (was rebuilding the full graph every token)

In-place mask adds, shrinking compute buffer ~6.8 to ~4.2 GiB at ub2048/62k

Multi-stream: MSA now runs with -np N when kv_unified=false. Decode stays batched across streams (still 1 FA call), prefill loops per stream. dense fallback only for --kv-unified + multi-seq

Measured effect on expert offload bound setup: decode 6.2(4WAY)–7.15(MSA_decode) -> 7.7~7.8 t/s, flat from 5k to 60k+. prefill around 10% faster. buffer about 20% smaller, multi-user support.

* set default cache type to F32

* Fix potential DSA double indexer cache  allocation bug, only allocate in-cache k_idx for archs that opt in

* remove F16 downcasts in MSA attention, force F32 indexer score accum

* Add Minimax eos to llama vocab

* Guard edge case where idx cache can become stale after a tail trim

* Update llama-kv-cache.h

* Update llama-kv-cache.cpp

* Update llama-kv-cache.cpp

* Update llama-kv-cache.h

* Update llama-kv-cache.cpp

* Review driven changes

* style fix

* indexer hparams are required

* fix tests

* fix lint

---------

Co-authored-by: Daniel Han <redacted>
Co-authored-by: Sigbjørn Skjæret <redacted>
Co-authored-by: Xuan Son Nguyen <redacted>
21 files changed:
conversion/__init__.py
conversion/base.py
conversion/minimax.py
gguf-py/gguf/constants.py
gguf-py/gguf/gguf_writer.py
gguf-py/gguf/tensor_mapping.py
src/llama-arch.cpp
src/llama-arch.h
src/llama-context.cpp
src/llama-graph.cpp
src/llama-hparams.cpp
src/llama-hparams.h
src/llama-kv-cache.cpp
src/llama-kv-cache.h
src/llama-model-saver.cpp
src/llama-model.cpp
src/llama-model.h
src/llama-vocab.cpp
src/models/minimax-m3.cpp [new file with mode: 0644]
src/models/models.h
tests/test-llama-archs.cpp

index 7936f1159cb862a86c38cd7d817b2228e284c915..0b08e6e57008ae85c116551864c7d3d868d4581e 100644 (file)
@@ -158,6 +158,8 @@ TEXT_MODEL_MAP: dict[str, str] = {
     "MiniCPMForCausalLM": "minicpm",
     "MiniCPMV4_6ForConditionalGeneration": "minicpm",
     "MiniMaxM2ForCausalLM": "minimax",
+    "MiniMaxM3SparseForCausalLM": "minimax",
+    "MiniMaxM3SparseForConditionalGeneration": "minimax",
     "Ministral3ForCausalLM": "mistral3",
     "Mistral3ForConditionalGeneration": "mistral3",
     "MistralForCausalLM": "llama",
index 051b8b4e59fdd6304a8b0213adecb9e763beb9e6..a7cd3fd904aaebc9a881d2fc20e209a2b82e5221 100644 (file)
@@ -1156,7 +1156,7 @@ class TextModel(ModelBase):
                 or "projector." in name or "pre_mm_projector_norm" in name \
                 or "image_newline" in name or "view_seperator" in name \
                 or "patch_embed" in name or "patch_embedding" in name \
-                or "patch_merger." in name or "model.connector." in name:
+                or "patch_merger." in name or "patch_merge_mlp." in name or "model.connector." in name:
             return None
 
         return super().filter_tensors(item)
@@ -1203,7 +1203,7 @@ class TextModel(ModelBase):
             self.gguf_writer.add_embedding_length(n_embd)
             logger.info(f"gguf: embedding length = {n_embd}")
 
-        if (n_ff := self.find_hparam(["prefix_dense_intermediate_size", "intermediate_size", "n_inner", "hidden_dim"], optional=True)) is not None:
+        if (n_ff := self.find_hparam(["prefix_dense_intermediate_size", "dense_intermediate_size", "intermediate_size", "n_inner", "hidden_dim"], optional=True)) is not None:
             self.gguf_writer.add_feed_forward_length(n_ff)
             logger.info(f"gguf: feed forward length = {n_ff}")
 
index 4857775cbfb9caf4a30fabf54b13ee807a4b513f..cbbdfe3ae82d2f539480a0ea6c886b62a6fb5149 100644 (file)
@@ -23,7 +23,7 @@ class MiniMaxM2Model(TextModel):
 
     def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None):
         # merge expert weights
-        if 'experts' in name:
+        if "block_sparse_moe.experts." in name:
             n_experts = self.find_hparam(["num_local_experts", "num_experts"])
             assert bid is not None
 
@@ -52,3 +52,38 @@ class MiniMaxM2Model(TextModel):
             return
 
         yield from super().modify_tensors(data_torch, name, bid)
+
+
+@ModelBase.register("MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration")
+class MiniMaxM3Model(MiniMaxM2Model):
+    model_arch = gguf.MODEL_ARCH.MINIMAXM3
+
+    def set_gguf_parameters(self):
+        super().set_gguf_parameters()
+
+        self.gguf_writer.add_expert_shared_count(self.find_hparam(["n_shared_experts"]))
+        self.gguf_writer.add_expert_weights_scale(self.find_hparam(["routed_scaling_factor"]))
+        self.gguf_writer.add_expert_weights_norm(True)
+
+        sac = self.find_hparam(["sparse_attention_config"])
+        self.gguf_writer.add_indexer_head_count(sac["sparse_num_index_heads"])
+        self.gguf_writer.add_indexer_key_length(sac["sparse_index_dim"])
+        self.gguf_writer.add_indexer_top_k(sac["sparse_topk_blocks"])
+        self.gguf_writer.add_indexer_block_size(sac["sparse_block_size"])
+        self.gguf_writer.add_indexer_local_blocks(sac["sparse_local_block"])
+
+        moe_layer_freq = self.find_hparam(["moe_layer_freq"])
+        n_dense = 0
+        for v in moe_layer_freq:
+            if v == 0:
+                n_dense += 1
+            else:
+                break
+        self.gguf_writer.add_leading_dense_block_count(n_dense)
+
+    def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None):
+        # Gemma-style (1 + w) RMSNorm: bake the +1 in so llama.cpp can use plain RMSNorm
+        if name.endswith("norm.weight"):
+            data_torch = data_torch + 1.0
+
+        yield from super().modify_tensors(data_torch, name, bid)
index d55253e0eb4bad24bfa4de158dfddee86fbfd652..66d50cca268401e3cb84740713be75b72640809e 100644 (file)
@@ -200,6 +200,8 @@ class Keys:
             HEAD_COUNT = "{arch}.attention.indexer.head_count"
             KEY_LENGTH = "{arch}.attention.indexer.key_length"
             TOP_K      = "{arch}.attention.indexer.top_k"
+            BLOCK_SIZE   = "{arch}.attention.indexer.block_size"    # MSA
+            LOCAL_BLOCKS = "{arch}.attention.indexer.local_blocks"  # MSA
             TYPES      = "{arch}.attention.indexer.types"
 
     class HyperConnection:
@@ -528,6 +530,7 @@ class MODEL_ARCH(IntEnum):
     APERTUS          = auto()
     COGVLM           = auto()
     MINIMAXM2        = auto()
+    MINIMAXM3        = auto()
     RND1             = auto()
     PANGU_EMBED      = auto()
     MISTRAL3         = auto()
@@ -774,6 +777,9 @@ class MODEL_TENSOR(IntEnum):
     INDEXER_PROJ         = auto()
     INDEXER_ATTN_K       = auto()
     INDEXER_ATTN_Q_B     = auto()
+    INDEXER_Q_PROJ       = auto()
+    INDEXER_K_PROJ       = auto()
+    INDEXER_Q_NORM       = auto()
     INDEXER_COMPRESSOR_WKV = auto()
     INDEXER_COMPRESSOR_WGATE = auto()
     INDEXER_COMPRESSOR_APE = auto()
@@ -1110,6 +1116,7 @@ MODEL_ARCH_NAMES: dict[MODEL_ARCH, str] = {
     MODEL_ARCH.GROVEMOE:         "grovemoe",
     MODEL_ARCH.APERTUS:          "apertus",
     MODEL_ARCH.MINIMAXM2:        "minimax-m2",
+    MODEL_ARCH.MINIMAXM3:        "minimax-m3",
     MODEL_ARCH.COGVLM:           "cogvlm",
     MODEL_ARCH.RND1:             "rnd1",
     MODEL_ARCH.PANGU_EMBED:      "pangu-embedded",
@@ -1355,6 +1362,9 @@ TENSOR_NAMES: dict[MODEL_TENSOR, str] = {
     MODEL_TENSOR.INDEXER_PROJ:              "blk.{bid}.indexer.proj",
     MODEL_TENSOR.INDEXER_ATTN_K:            "blk.{bid}.indexer.attn_k",
     MODEL_TENSOR.INDEXER_ATTN_Q_B:          "blk.{bid}.indexer.attn_q_b",
+    MODEL_TENSOR.INDEXER_Q_PROJ:            "blk.{bid}.indexer.q_proj",
+    MODEL_TENSOR.INDEXER_K_PROJ:            "blk.{bid}.indexer.k_proj",
+    MODEL_TENSOR.INDEXER_Q_NORM:            "blk.{bid}.indexer.q_norm",
     MODEL_TENSOR.INDEXER_COMPRESSOR_WKV:    "blk.{bid}.indexer_compressor_kv",
     MODEL_TENSOR.INDEXER_COMPRESSOR_WGATE:  "blk.{bid}.indexer_compressor_gate",
     MODEL_TENSOR.INDEXER_COMPRESSOR_APE:    "blk.{bid}.indexer_compressor_ape",
@@ -4163,6 +4173,34 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
         MODEL_TENSOR.FFN_UP_EXP,
         MODEL_TENSOR.FFN_EXP_PROBS_B,
     ],
+    MODEL_ARCH.MINIMAXM3: [
+        MODEL_TENSOR.TOKEN_EMBD,
+        MODEL_TENSOR.OUTPUT_NORM,
+        MODEL_TENSOR.OUTPUT,
+        MODEL_TENSOR.ATTN_NORM,
+        MODEL_TENSOR.ATTN_Q,
+        MODEL_TENSOR.ATTN_Q_NORM,
+        MODEL_TENSOR.ATTN_K,
+        MODEL_TENSOR.ATTN_K_NORM,
+        MODEL_TENSOR.ATTN_V,
+        MODEL_TENSOR.ATTN_OUT,
+        MODEL_TENSOR.FFN_NORM,
+        MODEL_TENSOR.FFN_GATE_INP,
+        MODEL_TENSOR.FFN_EXP_PROBS_B,
+        MODEL_TENSOR.FFN_GATE_EXP,
+        MODEL_TENSOR.FFN_DOWN_EXP,
+        MODEL_TENSOR.FFN_UP_EXP,
+        MODEL_TENSOR.FFN_GATE_SHEXP,
+        MODEL_TENSOR.FFN_DOWN_SHEXP,
+        MODEL_TENSOR.FFN_UP_SHEXP,
+        MODEL_TENSOR.FFN_GATE,
+        MODEL_TENSOR.FFN_DOWN,
+        MODEL_TENSOR.FFN_UP,
+        MODEL_TENSOR.INDEXER_Q_PROJ,
+        MODEL_TENSOR.INDEXER_K_PROJ,
+        MODEL_TENSOR.INDEXER_Q_NORM,
+        MODEL_TENSOR.INDEXER_K_NORM,
+    ],
     MODEL_ARCH.COGVLM: [
         MODEL_TENSOR.TOKEN_EMBD,
         MODEL_TENSOR.OUTPUT_NORM,
index bb21596701d489de5b1b0d518cd6f66093c80c27..ba08f8d650044ab57c55f6ac1bd78d6dbf798851 100644 (file)
@@ -793,6 +793,12 @@ class GGUFWriter:
     def add_indexer_top_k(self, top_k: int) -> None:
         self.add_uint32(Keys.Attention.Indexer.TOP_K.format(arch=self.arch), top_k)
 
+    def add_indexer_block_size(self, block_size: int) -> None:
+        self.add_uint32(Keys.Attention.Indexer.BLOCK_SIZE.format(arch=self.arch), block_size)
+
+    def add_indexer_local_blocks(self, local_blocks: int) -> None:
+        self.add_uint32(Keys.Attention.Indexer.LOCAL_BLOCKS.format(arch=self.arch), local_blocks)
+
     def add_indexer_types(self, value: Sequence[bool]) -> None:
         key = Keys.Attention.Indexer.TYPES.format(arch=self.arch)
         self.add_array(key, value)
index b5707f11f5c46dbe898a0d2b2ecd3ec9de937787..59623accfdb66ab2453d149262bc063cab494867 100644 (file)
@@ -1264,7 +1264,8 @@ class TensorNameMap:
         ),
 
         MODEL_TENSOR.INDEXER_K_NORM: (
-            "model.layers.{bid}.self_attn.indexer.k_norm", # DSA
+            "model.layers.{bid}.self_attn.indexer.k_norm",  # DSA
+            "model.layers.{bid}.self_attn.index_k_norm",    # MSA
         ),
 
         MODEL_TENSOR.INDEXER_PROJ: (
@@ -1279,6 +1280,18 @@ class TensorNameMap:
             "model.layers.{bid}.self_attn.indexer.wq_b", # DSA
         ),
 
+        MODEL_TENSOR.INDEXER_Q_PROJ: (
+            "model.layers.{bid}.self_attn.index_q_proj", # MSA
+        ),
+
+        MODEL_TENSOR.INDEXER_K_PROJ: (
+            "model.layers.{bid}.self_attn.index_k_proj", # MSA
+        ),
+
+        MODEL_TENSOR.INDEXER_Q_NORM: (
+            "model.layers.{bid}.self_attn.index_q_norm", # MSA
+        ),
+
         ############################################################################
         # TODO: these do not belong to block_mappings_cfg - move them to mappings_cfg
         MODEL_TENSOR.ENC_OUTPUT_NORM: (
index 9aa3dace5ce0209e718e16af50ace5503a2526ba..39bf2c79590b1be0183cedb1b492378ffa11294d 100644 (file)
@@ -127,6 +127,7 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
     { LLM_ARCH_GROVEMOE,         "grovemoe"         },
     { LLM_ARCH_APERTUS,          "apertus"          },
     { LLM_ARCH_MINIMAX_M2,       "minimax-m2"       },
+    { LLM_ARCH_MINIMAX_M3,       "minimax-m3"       },
     { LLM_ARCH_COGVLM,           "cogvlm"           },
     { LLM_ARCH_RND1,             "rnd1"             },
     { LLM_ARCH_PANGU_EMBED,      "pangu-embedded"   },
@@ -253,6 +254,8 @@ static const std::map<llm_kv, const char *> LLM_KV_NAMES = {
     { LLM_KV_ATTENTION_INDEXER_HEAD_COUNT,           "%s.attention.indexer.head_count"           },
     { LLM_KV_ATTENTION_INDEXER_KEY_LENGTH,           "%s.attention.indexer.key_length"           },
     { LLM_KV_ATTENTION_INDEXER_TOP_K,                "%s.attention.indexer.top_k"                },
+    { LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE,           "%s.attention.indexer.block_size"           },
+    { LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS,         "%s.attention.indexer.local_blocks"         },
     { LLM_KV_ATTENTION_INDEXER_TYPES,                "%s.attention.indexer.types"                },
     { LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT,           "%s.attention.output_group_count"           },
     { LLM_KV_ATTENTION_OUTPUT_LORA_RANK,             "%s.attention.output_lora_rank"             },
@@ -597,6 +600,9 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
     { LLM_TENSOR_INDEXER_PROJ,                           "blk.%d.indexer.proj" },
     { LLM_TENSOR_INDEXER_ATTN_K,                         "blk.%d.indexer.attn_k" },
     { LLM_TENSOR_INDEXER_ATTN_Q_B,                       "blk.%d.indexer.attn_q_b" },
+    { LLM_TENSOR_INDEXER_Q_PROJ,                         "blk.%d.indexer.q_proj" },
+    { LLM_TENSOR_INDEXER_K_PROJ,                         "blk.%d.indexer.k_proj" },
+    { LLM_TENSOR_INDEXER_Q_NORM,                         "blk.%d.indexer.q_norm" },
     { LLM_TENSOR_INDEXER_COMPRESSOR_WKV,                 "blk.%d.indexer_compressor_kv" },
     { LLM_TENSOR_INDEXER_COMPRESSOR_WGATE,               "blk.%d.indexer_compressor_gate" },
     { LLM_TENSOR_INDEXER_COMPRESSOR_APE,                 "blk.%d.indexer_compressor_ape" },
@@ -832,6 +838,9 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
     {LLM_TENSOR_INDEXER_PROJ,               {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
     {LLM_TENSOR_INDEXER_ATTN_K,             {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
     {LLM_TENSOR_INDEXER_ATTN_Q_B,           {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
+    {LLM_TENSOR_INDEXER_Q_PROJ,             {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
+    {LLM_TENSOR_INDEXER_K_PROJ,             {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
+    {LLM_TENSOR_INDEXER_Q_NORM,             {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}},
     {LLM_TENSOR_INDEXER_COMPRESSOR_WKV,     {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
     {LLM_TENSOR_INDEXER_COMPRESSOR_WGATE,   {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
     {LLM_TENSOR_INDEXER_COMPRESSOR_APE,     {LLM_TENSOR_LAYER_REPEATING, GGML_OP_GET_ROWS}},
@@ -1001,6 +1010,7 @@ bool llm_arch_supports_sm_tensor(const llm_arch & arch) {
         case LLM_ARCH_LFM2:
         case LLM_ARCH_LFM2MOE:
         case LLM_ARCH_MINIMAX_M2:
+        case LLM_ARCH_MINIMAX_M3:
         case LLM_ARCH_MISTRAL4:
         case LLM_ARCH_KIMI_LINEAR:
             return false;
index 39c55a66a94338fdca8c46c7725fb0fac1597019..2e3916a0beee7202ac2801e69950a4bd1cc68cd0 100644 (file)
@@ -146,6 +146,7 @@ enum llm_arch {
     LLM_ARCH_TALKIE,
     LLM_ARCH_MELLUM,
     LLM_ARCH_EAGLE3,
+    LLM_ARCH_MINIMAX_M3,
     LLM_ARCH_DFLASH,
     LLM_ARCH_UNKNOWN,
 };
@@ -258,6 +259,8 @@ enum llm_kv {
     LLM_KV_ATTENTION_INDEXER_HEAD_COUNT,
     LLM_KV_ATTENTION_INDEXER_KEY_LENGTH,
     LLM_KV_ATTENTION_INDEXER_TOP_K,
+    LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE,
+    LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS,
     LLM_KV_ATTENTION_INDEXER_TYPES,
     LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT,
     LLM_KV_ATTENTION_OUTPUT_LORA_RANK,
@@ -597,6 +600,9 @@ enum llm_tensor {
     LLM_TENSOR_INDEXER_PROJ,
     LLM_TENSOR_INDEXER_ATTN_K,
     LLM_TENSOR_INDEXER_ATTN_Q_B,
+    LLM_TENSOR_INDEXER_Q_PROJ,
+    LLM_TENSOR_INDEXER_K_PROJ,
+    LLM_TENSOR_INDEXER_Q_NORM,
     LLM_TENSOR_INDEXER_COMPRESSOR_WKV,
     LLM_TENSOR_INDEXER_COMPRESSOR_WGATE,
     LLM_TENSOR_INDEXER_COMPRESSOR_APE,
index eed041eef4e7a9ac418d3a6ee616faf11133ec22..c512477c0eab993904d08eff2a48f7dab107eb2e 100644 (file)
@@ -2338,7 +2338,8 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const {
         model.arch == LLM_ARCH_KIMI_LINEAR ||
         model.arch == LLM_ARCH_QWEN35 ||
         model.arch == LLM_ARCH_QWEN35MOE ||
-        model.arch == LLM_ARCH_DEEPSEEK4) {
+        model.arch == LLM_ARCH_DEEPSEEK4 ||
+        model.arch == LLM_ARCH_MINIMAX_M3) {
         return std::max<uint32_t>(n_tokens * 40, 32u * model.n_tensors());
     }
     uint32_t res = std::max<uint32_t>(1024u, 8u*model.n_tensors());
index c8ecb0a2854cc2a43ca3102ebafea89e386d917f..6d1c8f4e42a83a746e8c3ead00d216b2335f2df4 100644 (file)
@@ -1709,6 +1709,17 @@ ggml_tensor * llm_graph_context::build_ffn(
                 cur = ggml_swiglu(ctx0, cur);
                 cb(cur, "ffn_swiglu", il);
             } break;
+        case LLM_FFN_SWIGLU_OAI_MOE:
+            if (gate && type_gate == LLM_FFN_PAR) {
+                // same alpha/limit constants as gpt-oss
+                const float alpha = 1.702f;
+                const float limit = 7.0f;
+                cur = ggml_swiglu_oai(ctx0, cur, tmp, alpha, limit);
+                cb(cur, "ffn_swiglu_oai", il);
+                type_gate = LLM_FFN_SEQ;
+            } else {
+                GGML_ABORT("LLM_FFN_SWIGLU_OAI_MOE requires a parallel gate");
+            } break;
         case LLM_FFN_GEGLU:
             {
                 cur = ggml_geglu(ctx0, cur);
@@ -2668,7 +2679,7 @@ ggml_tensor * llm_graph_context::build_attn(
         ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il));
     }
 
-    const auto & kq_mask = inp->get_kq_mask();
+    ggml_tensor * kq_mask = inp->get_kq_mask();
 
     ggml_tensor * q = q_cur;
     ggml_tensor * k = mctx_cur->get_k(ctx0, il);
index 846d4c69a6265b1cb7663605befe757c6b2f75bd..50af97f358c339369b637587aed6923709c866c1 100644 (file)
@@ -180,6 +180,16 @@ 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 747754fc0d0bbce1ef5983ba4100414186627f87..727df6ca21e288120a85ff8979e21db76848f1f1 100644 (file)
@@ -226,6 +226,11 @@ struct llama_hparams {
     uint32_t indexer_n_head    = 0;
     uint32_t indexer_head_size = 0;
     uint32_t indexer_top_k     = 0;
+    // 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
@@ -350,6 +355,9 @@ 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;
index e25464c597ace910245279496e6bb9c826bb559c..44cb1668dacf74d899a1658a9a51ba825a281510 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(2u*(1 + n_stream)*n_layer*ggml_tensor_overhead()),
+                /*.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_buffer =*/ NULL,
                 /*.no_alloc   =*/ true,
             };
@@ -242,9 +242,25 @@ 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_stream, v_stream, });
+        layers.push_back({ il, k, v, k_idx, k_stream, v_stream, k_idx_stream });
     }
 
     if (reuse) {
@@ -293,13 +309,24 @@ 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     = 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);
+        }
 
-        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));
+        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());
     }
 
     // TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
@@ -392,6 +419,39 @@ 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]];
@@ -846,6 +906,10 @@ 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]);
+                }
             }
         }
     }
@@ -994,6 +1058,44 @@ 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 ->
@@ -1002,11 +1104,6 @@ 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
@@ -1113,6 +1210,15 @@ 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);
 
@@ -1156,7 +1262,8 @@ 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);
 
-            seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1);
+            // 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));
         }
     }
 
@@ -1176,6 +1283,12 @@ 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;
 }
 
@@ -1292,6 +1405,23 @@ 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);
 
@@ -1393,6 +1523,28 @@ 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;
 
@@ -1827,6 +1979,18 @@ 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,
@@ -2139,6 +2303,36 @@ 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;
@@ -2387,6 +2581,68 @@ 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;
@@ -2588,6 +2844,10 @@ 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]);
 }
@@ -2596,6 +2856,10 @@ 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 531d99dbdec185d92a7b4126b2f61e432949e161..d5a92f4405b572ebc5357516b8a1420d97bd5ca5 100644 (file)
@@ -173,10 +173,12 @@ 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
@@ -228,9 +230,11 @@ 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
@@ -259,6 +263,9 @@ 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;
 
@@ -291,6 +298,7 @@ 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,
@@ -370,6 +378,7 @@ 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
@@ -379,6 +388,7 @@ 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 d26e2ff7af629d8f76a0e586c88e636839ddac16..3812c594e7951e8040a1d0806fb70a11e66a4309 100644 (file)
@@ -281,6 +281,8 @@ void llama_model_saver::add_kv_from_model() {
     add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT,      hparams.indexer_n_head);
     add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH,      hparams.indexer_head_size);
     add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K,           hparams.indexer_top_k);
+    add_kv(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE,      hparams.indexer_block_size);
+    add_kv(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS,    hparams.indexer_local_blocks);
     add_kv(LLM_KV_ATTENTION_INDEXER_TYPES,           hparams.is_indexer_full_impl, true);
     add_kv(LLM_KV_ATTENTION_RECURRENT_LAYERS,        hparams.is_recr_impl, true);
 
index b100f60181502d7aa17c71ddb8ebf207602f49a7..51796921081fcf68764c90fe3e64c7718d8cfadf 100644 (file)
@@ -285,6 +285,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params
             return new llama_model_apertus(params);
         case LLM_ARCH_MINIMAX_M2:
             return new llama_model_minimax_m2(params);
+        case LLM_ARCH_MINIMAX_M3:
+            return new llama_model_minimax_m3(params);
         case LLM_ARCH_COGVLM:
             return new llama_model_cogvlm(params);
         case LLM_ARCH_PANGU_EMBED:
@@ -818,6 +820,7 @@ const char * llm_type_name(llm_type type) {
         case LLM_TYPE_122B_A10B:     return "122B.A10B";
         case LLM_TYPE_196B_A11B:     return "196B.A11B";
         case LLM_TYPE_230B_A10B:     return "230B.A10B";
+        case LLM_TYPE_428B_A23B:     return "428B.A23B";
         case LLM_TYPE_235B_A22B:     return "235B.A22B";
         case LLM_TYPE_300B_A47B:     return "300B.A47B";
         case LLM_TYPE_310B_A15B:     return "310B.A15B";
@@ -2550,6 +2553,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) {
         case LLM_ARCH_GROVEMOE:
         case LLM_ARCH_APERTUS:
         case LLM_ARCH_MINIMAX_M2:
+        case LLM_ARCH_MINIMAX_M3:
         case LLM_ARCH_COGVLM:
         case LLM_ARCH_PANGU_EMBED:
         case LLM_ARCH_AFMOE:
index 45b054cedf1d1e6accc7cf8aafcbae374614e64f..36d0480e5eb706b50b8a88fe5de2cce72a859c2a 100644 (file)
@@ -134,6 +134,7 @@ enum llm_type {
     LLM_TYPE_122B_A10B, // Qwen3.5
     LLM_TYPE_196B_A11B, // Step3.5-Flash
     LLM_TYPE_230B_A10B, // Minimax M2
+    LLM_TYPE_428B_A23B, // Minimax M3
     LLM_TYPE_235B_A22B,
     LLM_TYPE_300B_A47B, // Ernie MoE big
     LLM_TYPE_310B_A15B, // /MiMo-V2-Flash
@@ -515,6 +516,12 @@ struct llama_layer {
     struct ggml_tensor * indexer_attn_k   = nullptr;
     struct ggml_tensor * indexer_attn_q_b = nullptr; // note: for lora a/b, not bias
 
+    // MSA
+    struct ggml_tensor * index_q_proj = nullptr;
+    struct ggml_tensor * index_k_proj = nullptr;
+    struct ggml_tensor * index_q_norm = nullptr;
+    struct ggml_tensor * index_k_norm = nullptr;
+
     // gemma4 layer output scale, reused for talkie embedding skip scale
     struct ggml_tensor * out_scale = nullptr;
 
index 7b312d1d88e12ce84974786b8d74877327aaed18..9164a4dd888d215b5cd595a0036c26bd2474d3bc 100644 (file)
@@ -2809,6 +2809,7 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) {
                     || t.first == "<turn|>"          // gemma4
                     || t.first == "<|tool_response>" // gemma4
                     || t.first == "<|end▁of▁sentence|>" // deepseek-ocr
+                    || t.first == "[e~[" // minimax-m2/m3
                ) {
                 special_eog_ids.insert(t.second);
                 if ((attr & LLAMA_TOKEN_ATTR_CONTROL) == 0) {
diff --git a/src/models/minimax-m3.cpp b/src/models/minimax-m3.cpp
new file mode 100644 (file)
index 0000000..6068fc6
--- /dev/null
@@ -0,0 +1,562 @@
+#include "models.h"
+#include "llama-kv-cache.h"
+#include <cmath>
+#include <vector>
+#include <algorithm>
+#include <cstdint>
+
+// 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.
+
+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);
+    ml.get_key(LLM_KV_LEADING_DENSE_BLOCK_COUNT,   hparams.n_layer_dense_lead, false);
+    ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH,  hparams.n_ff_exp);
+    ml.get_key(LLM_KV_EXPERT_SHARED_COUNT,         hparams.n_expert_shared);
+    ml.get_key(LLM_KV_EXPERT_WEIGHTS_SCALE,        hparams.expert_weights_scale, false);
+    ml.get_key(LLM_KV_EXPERT_WEIGHTS_NORM,         hparams.expert_weights_norm, false);
+    ml.get_key(LLM_KV_EXPERT_GATING_FUNC,          hparams.expert_gating_func);
+    ml.get_key(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT,    hparams.indexer_n_head);
+    ml.get_key(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH,    hparams.indexer_head_size);
+    ml.get_key(LLM_KV_ATTENTION_INDEXER_TOP_K,         hparams.indexer_top_k);
+    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;
+        default: type = LLM_TYPE_UNKNOWN;
+    }
+}
+
+void llama_model_minimax_m3::load_arch_tensors(llama_model_loader &) {
+    LLAMA_LOAD_LOCALS;
+    const int64_t n_expert_shared = hparams.n_expert_shared;
+    const int64_t n_ff_exp        = hparams.n_ff_exp;
+
+    tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0);
+
+    // output
+    output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0);
+    output      = create_tensor(tn(LLM_TENSOR_OUTPUT,      "weight"), {n_embd, n_vocab}, 0);
+
+    for (int i = 0; i < n_layer; ++i) {
+        auto & layer = layers[i];
+
+        create_tensor_qkv(layer, i, n_embd, n_embd_head_k * n_head, n_embd_gqa, n_embd_gqa, 0);
+        layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), { n_embd_head_k * n_head, n_embd }, 0);
+
+        layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0);
+        // per-head QK-norm: a single head_dim vector applied to every head
+        layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), {n_embd_head_k}, 0);
+        layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), {n_embd_head_k}, 0);
+
+        layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0);
+
+        if (i < (int) hparams.n_layer_dense_lead) {
+            // leading dense layers
+            layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd,   n_ff}, 0);
+            layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), {  n_ff, n_embd}, 0);
+            layer.ffn_up   = create_tensor(tn(LLM_TENSOR_FFN_UP,   "weight", i), {n_embd,   n_ff}, 0);
+        } else {
+            // routed experts
+            layer.ffn_gate_inp    = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP,    "weight", i), {n_embd, n_expert}, 0);
+            layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias",   i), {n_expert}, 0);
+            layer.ffn_gate_exps   = create_tensor(tn(LLM_TENSOR_FFN_GATE_EXPS,   "weight", i), {n_embd, n_ff_exp, n_expert}, 0);
+            layer.ffn_down_exps   = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS,   "weight", i), {n_ff_exp, n_embd, n_expert}, 0);
+            layer.ffn_up_exps     = create_tensor(tn(LLM_TENSOR_FFN_UP_EXPS,     "weight", i), {n_embd, n_ff_exp, n_expert}, 0);
+
+            // shared expert
+            layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, 0);
+            layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), {        n_ff_exp * n_expert_shared, n_embd}, 0);
+            layer.ffn_up_shexp   = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP,   "weight", i), {n_embd, n_ff_exp * n_expert_shared}, 0);
+
+            // indexer
+            layer.index_q_proj = create_tensor(tn(LLM_TENSOR_INDEXER_Q_PROJ, "weight", i), {n_embd, hparams.indexer_n_head * hparams.indexer_head_size}, 0);
+            layer.index_k_proj = create_tensor(tn(LLM_TENSOR_INDEXER_K_PROJ, "weight", i), {n_embd, hparams.indexer_head_size}, 0);
+            layer.index_q_norm = create_tensor(tn(LLM_TENSOR_INDEXER_Q_NORM, "weight", i), {hparams.indexer_head_size}, 0);
+            layer.index_k_norm = create_tensor(tn(LLM_TENSOR_INDEXER_K_NORM, "weight", i), {hparams.indexer_head_size}, 0);
+        }
+    }
+}
+
+std::unique_ptr<llm_graph_context> llama_model_minimax_m3::build_arch_graph(const llm_graph_params & params) const {
+    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 {
+public:
+    llm_graph_input_msa_local(int blk, int local, int64_t nblk) : blk(blk), local(local), nblk(nblk) {}
+
+    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;
+                }
+            }
+        }
+        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
+    bool can_reuse(const llm_graph_params & params) override {
+        const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);
+
+        bool res = true;
+        res &= bias->ne[1] == params.ubatch.n_tokens;
+        res &= bias->ne[0] * blk == (int64_t) mctx->get_n_kv();
+        return res;
+    }
+
+    ggml_tensor * bias = nullptr;
+    int     blk;
+    int     local;
+    int64_t nblk;
+};
+
+// pooled score of a block with no visible token: -inf from the mask, or -FLT_MAX from the
+// max-pool identity when every element of the block is -inf
+static inline bool msa_score_masked(float x) { return x <= -1e30f; }
+
+// MSA block selection (batch regime)
+// CPU custom op, the token-level expansion and the combination with the causal mask happen on the GPU.
+static void msa_block_mask_op(struct ggml_tensor * dst, int ith, int nth, void * userdata) {
+    const struct ggml_tensor * bs   = dst->src[0];
+    const struct ggml_tensor * bias = dst->src[1];
+    const msa_params * p = (const msa_params *) userdata;
+
+    const int nblk = (int) bs->ne[0];
+    const int Hd   = (int) bs->ne[1];
+    const int S    = (int) bs->ne[2];
+
+    GGML_ASSERT(bs->type   == GGML_TYPE_F32 && ggml_is_contiguous(bs));
+    GGML_ASSERT(bias->type == GGML_TYPE_F32 && ggml_is_contiguous(bias));
+    GGML_ASSERT(dst->type  == GGML_TYPE_F16 && ggml_is_contiguous(dst));
+    GGML_ASSERT(dst->ne[0] == nblk && dst->ne[1] == S && dst->ne[2] == Hd);
+    GGML_ASSERT(bias->ne[0] == nblk && bias->ne[1] == S);
+
+    const int topk = p->topk_blocks < nblk ? p->topk_blocks : nblk;
+
+    const ggml_fp16_t f16_zero = ggml_fp32_to_fp16(0.0f);
+    const ggml_fp16_t f16_ninf = ggml_fp32_to_fp16(-INFINITY);
+
+    std::vector<float> rank(nblk);
+    std::vector<char>  valid(nblk);
+    std::vector<int>   ord(nblk);
+
+    ggml_fp16_t * out = (ggml_fp16_t *) dst->data;
+
+    for (int i = ith; i < S; i += nth) {
+        const float * bias_col = (const float *) bias->data + (size_t) i * nblk;
+        for (int h = 0; h < Hd; ++h) {
+            const float * bs_col = (const float *) bs->data + ((size_t) i * Hd + h) * nblk;
+
+            for (int bk = 0; bk < nblk; ++bk) {
+                // a block is selectable if it has a visible token or is locally forced
+                valid[bk] = !msa_score_masked(bs_col[bk]) || bias_col[bk] > 0.0f;
+                rank [bk] = bias_col[bk] > 0.0f ? bias_col[bk] : bs_col[bk];
+                ord  [bk] = bk;
+            }
+
+            std::partial_sort(ord.begin(), ord.begin() + topk, ord.end(),
+                              [&](int a, int b) { return rank[a] > rank[b]; });
+
+            ggml_fp16_t * dst_col = out + ((size_t) h * S + i) * nblk;
+            for (int bk = 0; bk < nblk; ++bk) {
+                dst_col[bk] = f16_ninf;
+            }
+            for (int t = 0; t < topk; ++t) {
+                const int bk = ord[t];
+                if (!valid[bk]) {
+                    break;   // sorted desc: first invalid -> fewer than topk selectable blocks
+                }
+                dst_col[bk] = f16_zero;
+            }
+        }
+    }
+}
+
+// One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3])
+ggml_tensor * llama_model_minimax_m3::graph::build_attn_msa_fa(
+        ggml_tensor * q_cur,   // [D, HQ, T]
+        ggml_tensor * k,       // [D, n_keys, 1, C]
+        ggml_tensor * v,       // [D, n_keys, 1, C]
+        ggml_tensor * mask,    // [n_keys, R, 1, C] f16, contiguous
+        int64_t Gp, float kq_scale, int il) const {
+
+    const int64_t D  = q_cur->ne[0];
+    const int64_t HQ = q_cur->ne[1];
+    const int64_t T  = q_cur->ne[2];
+    const int64_t C  = k->ne[3];
+    const int64_t R  = HQ*T/(Gp*C);
+    GGML_ASSERT(Gp*C*R == HQ*T);
+    GGML_ASSERT(mask->type == GGML_TYPE_F16);
+
+    // [D, HQ, T] -> [D, Gp, C, R] -> [D, R, Gp, C]
+    // batch  (C=HKV,   R=T): channel = group
+    // decode (C=HKV*ns, R=1): channel = (group, stream), group innermost
+    ggml_tensor * q = ggml_reshape_4d(ctx0, q_cur, D, Gp, C, R);
+    q = ggml_permute(ctx0, q, 0, 2, 3, 1);
+
+    ggml_tensor * o = ggml_flash_attn_ext(ctx0, q, k, v, mask, kq_scale,
+                                          hparams.f_max_alibi_bias, 0.0f);
+    ggml_flash_attn_ext_set_prec(o, GGML_PREC_F32);
+    cb(o, "msa_fattn", il);
+
+    // [D, Gp, R, C] -> [D, Gp, C, R] -> [n_embd, T]
+    o = ggml_permute(ctx0, o, 0, 1, 3, 2);
+    if (!ggml_is_contiguous(o)) {
+        o = ggml_cont(ctx0, o);   // no-op layout at decode (R == 1), copy at batch
+    }
+    return ggml_reshape_2d(ctx0, o, D*HQ, T);
+}
+
+llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) {
+    const int64_t n_embd_head = hparams.n_embd_head_v();
+    const auto & mm = static_cast<const llama_model_minimax_m3 &>(model);
+
+    GGML_ASSERT(n_embd_head == hparams.n_embd_head_k());
+    // partial rotary: head_dim != n_rot, so don't assert n_embd_head == n_rot
+
+    ggml_tensor * cur;
+    ggml_tensor * inpL;
+
+    inpL = build_inp_embd(model.tok_embd);
+
+    ggml_tensor * inp_pos = build_inp_pos();
+    auto inp_attn = build_attn_inp_kv();
+
+    // 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
+    // to absolute KV cache slots, which equal positions only for append-only per-stream
+    // caches either a single sequence, or multiple sequences with kv_unified == false (each
+    // stream then has its own slot space). A unified cache with multiple sequences
+    // interleaves slots and would silently break block anchoring so it falls back to dense.
+    const bool fa_on       = cparams.flash_attn;
+    const bool streams_ok  = cparams.n_seq_max == 1 || !cparams.kv_unified;
+    const bool msa_enabled = fa_on && streams_ok;
+
+    static bool warned_no_fa = false;
+    if (!fa_on && !warned_no_fa) {
+        LLAMA_LOG_WARN("%s: flash attention disabled; MSA requires it -> running DENSE attention "
+                       "(output may be degraded). Enable flash attention for MSA.\n", __func__);
+        warned_no_fa = true;
+    }
+    static bool warned_unified = false;
+    if (fa_on && !streams_ok && !warned_unified) {
+        LLAMA_LOG_WARN("%s: unified KV cache with n_seq_max > 1; MSA needs per-sequence streams "
+                       "-> running DENSE attention. Output may be degraded. Drop --kv-unified to enable MSA.\n", __func__);
+        warned_unified = true;
+    }
+
+    // hoisted per-graph MSA state (shared by every sparse layer)
+    llm_graph_input_msa_local * msa_loc = nullptr;
+    ggml_tensor * msa_kqm = nullptr;
+    ggml_tensor * msa_mf  = nullptr;
+    int64_t n_kv = 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) {
+        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;
+        msa_decode = n_tps == 1;
+
+        msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32);
+
+        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));
+    }
+
+    ggml_tensor * inp_out_ids = build_inp_out_ids();
+
+    for (int il = 0; il < n_layer; ++il) {
+        ggml_tensor * inpSA = inpL;
+
+        // self-attention
+        {
+            cur = build_norm(inpL, model.layers[il].attn_norm, NULL, LLM_NORM_RMS, il);
+            cb(cur, "attn_norm", il);
+
+            auto [Qcur, Kcur, Vcur] = build_qkv(model.layers[il], cur,
+                    n_embd_head, n_head, n_head_kv, il);
+
+            // per-head QK RMSNorm (weights already include Gemma's +1)
+            Qcur = build_norm(Qcur, model.layers[il].attn_q_norm, NULL, LLM_NORM_RMS, il);
+            cb(Qcur, "Qcur_normed", il);
+            Kcur = build_norm(Kcur, model.layers[il].attn_k_norm, NULL, LLM_NORM_RMS, il);
+            cb(Kcur, "Kcur_normed", il);
+
+            // partial rotary: only the first n_rot dims are rotated
+            Qcur = ggml_rope_ext(
+                ctx0, Qcur, inp_pos, nullptr,
+                n_rot, rope_type, n_ctx_orig, freq_base, freq_scale,
+                ext_factor, attn_factor, beta_fast, beta_slow);
+            Kcur = ggml_rope_ext(
+                ctx0, Kcur, inp_pos, nullptr,
+                n_rot, rope_type, n_ctx_orig, freq_base, freq_scale,
+                ext_factor, attn_factor, beta_fast, beta_slow);
+
+            cb(Qcur, "Qcur", il);
+            cb(Kcur, "Kcur", il);
+            cb(Vcur, "Vcur", il);
+
+            const bool is_sparse = msa_enabled && il >= (int) hparams.n_layer_dense_lead;
+
+            if (!is_sparse) {
+                cur = build_attn(inp_attn, model.layers[il].wo, NULL, model.layers[il].wo_s,
+                        Qcur, Kcur, Vcur, nullptr, nullptr, nullptr,
+                        1.0f/sqrtf(float(n_embd_head)), il);
+            } else {
+                const int64_t n_idx_dim = hparams.indexer_head_size;   // 128
+
+                GGML_ASSERT(!inp_attn->self_k_rot && !inp_attn->self_v_rot && "MSA: attn-rot not supported");
+
+                // Index Branch, project, norm, partial RoPE, cache
+                ggml_tensor * iq = build_lora_mm(model.layers[il].index_q_proj, cur);
+                ggml_tensor * ik = build_lora_mm(model.layers[il].index_k_proj, cur);
+                iq = ggml_reshape_3d(ctx0, iq, n_idx_dim, Hd, n_tokens);
+                ik = ggml_reshape_3d(ctx0, ik, n_idx_dim, 1,  n_tokens);
+                iq = build_norm(iq, model.layers[il].index_q_norm, NULL, LLM_NORM_RMS, il);  // +1 baked
+                ik = build_norm(ik, model.layers[il].index_k_norm, NULL, LLM_NORM_RMS, il);
+                iq = ggml_rope_ext(ctx0, iq, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig,
+                                   freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow);
+                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);
+
+                // Main branch: store K/V, take cache views
+                ggml_build_forward_expand(gf, Qcur);
+                ggml_build_forward_expand(gf, Kcur);
+                ggml_build_forward_expand(gf, Vcur);
+                ggml_build_forward_expand(gf, mctx_cur->cpy_k(ctx0, Kcur, inp_attn->get_k_idxs(), il));
+                ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, Vcur, inp_attn->get_v_idxs(), il));
+                ggml_tensor * k = mctx_cur->get_k(ctx0, il);
+                ggml_tensor * v = mctx_cur->get_v(ctx0, il);
+                GGML_ASSERT(!(v->nb[1] > v->nb[2]) && "MSA assumes v_trans=false (FA on)");
+
+                const int64_t D   = k->ne[0];
+                const int64_t HKV = k->ne[1];
+                const int64_t Gp  = n_head/HKV;
+                GGML_ASSERT(HKV == Hd && "MSA: one indexer head per GQA group");
+                GGML_ASSERT(k->ne[3] == ns);
+                const int K = mm.msa_p.topk_blocks < (int) nblk ? mm.msa_p.topk_blocks : (int) nblk;
+
+                const float kq_scale = 1.0f/sqrtf(float(n_embd_head));
+
+                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);
+                    ggml_tensor * iq4 = ggml_reshape_4d(ctx0, iq, n_idx_dim, Hd, 1, ns);
+                    ggml_tensor * sc  = ggml_mul_mat(ctx0, ikv4, iq4);
+                    ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
+                    sc = ggml_add_inplace(ctx0, sc, msa_mf);
+                    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);
+
+                    // 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)
+                    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 * 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 * 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);
+
+                    // fold (group, stream) onto the FA channel dim
+                    const ggml_type kt = ggml_is_quantized(k->type) ? GGML_TYPE_F16 : k->type;
+                    const ggml_type vt = ggml_is_quantized(v->type) ? GGML_TYPE_F16 : v->type;
+                    ggml_tensor * kfa = ggml_reshape_4d(ctx0, kg, D, (int64_t) blk*K, 1, Hd*ns);
+                    ggml_tensor * vfa = ggml_reshape_4d(ctx0, vg, D, (int64_t) blk*K, 1, Hd*ns);
+                    if (kfa->type != kt) { kfa = ggml_cast(ctx0, kfa, kt); }
+                    if (vfa->type != vt) { vfa = ggml_cast(ctx0, vfa, vt); }
+                    // the FA mask must be F16
+                    ggml_tensor * mfa = ggml_cast(ctx0, ggml_reshape_4d(ctx0, mg, (int64_t) blk*K, 1, 1, Hd*ns), GGML_TYPE_F16);
+
+                    cur = build_attn_msa_fa(Qcur, kfa, vfa, mfa, Gp, kq_scale, il);
+                } else {
+                    // batch: per-stream loop
+                    std::vector<ggml_tensor *> outs(ns);
+                    for (int64_t st = 0; st < ns; ++st) {
+                        ggml_tensor * iq_s = ggml_view_3d(ctx0, iq, n_idx_dim, Hd, n_tps,
+                                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_2d(ctx0, msa_loc->bias, nblk, n_tps,
+                                msa_loc->bias->nb[1], st*n_tps*msa_loc->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,
+                                k->nb[1], k->nb[2], k->nb[3], st*k->nb[3]);
+                        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)
+                        // scores are unscaled, only the top-k ordering matters
+                        ggml_tensor * sc = ggml_mul_mat(ctx0, ik_s,
+                                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);
+                        ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
+                        cb(bs, "msa_bs", il);
+
+                        // block-level 0/-inf keep mask on the CPU, tiny transfer
+                        ggml_tensor * srcs[2] = { bs, bias_s };
+                        ggml_tensor * bm = ggml_custom_4d(ctx0, GGML_TYPE_F16,
+                                nblk, n_tps, Hd, 1,
+                                srcs, 2, msa_block_mask_op, GGML_N_TASKS_MAX,
+                                const_cast<msa_params *>(&mm.msa_p));
+                        cb(bm, "msa_block_mask", il);
+
+                        // expand block -> token granularity on the GPU (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);
+                        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);
+                        cb(mask4, "msa_mask4", il);
+
+                        // cache views with groups on ne[3];
+                        ggml_tensor * kfa = ggml_permute(ctx0, k_s, 0, 3, 1, 2);
+                        ggml_tensor * vfa = ggml_permute(ctx0, v_s, 0, 3, 1, 2);
+
+                        outs[st] = build_attn_msa_fa(q_s, kfa, vfa, mask4, Gp, kq_scale, il);
+                    }
+                    cur = outs[0];
+                    for (int64_t st = 1; st < ns; ++st) {
+                        cur = ggml_concat(ctx0, cur, outs[st], 1);
+                    }
+                }
+
+                cb(cur, "kqv_out", il);
+                if (model.layers[il].wo) {
+                    cur = build_lora_mm(model.layers[il].wo, cur, model.layers[il].wo_s);
+                }
+            }
+        }
+
+        if (il == n_layer - 1 && inp_out_ids) {
+            cur   = ggml_get_rows(ctx0,   cur, inp_out_ids);
+            inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids);
+        }
+
+        ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA);
+        cb(ffn_inp, "ffn_inp", il);
+
+        cur = build_norm(ffn_inp, model.layers[il].ffn_norm, NULL, LLM_NORM_RMS, il);
+        cb(cur, "ffn_norm", il);
+
+        if ((uint32_t) il < hparams.n_layer_dense_lead) {
+            // leading dense FFN (swigluoai)
+            cur = build_ffn(cur,
+                    model.layers[il].ffn_up,   NULL, NULL,
+                    model.layers[il].ffn_gate, NULL, NULL,
+                    model.layers[il].ffn_down, NULL, NULL,
+                    NULL,
+                    LLM_FFN_SWIGLU_OAI_MOE, LLM_FFN_PAR, il);
+            cb(cur, "ffn_out", il);
+        } else {
+            // routed experts (swigluoai MoE)
+            ggml_tensor * moe_out = build_moe_ffn(cur,
+                    model.layers[il].ffn_gate_inp,
+                    model.layers[il].ffn_up_exps,
+                    model.layers[il].ffn_gate_exps,
+                    model.layers[il].ffn_down_exps,
+                    model.layers[il].ffn_exp_probs_b,
+                    n_expert, n_expert_used,
+                    LLM_FFN_SWIGLU_OAI_MOE, hparams.expert_weights_norm,
+                    hparams.expert_weights_scale,
+                    (llama_expert_gating_func_type) hparams.expert_gating_func,
+                    il);
+            cb(moe_out, "ffn_moe_out", il);
+
+            // shared expert (swigluoai)
+            ggml_tensor * ffn_shexp = build_ffn(cur,
+                    model.layers[il].ffn_up_shexp,   NULL, NULL,
+                    model.layers[il].ffn_gate_shexp, NULL, NULL,
+                    model.layers[il].ffn_down_shexp, NULL, NULL,
+                    NULL,
+                    LLM_FFN_SWIGLU_OAI_MOE, LLM_FFN_PAR, il);
+            cb(ffn_shexp, "ffn_shexp", il);
+
+            cur = ggml_add(ctx0, moe_out, ffn_shexp);
+            cb(cur, "ffn_out", il);
+        }
+
+        cur = ggml_add(ctx0, cur, ffn_inp);
+
+        cur = build_cvec(cur, il);
+        cb(cur, "l_out", il);
+
+        // input for next layer
+        inpL = cur;
+    }
+
+    cur = inpL;
+
+    cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1);
+    cb(cur, "result_norm", -1);
+    res->t_embd = cur;
+
+    // lm_head
+    cur = build_lora_mm(model.output, cur, model.output_s);
+    cb(cur, "result_output", -1);
+    res->t_logits = cur;
+
+    ggml_build_forward_expand(gf, cur);
+}
index 76daa8cc199458131433d2cfeabc70efc7aebae5..916459e127828dbb60211bf739da040e27ffc6a3 100644 (file)
@@ -1902,6 +1902,29 @@ struct llama_model_minimax_m2 : public llama_model_base {
     std::unique_ptr<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
 };
 
+struct msa_params {
+    int blk;
+    int topk_blocks;
+    int local;
+};
+
+struct llama_model_minimax_m3 : public llama_model_base {
+    llama_model_minimax_m3(const struct llama_model_params & params) : llama_model_base(params) {}
+    void load_arch_hparams(llama_model_loader & ml) override;
+    void load_arch_tensors(llama_model_loader & ml) override;
+    msa_params msa_p;
+    struct graph : public llm_graph_context {
+        graph(const llama_model & model, const llm_graph_params & params);
+
+        ggml_tensor * build_attn_msa_fa(
+                ggml_tensor * q_cur,   // [D, HQ, S] f32
+                ggml_tensor * k,       // [D, n_keys, 1, C]  C = HKV or HKV*n_stream
+                ggml_tensor * v,       // [D, n_keys, 1, C]
+                ggml_tensor * mask,    // [n_keys, R, 1, C] f16, R = HQ*T/(Gp*C)
+                int64_t Gp, float kq_scale, int il) const;
+    };
+    std::unique_ptr<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
+};
 
 struct llama_model_cogvlm : public llama_model_base {
     llama_model_cogvlm(const struct llama_model_params & params) : llama_model_base(params) {}
index 86c3051c5fe51d054a92e76cdc3bf550a7fd9196..d02e65c9ead01ef49b1a0a894484c859f765a486 100644 (file)
@@ -168,6 +168,9 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
         ms.add_kv(LLM_KV_ROPE_DIMENSION_COUNT,       uint32_t(64));
         ms.add_kv(LLM_KV_ATTENTION_KEY_LENGTH_MLA,   uint32_t(192));
         ms.add_kv(LLM_KV_ATTENTION_VALUE_LENGTH_MLA, uint32_t(128));
+    } else if (arch == LLM_ARCH_MINIMAX_M3) {
+        // partial rotary: n_rot must not exceed the indexer key length (64)
+        ms.add_kv(LLM_KV_ROPE_DIMENSION_COUNT,       uint32_t(64));
     }
     ms.add_kv(LLM_KV_ATTENTION_CLAMP_KQV,              1.0f);
     ms.add_kv(LLM_KV_ATTENTION_LAYERNORM_EPS,          1e-5f);
@@ -198,9 +201,13 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) {
         ms.add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, uint32_t(2));
     }
 
-    ms.add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, uint32_t(1));
-    ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, uint32_t(64));
-    ms.add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K,      uint32_t(8));
+    // MSA requires one indexer head per GQA (KV) head, unlike the DSA archs where the
+    // indexer head count is independent of the main attention head count.
+    ms.add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT,   arch == LLM_ARCH_MINIMAX_M3 ? n_head : uint32_t(1));
+    ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH,   uint32_t(64));
+    ms.add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K,        uint32_t(8));
+    ms.add_kv(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE,   uint32_t(4));
+    ms.add_kv(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, uint32_t(1));
     ms.add_kv(LLM_KV_ROPE_DIMENSION_SECTIONS, std::vector<uint32_t>({n_embd_head/4, n_embd_head/4, n_embd_head/4, n_embd_head/4}));
     ms.add_kv(LLM_KV_TOKENIZER_MODEL,         "no_vocab");
     // ms.add_kv(LLM_KV_DENSE_2_FEAT_OUT,     n_embd);
@@ -355,6 +362,7 @@ static bool moe_mandatory(const llm_arch arch) {
         case LLM_ARCH_LLADA_MOE:
         case LLM_ARCH_GROVEMOE:
         case LLM_ARCH_MINIMAX_M2:
+        case LLM_ARCH_MINIMAX_M3:
         case LLM_ARCH_RND1:
         case LLM_ARCH_PADDLEOCR:
         case LLM_ARCH_MIMO2: