From b1d4c6552489d27eb732c2804f189d9cc6ba99bd Mon Sep 17 00:00:00 2001 From: timkhronos Date: Sun, 26 Jul 2026 19:43:45 +0200 Subject: [PATCH] model: Add MiniMax-M3 (MSA: MiniMax Sparse Attention) support (#24908) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * 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 <1629204+CISC@users.noreply.github.com> * Remove whitespaces from constants.py Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Tighten comment in minimax.py Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * 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 <1629204+CISC@users.noreply.github.com> * Update gguf-py/gguf/gguf_writer.py Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Update gguf-py/gguf/tensor_mapping.py Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Update conversion/minimax.py Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Update conversion/minimax.py Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Remove whitespace in src/llama-kv-cache.cpp Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Remove Whitespace in Update src/llama-model.h Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * Remove whitespace in src/llama-hparams.h Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> * 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 Co-authored-by: Sigbjørn Skjæret <1629204+CISC@users.noreply.github.com> Co-authored-by: Xuan Son Nguyen --- conversion/__init__.py | 2 + conversion/base.py | 4 +- conversion/minimax.py | 37 ++- gguf-py/gguf/constants.py | 38 +++ gguf-py/gguf/gguf_writer.py | 6 + gguf-py/gguf/tensor_mapping.py | 15 +- src/llama-arch.cpp | 10 + src/llama-arch.h | 6 + src/llama-context.cpp | 3 +- src/llama-graph.cpp | 13 +- src/llama-hparams.cpp | 10 + src/llama-hparams.h | 8 + src/llama-kv-cache.cpp | 292 ++++++++++++++++- src/llama-kv-cache.h | 10 + src/llama-model-saver.cpp | 2 + src/llama-model.cpp | 4 + src/llama-model.h | 7 + src/llama-vocab.cpp | 1 + src/models/minimax-m3.cpp | 562 +++++++++++++++++++++++++++++++++ src/models/models.h | 23 ++ tests/test-llama-archs.cpp | 14 +- 21 files changed, 1044 insertions(+), 23 deletions(-) create mode 100644 src/models/minimax-m3.cpp diff --git a/conversion/__init__.py b/conversion/__init__.py index 7936f1159..0b08e6e57 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -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", diff --git a/conversion/base.py b/conversion/base.py index 051b8b4e5..a7cd3fd90 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -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}") diff --git a/conversion/minimax.py b/conversion/minimax.py index 4857775cb..cbbdfe3ae 100644 --- a/conversion/minimax.py +++ b/conversion/minimax.py @@ -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) diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index d55253e0e..66d50cca2 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -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, diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index bb2159670..ba08f8d65 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -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) diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index b5707f11f..59623accf 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -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: ( diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 9aa3dace5..39bf2c795 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -127,6 +127,7 @@ static const std::map 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_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_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_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; diff --git a/src/llama-arch.h b/src/llama-arch.h index 39c55a66a..2e3916a0b 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -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, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index eed041eef..c512477c0 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -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(n_tokens * 40, 32u * model.n_tensors()); } uint32_t res = std::max(1024u, 8u*model.n_tensors()); diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index c8ecb0a28..6d1c8f4e4 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -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); diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index 846d4c69a..50af97f35 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -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 diff --git a/src/llama-hparams.h b/src/llama-hparams.h index 747754fc0..727df6ca2 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -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; diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index e25464c59..44cb1668d 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -112,7 +112,7 @@ llama_kv_cache::llama_kv_cache( auto it = ctx_map.find(buft); if (it == ctx_map.end()) { ggml_init_params params = { - /*.mem_size =*/ size_t(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 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; - 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)); + constexpr float mib = 1024.0f * 1024.0f; + + const std::string k_log = format(", K (%s): %7.2f MiB", ggml_type_name(type_k), (float) memory_size_k / mib); + const std::string v_log = format(", V (%s): %7.2f MiB", ggml_type_name(type_v), (float) memory_size_v / mib); + + std::string k_idx_log; + if (memory_size_k_idx > 0) { + k_idx_log = format(", K_idx (%s): %7.2f MiB", ggml_type_name(GGML_TYPE_F32), (float) memory_size_k_idx / mib); + } + + LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs)%s%s%s\n", __func__, + (float) memory_size_total / mib, kv_size, (int) layers.size(), n_seq_max, n_stream, + k_log.c_str(), v_log.c_str(), k_idx_log.c_str()); } // 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::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); } diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index 531d99dbd..d5a92f440 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -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 k_stream; std::vector v_stream; + std::vector 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 diff --git a/src/llama-model-saver.cpp b/src/llama-model-saver.cpp index d26e2ff7a..3812c594e 100644 --- a/src/llama-model-saver.cpp +++ b/src/llama-model-saver.cpp @@ -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); diff --git a/src/llama-model.cpp b/src/llama-model.cpp index b100f6018..517969210 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -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: diff --git a/src/llama-model.h b/src/llama-model.h index 45b054ced..36d0480e5 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -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; diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index 7b312d1d8..9164a4dd8 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -2809,6 +2809,7 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { || t.first == "" // 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 index 000000000..6068fc6b8 --- /dev/null +++ b/src/models/minimax-m3.cpp @@ -0,0 +1,562 @@ +#include "models.h" +#include "llama-kv-cache.h" +#include +#include +#include +#include + +// 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 llama_model_minimax_m3::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +// per-query local-force bias for MSA selection +// local window always wins a slot +class llm_graph_input_msa_local : public llm_graph_input_i { +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 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(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 rank(nblk); + std::vector valid(nblk); + std::vector 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(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(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 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(&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); +} diff --git a/src/models/models.h b/src/models/models.h index 76daa8cc1..916459e12 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -1902,6 +1902,29 @@ struct llama_model_minimax_m2 : public llama_model_base { std::unique_ptr 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 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) {} diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index 86c3051c5..d02e65c9e 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -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({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: