MiniMaxM3SparseForConditionalGeneration model_type · minimax_m3_vl MSA 块级稀疏注意力 GQA 64 : 4 原生多模态

MiniMax-M3 架构解剖

M3 是一个 VL 外壳套一个稀疏 MoE 文本骨架config.json 顶层只有 text_configvision_config 两段,权重也干净地分成 language_model.*(22893 个张量)与 vision_tower.* / projector(523 个)。 真正的新东西是文本侧的 MSA:注意力仍是普通 GQA(64 头 : 4 KV 头,head_dim 128), 但每个 query token 先由一个 4 头的轻量 indexer 给每 128 个 token 一块打分, 只取 16 块 = 2048 个 KV token 进真正的注意力。 索引的 K 只有一个头、单独存一份 page-128 cache, 整套东西在 ATOM 里落在 atom/models/minimax_m3.pyatom/model_ops/minimax_m3/ 两处。

427.0 B总参数(含视觉塔)
24.7 B单 token 激活
3 + 57稠密层 + MoE 层
57 / 60稀疏注意力层
6144hidden_size
128 / 4专家数 / 每 token
16 × 128top-k 块 → 2048 token
524 288max_position(对外 1 M)

定位

不是 MLA + DSA 那一路,是 GQA + 块级 top-k

把 M3 和 GLM-5.2 / 5.3 那一系放在一起最容易读懂它:两边都在做「先索引、再稀疏注意力」, 但三个关键选择完全不同。

M3 的三个选择

  • 注意力本体是 GQA,不是 MLA。q/k/v 三个独立投影,KV cache 存的就是 4 个头 × 128 维的真实 K 和 V,没有低秩压缩、没有解压 GEMM。
  • top-k 的粒度是 128-token 块,不是单 token。 sparse_block_size = 128 被硬绑在 KV page size 上, 一个稀疏块正好是一个 page,选中的块号可以直接当块表用。
  • 索引器只打分,不带 value。sparse_disable_index_value 对所有稀疏层置 1,indexer 没有 v 投影也没有 per-head 权重, score 就是 max over 块内 128 个位置的 q·kᵀ/√128

落到 ATOM 的位置

  • atom/models/minimax_m3.py — 925 行,独立模型文件(不像 GLM 挂在 deepseek_v2.py 上)。
  • atom/model_ops/minimax_m3/index_topk.py — 1155 行 Triton:打分 + bitonic top-k + 顺手吐出 page-16 块表。
  • atom/model_ops/minimax_m3/sparse_attn.py — 1423 行:block-sparse GQA 前向。
  • SparseMHAPagedAttentionImplattention_mha.py)只覆写了 rope_cachedispatch_backend 两个钩子, KV cache 的分配 / 绑定完全走标准 MHA 路径。

视觉塔目前不在 ATOM 里

文件末尾一行 MiniMaxM3SparseForConditionalGeneration = MiniMaxM3SparseForConditionalGenerationTextOnly:架构名存在,但走的是文本视图, vision_tower. / multi_modal_projector. / patch_merge_mlp. 三个前缀在 skip_weight_prefixes 里被直接跳过。所以 recipe 里的 --language-model-only 不是可选项。


规模

427.0 B 落盘,单 token 只碰 24.7 B

下面每一行都是从 59 个 safetensors 分片的头部读出来的真实形状累加,不是从 config 推的。 专家权重吃掉 96.74 %,剩下所有东西加起来不到 14 B。

routed_experts57 层 × 128 × 3 × [3072, 6144]
413.122 B96.74 %
attention60 层 × (q 8192 + k/v 512 + o)
6.417 B1.50 %
shared_expert57 层 × 3 × [3072, 6144]
3.228 B0.76 %
lm_head[200064, 6144]
1.229 B0.29 %
embed_tokens[200064, 6144]
1.229 B0.29 %
dense_mlp3 层 × 3 × [12288, 6144]
0.679 B0.16 %
vision_tower32 层 × 1280
0.631 B0.15 %
projector + merge1280→6144, 24576→6144
0.234 B0.05 %
sparse_indexer57 层 × (512 + 128) × 6144
0.224 B0.05 %
router + normsgate [128,6144] fp32 等
0.046 B0.01 %
合计 427.040 B 落盘 854.18 GB(bf16) 条长以 routed_experts 为满格

激活量是怎么凑出来的

  • 每个 MoE 层:注意力 106.95 M + indexer 3.93 M + 4 个专家 226.49 M + shared 56.62 M + router 0.79 M ≈ 394.8 M
  • × 57 层 = 22.50 B;3 个稠密层 × (106.95 M + 226.49 M) = 1.00 B; lm_head 1.229 B。
  • 合计 24.73 B(不含 embedding 查表)。稀疏比 24.73 / 427.04 ≈ 1 : 17.3

一个对不上的数

model.safetensors.index.jsonmetadata.total_size 写的是 869.16 GB,而 59 个分片 du --apparent-size 求和是 854.18 GB,与把所有 header 里的形状按 dtype 累加得到的 854.17 GB 一致。多出来的 15 GB 不对应任何一个实际张量,算显存时不要用 metadata 那个数。


结构全图

从 input_ids 到 logits,再放大其中一层

右侧展开的是 layer 3–59 中任意一层。值得单独看一眼的是那次 [6144 → 9856] 的打包 GEMM:q(8192) + k(512) + v(512) + index_q(512) + index_k(128) 五路一次算完,磁盘上它们是 5 个独立张量,靠 packed_modules_mapping 在加载时拼进 qkv_proj。 即使是复用上一层 top-k 的那 42 层,这次 GEMM 也照算不误——省掉的只是 indexer 的 norm / rope / 打分 / top-k。

图 1 · 结构全图

input_ids [T] embed_tokens [200064, 6144] layers 0 – 2 · 稠密 全注意力 + MLP 12288 无 indexer 权重 layers 3 – 59 · MoE + MSA 稀疏注意力 + 128 专家 57 层,结构逐层相同 model.norm · GemmaRMSNorm lm_head [200064, 6144] logits [T, 200064] 残差流全程 [T, 6144] tie_word_embeddings = false lm_head 与 embed 各占 1.229 B 一层 MoE + MSA 解码层展开(layer 3 – 59) input_layernorm · GemmaRMSNorm(6144) qkv_proj(打包,单次 GEMM) [6144 → 9856] 按列切 5 段 q [T, 64, 128] 8192 k [T, 4, 128] 512 v [T, 4, 128] 512 index_q [T, 4, 128] 512 index_k [T, 1, 128] 128 q_norm / k_norm · GemmaRMSNorm(128) 逐头 RoPE 部分旋转 rotary_dim = 64 / 128 · θ = 5e6 index_q_norm / index_k_norm(128) 与主路共享同一份 cos/sin 写 KV cache · page-16 SHUFFLE K [8N,4,hd/x,16,x] · V [8N,4,16/x,hd,x] 写 index cache · page-128 [N, 128, 128] 单头,4 个索引头共享 打分 q·kᵀ / √128 每 128 token 一块取 max score [4, T, ⌈S/128⌉] top-k = 16 块(bitonic) topk_idx [4, T, 16] → 块表 [T, 128] block-sparse paged attention(AITER ASM / gluon) 每个 q 头只读 16 × 128 = 2048 个 KV token out [T, 64, 128] q 不入 cache KV 读回 选中块下标 o_proj [8192 → 6144] RowParallel post_attention_layernorm(6144) block_sparse_moe gate [6144 → 128] fp32 权重 · sigmoid + e_score_correction_bias [128] fp32 top-4 / 128 renormalize × routed_scaling 2.0 选中的 4 个专家 w1, w3 [3072, 6144] · w2 [6144, 3072] shared_experts × 1 always-on swiglu_oai α=1.702 · limit=7.0 hidden [T, 6144] → 下一层残差
残差加法与 all-reduce 没有画出来:ATOM 把「上一层的残差 + all-reduce + GemmaRMSNorm(+ 可选 fp8 量化)」 融成一个算子 fused_allreduce_gemma_rms_norm[_quant], 所以图中每个 layernorm 方框实际上都同时承担了 TP 归约与残差累加。

核心机制 · MSA

4 个索引头,一份共享的索引 K

索引器的形状很值得单独念一遍:index_q_proj [512, 6144]4 个头 × 128 维,而 index_k_proj [128, 6144] 只有 1 个头。打分内核里 pid_h 遍历 4 个索引头, 但每次都从同一个 page 载入同一份 K——索引是 MQA,主注意力是 GQA。 4 个索引头正好对上 4 个 KV 头,于是每个 KV 头组(16 个 q 头)拿到属于自己的那 16 个块。

图 2 · MSA 三步

① 打分 index score index_q [T,4,128] × index cache [N,128,128] q·kᵀ / √128,每 128-token 块内取 max → score [4, T, ⌈S/128⌉] fp32 ② 选块 top-k bitonic top-k,k = 16(含 1 个强制 local 块) score 被改写为 1e29 以强制入选 → topk_idx [4, T, 16] int32 ③ 取 KV block-sparse PA 每个 128-块展开成 8 个物理 16-page sparse block table [T, 16 × 8 = 128] → 每个 q 头只读 2048 个 token 4 个索引头 块下标 同样的 2048 个 token,在不同上下文长度下是完全不同的一件事 S = 2048 100 % 16 块 S = 8 192 25 % 16 块 S = 65 536 3.125 % 16 块 S = 262 144 0.781 % 16 块 S = 1 048 576 (1M) 0.195 % 16 块 上下文不足 2048 时 top-k 覆盖全部块,MSA 退化为全注意力;真正省下来的算力只在长上下文出现。
sparse_init_block = 0sparse_local_block = 1:M3 不强制保留开头的块, 只强制保留当前 token 所在的那一块。强制的实现方式是在 top-k 之前把该块的 score 直接改写成 1e29(init 用 1e30),而不是单独留出名额——所以 16 是强制块的总数。

打分内核的两个细节

  • 缩放用的是主注意力的 sm_scale = 128-1/2,并且预乘了 log2(e)——整条链路是 base-2 softmax。
  • 因果掩码只在「当前 q 块可能落进这个 K 块」时才施加,其余块整块跳过掩码计算。 块内被掩掉的位置取 -inf,不会污染 max

为什么 block-size 必须是 128

index cache 布局是 [num_blocks, 128, 128],第二维就是 page 内的 token 位置。 打分内核用 block_table[blk] 直接换算 page 号, top-k 出来的块号也直接当块表用。--block-size 128 一改,这三处同时失配。

图 3 · 层调度

FFN
attention
index top-k
0
10
20
30
40
50
59
稠密 MLP(12288) MoE(128 专家 / top-4) 全注意力 MSA 稀疏注意力 真的算 top-k(15 层) 复用上一层的 top-k(42 层)
三条 *_freq 数组在 config 里是同一个形状:moe_layer_freqsparse_attention_freqsparse_disable_index_value 都是 [0,0,0,1,1,…,1]——前 3 层稠密且全注意力,后 57 层既是 MoE 也是稀疏层, 三者边界完全重合。第三行不是 checkpoint 的性质而是 ATOM 的运行期调度: use_index_cache=trueindex_topk_freq=4 时, 按稀疏层序号(不是绝对层号)每 4 层只有第 1 层跑 indexer 的 norm/rope/打分/top-k, 其余 3 层直接复用缓存下来的块下标。57 个稀疏层因此只剩 15 层真正算 top-k。

MoE

128 个专家里挑 4 个,外加一个永远在线的 shared

路由是 sigmoid 打分 + 修正偏置那一套:gate 权重在 checkpoint 里是 fp32[128, 6144]),e_score_correction_bias 也是 fp32。选出 top-4 后 renormalize,再整体乘 routed_scaling_factor = 2.0,最后加上 shared expert 的输出。

字段形状 / 说明
num_local_experts128w1 / w3 [3072, 6144],w2 [6144, 3072]
num_experts_per_tok4每 token 激活 4 × 56.62 M = 226.49 M
n_shared_experts1shared_intermediate_size 3072,与单个路由专家等宽
intermediate_size3072路由专家;dense_intermediate_size 12288 只用于 layer 0–2
scoring_funcsigmoiduse_routing_bias = true
routed_scaling_factor2.0只乘路由输出,shared 不乘
hidden_actswigluoaigate·σ(1.702·gate)·(up+1),gate 与 up 各自 clamp 到 ±7.0

ATOM 里一处刻意的反优化

MiniMaxM3MoE.__init__ 把后端的 intermediate_pad 强行置 0, 注释写得很直白:专家权重在加载时已经补齐,运行期再走「跳过 pad」的快路径会掉精度, 所以宁可整块算完 padded intermediate。看 MoE 性能时这一行值得记住。


归一化与位置编码

Gemma 式 RMSNorm,逐头 QK-norm,一半维度做 RoPE

配置含义
归一化use_gemma_norm = trueATOM 用 GemmaRMSNormx·(1+w) 而不是 x·w,且乘完再回 bf16
QK-normuse_qk_norm · per_headq_norm / k_norm 权重只有 [128],对每个头独立归一化,60 层全有
RoPEpartial_rotary_factor = 0.5rotary_dim = 64,head_dim 的后 64 维不旋转;rope_theta = 5e6
索引路 RoPE共享主路index_rotary_emb = rotary_emb,连 cos/sin 缓存都是同一份拼好的张量
attention 门控attention_output_gate = falseM2 系有的输出门在 M3 上关掉了
max_position524 288config 里是 512 K;对外宣称的 1 M 上下文靠推理侧外推

这一整套(qk-norm → 部分 RoPE → 写 KV → 写 index K)在 ATOM 里是一个 aiter 算子aiter.fused_qknorm_idxrqknorm,直接吃打包好的 qkv 张量, 一次写完 SHUFFLE KV cache、index cache 和 fp8 的 per-token dequant scale。 只有复用 top-k 的那 42 层会退回到 Triton 版 triton_fused_norm_rope_cache(不带索引路)。


显存

缓存账:主 KV 之外还要单独养一份索引 K

图 4 · 两套 KV 布局

主 KV cache(由标准 MHA 路径分配,与其它 MHA 模型完全相同) page-128 SHUFFLE K [N, 4, hd/x, 128, x] V [N, 4, 128/x, hd, x] k_scale / v_scale [N, 4, 128](仅 fp8) page-16 SHUFFLE(ASM / gluon 内核实际索引方式) K [8N, 4, hd/x, 16, x] V [8N, 4, 16/x, hd, x] k_scale / v_scale [8N, 4, 16] 128 = 8 × 16 → 纯 view 重解释,零拷贝 index cache · page-128 [N, 128, 128] 单头 4 个索引头共读同一份 K block_size 必须 = 128 一个稀疏块 ⇔ 一个逻辑 page top-k 下标可直接当块表用 每 token 的缓存账(bf16) 主 KV · 60 层 × 2048 B 120.00 KiB index K · 57 层 × 256 B 14.25 KiB 合计 134.25 KiB / token → 1 M 上下文单请求 134 GiB;kv-cache-dtype fp8 后主 KV 减半,合计 74.25 KiB / token → 74 GiB。
两套布局是同一块显存。ATOM 让 metadata builder 按普通 MHA 分配 page-128 SHUFFLE, 到了注意力里再 view 成 page-16——这样绑定路径上没有任何 M3 专用代码, 代价是读代码时要记住同一个指针有两种形状。
缓存每 token 每层层数每 token1 M 上下文
主 KV(bf16)2 × 4 × 128 × 2 B = 2048 B60120.00 KiB120 GiB
index K(bf16)1 × 128 × 2 B = 256 B5714.25 KiB14 GiB
合计134.25 KiB134 GiB
主 KV 转 fp8 后1024 B6074.25 KiB74 GiB

index cache 的 dtype 由 atom_config.index_cache_dtype 单独控制,默认跟随 KV; 打分内核里有 k.dtype.is_fp8() 分支,索引路可以独立降到 fp8。


多模态

ViT 出来的 token 先被砍掉 3/4 才进语言模型

图 5 · 视觉前端

输入 tile 336 × 336 × 3 RGB, temporal 2 patch_embedding Conv3d(3→1280) kernel 2 × 14 × 14 [576, 1280] 24 × 24 个 patch vision_tower × 32 层 1280 · 16 头 × 80 · MLP 5120 gelu · 3D RoPE multi_modal_projector 1280 → 6144 → 6144 gelu, 带 bias [576, 6144] 每 patch 一个 token 2 × 2 空间合并 concat 4 → [144, 24576] patch_merge_mlp 24576 → 6144 → 6144 [144, 6144] 每 tile 144 个视觉 token 插入文本序列 image_token_index = 200025 · video_token_index = 200026 dynamic_res:最大 2016 × 2016 → 20736 patch → 5184 token image_grid_pinpoints 共 36 档,每档都是 336 的整数倍
顺序容易记反:先投影到 6144,再做 2×2 合并。 证据在形状里——patch_merge_mlp.linear_1.weight[6144, 24576],而 24576 = 4 × 6144,不是 4 × 1280。 patch_embedding 是 Conv3d [1280, 3, 2, 14, 14], 那个 2 是 temporal_patch_size:静态图按重复两帧处理,视频天然按 2 帧一组。

ATOM 后端

跑起来要注意的几件事

不能改的参数

  • --block-size 128:一个稀疏块 ⇔ 一个 page,改了直接失配。
  • --language-model-only:视觉塔未移植。
  • --no-trust-remote-code:ATOM 自己注册模型类, 用 HF 的 auto_map 反而会走错实现。
  • head_dim 固定 128:sparse 路径写死走 AITER 融合算子,没有 Triton 兜底 (use_triton_attn = False)。

可以调的开关

  • --hf-overrides '{"use_index_cache": true, "index_topk_freq": 4}': top-k 复用的开关与步长,见图 3。
  • index_topk_pattern:给一串 "S" / 非 S 逐层指定, 优先级高于 index_topk_freq
  • --kv-cache-dtype fp8:主 KV 减半,见缓存账。
  • --compilation-config '{"cudagraph_mode": "FULL_AND_PIECEWISE"}'
checkpoint 变体量化要点
MiniMax-M3bf16854.18 GB,59 分片
MiniMax-M3-MXFP4fp4 per_group 32 · e8m0 scale权重与激活都 MXFP4,观测器 PerBlockMXObserver
MiniMax-M3-MXFP8mxfp8 · weight_block [1, 32]动态激活量化;ignored_layers 里排除了 lm_head、embed、视觉塔与部分层的 block_sparse_moe.gate
MiniMax-M3-MXFP4-AttnFP8MXFP4 + 注意力 fp8MXFP4 基础上把注意力线性层单独提到 fp8
MiniMax-M3-EAGLE3bf16草稿模型,见下节

MXFP8 走的是运行期在线量化:recipe 里那串 --additional-config '{"online_quant_config": {"global_quant_config": "ptpc_fp8", …}}', 排除表和 checkpoint 里的 ignored_layers 是同一批名字。

配置GSM8K flexible-extractstrict-match
MXFP80.95030.9510
MXFP40.93990.9407
MXFP8 + kv fp80.94800.9487
MXFP4 + kv fp80.94390.9445

5-shot、5 次本地取平均,来自 recipes/atom_vllm/MiniMax-M3.md


投机解码

config 写了 MTP,盘上没有 MTP

text_config 里明明白白有 num_mtp_modules = 7num_nextn_predict_layers = 1,但 23416 个张量里没有一个属于 MTP—— ATOM 的模型文件里也没有 MTP 分支。M3 在 ATOM 上的投机解码走的是 Eagle3

aux hidden state 的取法

get_eagle3_aux_hidden_state_layers() 返回 (2, 30, 57)(early / mid / late)。取的是进入该层、过完 input_layernorm 之后的 residual,不是常见的 hidden_states + residual—— 因为 M3 的融合 all-reduce RMSNorm 会让那个和在 CUDA graph 下处于 TP 部分和状态甚至出 NaN。

草稿模型

  • LlamaForCausalLMEagle31 层
  • hidden 6144 · intermediate 18432 · 64 头 / 64 KV 头(纯 MHA,不做 GQA)。
  • vocab 200064 与主模型一致,rope_theta 同为 5e6, max_position_embeddings 直接写 1 048 576。

读码路线

按这个顺序看

#文件看什么
1config.json三条 *_freq 数组的边界是否重合(M3 是重合的)
2atom/models/minimax_m3.py_sparse_attention_layer_ordinals_should_skip_minimax_m3_index_topk:调度按稀疏层序号而非绝对层号
3atom/models/minimax_m3.pyMiniMaxM3SparseAttention.forward:5 路 split 的宽度就是那次打包 GEMM 的全部内容
4atom/model_ops/attention_mha.pySparseMHAPagedAttentionImpl.rope_cache / _to_page16_shuffle:两套布局在这里对上
5minimax_m3/index_topk.py_index_block_score_kernel:4 个索引头共读一份 K;_topk_index_kernel:强制块靠改写 score 实现
6minimax_m3/sparse_attn.pypage-16 块表如何喂给 gluon PA
recipes/atom_vllm/MiniMax-M3.md      # 起服务的完整命令 + GSM8K 基线
recipes/atom_sglang/MiniMax-M3.md    # SGLang 后端
recipes/mesh/MiniMax-M3.md           # mesh 部署
/shared/data/amd_int/models/MiniMax-M3*   # bf16 与 4 个量化变体