feat(models): add MiniMax-M3 text model (hybrid dense/MoE with block-sparse attention) - #799
Conversation
…se attention) Add the MiniMax-M3 text decoder (config model_type "minimax_m3"). Unlike the existing MiniMax-M2 family (src/models/minimax.rs, 256 experts, standard RMSNorm, dense attention), M3 is a hybrid dense/MoE decoder with Gemma-style RMSNorm (weight+1) everywhere including per-head Q/K norm, clamp-SwiGLU (swigluoai) experts and dense MLPs, partial RoPE, a sigmoid router with a selection-only bias plus one shared expert, and a block-sparse "MSA" attention indexer on the sparse layers. New module (src/models/minimax_m3*.rs), split to stay under the file-size limit: config (ModelArgs, SparseAttentionConfig) with serde defaults keyed to the real checkpoint so a VL wrapper can later construct it from a nested text_config; attention with per-head Gemma Q/K norm and partial RoPE; the clamp-SwiGLU experts reusing mlxcel_core::compiled_gpt_oss_swiglu_activation and switch_layers::SwitchLinear; the sigmoid MoE router where the routing bias steers selection only while mixture weights stay the unbiased sigmoids, with the shared expert packed as switch-tensor index num_local_experts (fixed score 1.0) when its width equals the routed width, else a separate MLP; and the block-sparse indexer (minimax_m3_indexer.rs) that scores top-k blocks of keys, expands the choice to an additive block mask over the causal mask, and caches its index key alongside the regular K. Layer plan is driven by moe_layer_freq (leading dense layers, then MoE) and sparse_attention_config.sparse_attention_freq (dense attention on the leading layers). The weight sanitizer normalizes language_model. prefixes to the flat model. layout and drops MTP/next-N tensors so loading never fails on the MTP metadata the config carries but this decoder does not implement. Decode stays dense while the selected window still covers at least half the live cache; when sparse_topk_blocks covers every block the additive mask is all zeros and attention is bit-identical to dense. Registered across detection, model_metadata, the LoadedModel dispatch, the tensor-parallel arch string, ModelType/ALL_MODEL_TYPES/metadata, and docs/supported-models.md. Validation: the only public checkpoint is the 427B MiniMaxAI/MiniMax-M3 VL repo, which exceeds this machine's 128 GB, so validation is synthetic reduced-config unit tests plus a real-text_config parse test (6 tests, all passing): config parse and derived layer plan, router bias-affects-selection-only, packed shared expert (index 128, score 1.0), partial RoPE leaving the non-rotary tail untouched, per-head Q/K norm independence, and the indexer degeneration (full block coverage matches dense). Real-checkpoint load/generate and server chat completion carry over to a follow-up runtime validation when a fitting quantized conversion exists. Refs #763
Implementation Review SummaryIntentAdd the MiniMax-M3 text decoder ( What is correct and validatedThe checkpoint-free surface is sound and matches project conventions:
Findings (all are risks against the deferred "real checkpoint loads/generates" criterion, not defects in the validated surface)Two independent reviews converge on the checkpoint weight-layout assumptions as the top risk, and each reaches a different conclusion about the actual tensor layout. Neither can load the 427B checkpoint here, so these items must be reconciled against the real weight map (the safetensors index) before the load path can be considered done.
DispositionNo auto-fix was applied. Every finding targets the real-checkpoint load path, which cannot be exercised here (the only public checkpoint is the 427B VL repo and the network is unavailable), and the two expert reviews propose different "correct" layouts. Editing the deliberate, config-informed loading logic to match an unverified layout would trade one unverifiable assumption for another with no test to adjudicate, and risks regressing the passing synthetic surface. The recommended next step is to fetch the checkpoint's safetensors weight map, reconcile the three HIGH items, add defensive load-time shape/row-count checks, then run the deferred real-weights load/generate validation. Verification
|
…layout
Ground-truth tensor names and shapes from the MiniMaxAI/MiniMax-M3 safetensors index (verbatim, layer 3) refuted three assumptions in the first pass; the loader could not load the actual checkpoint.
MoE layout: the real experts live under block_sparse_moe.experts.{0..127}.w1/w2/w3.weight (Mixtral naming: w1=gate_proj, w3=up_proj, w2=down_proj) with router block_sparse_moe.gate.weight and bias block_sparse_moe.e_score_correction_bias, not the mlp.switch_mlp/mlp.experts paths the first pass expected. The decoder now loads MoE layers from block_sparse_moe (dense layers keep plain mlp), and the experts are stacked through the shared switch_layers loader with detected proj names (Mixtral w1/w3/w2, falling back to gate_proj/up_proj/down_proj).
Shared expert: the width-equality packing heuristic was wrong. The checkpoint always stores the shared expert as a separate block_sparse_moe.shared_experts.{gate_proj,up_proj,down_proj} MLP and the routed experts are exactly 0..127. The packed path and its append-shared-expert helper are removed; the shared expert is now loaded when its tensors are present.
Indexer: index_k_proj is [128, 6144], a single shared index-key head (MQA), while index_q_proj [512, 6144] is 4 query heads x 128. The indexer now builds one key stream that all query heads score against, caches the single-head index key on the regular K buffer via a head-axis concat (guarded by sparse_index_dim == head_dim), and drops the old sparse_num_index_heads == num_key_value_heads constraint. index_q_norm/index_k_norm are [128] over the index vectors.
Robustness: per-head Q/K norm confirmed to use a single [head_dim] gamma shared across heads; the loader tolerates absent index_* tensors on the dense-attention layers (driven by sparse_attention_freq) and the sanitizer now drops vision_tower./multi_modal_projector./patch_merge_mlp. tensors for a text-only load of the VL checkpoint. Load-time output-dim checks on q/k/v and index_q/index_k projections fail with a clear error on a layout mismatch instead of corrupting the forward.
Tests updated to the real layout (9 passing): sanitizer prefix rewrite plus vision/MTP drop on verbatim keys, block_sparse_moe Mixtral (w1/w2/w3) load with a separate shared expert, MQA indexer shapes (4 query heads, 1 key head), dense-layer index absence, plus the retained router-bias, partial-RoPE, per-head-QK-norm, and indexer-degeneration tests.
Refs #763
Re-review: HIGH findings resolved (commit aa2304a)Verified the delta against the real checkpoint layout. All three HIGH items from the previous review are fixed, and I re-ran
Also verified: One earlier MEDIUM is now moot: with a single ground-truth index-key head, the block mask is architecturally a single stream, so the previous per-KV-group-selection concern does not apply. Remaining (non-blocking, unchanged by this commit, for runtime validation)
Verdict: no remaining CRITICAL or HIGH. The load path now matches the verified checkpoint layout and the reduced-config load/forward paths are exercised by tests; a real-weights load/generate run remains the final validation step. |
sparse_block_size == 0 divides by zero computing num_blocks in build_block_drop_mask and always trips should_apply_sparse's window check; sparse_topk_blocks == 0 reaches argpartition with kth = -1, an invalid partition index. Add SparseAttentionConfig::validate(), following the validate_quantization_scheme precedent in gemma4.rs, and call it once in BlockSparseIndexer::load rather than guarding the masking hot path on every call. Also correct the docs/supported-models.md MiniMax-M3 entry, which still described the pre-aa2304a packed shared-expert heuristic and a per-head index_k_proj; it now reflects the separate shared_experts MLP and the MQA (single shared key head) indexer. Validation: - cargo fmt --all -- --check passes - cargo test --release --features cuda --lib models::minimax_m3 -- --test-threads=1: 10 passed, 0 failed (9 existing plus the new rejection test) Refs #763
Summary
Adds the MiniMax-M3 text architecture (config
model_type: "minimax_m3"): a hybrid dense/MoE decoder with Gemma-style RMSNorm (weight+1) everywhere, clamp-SwiGLU (swigluoai) experts and dense MLPs, partial RoPE, a sigmoid router with a selection-only routing bias plus one shared expert, and a block-sparse "MSA" attention indexer on the sparse layers. This is the text decoder only; the VL wrapper (#764) and the MTP head are out of scope and build on this work. The loader was reconciled against verbatim tensor names and shapes from the realMiniMaxAI/MiniMax-M3safetensors index.Architecture implemented
Layer plan: the per-layer dense/MoE split is driven by
moe_layer_freq(the real checkpoint runs 3 leading dense layers, then MoE), and the per-layer sparse-attention gate bysparse_attention_config.sparse_attention_freq(dense attention aligns with the leading dense-MLP layers). Dense layers load a plainmlp; MoE layers loadblock_sparse_moe.Attention: GQA (64 query heads, 4 KV heads, head_dim 128; q_proj
[8192, 6144], k/v_proj[512, 6144]) with per-head Gemma Q/K RMSNorm (qk_norm_type: "per_head", a single[128]gamma shared across heads) and partial RoPE (rotary_dim64 of head_dim 128,rope_theta5e6). Load-time output-dim checks on q/k/v fail with a clear error on a GQA layout mismatch.Router and experts (
block_sparse_moe): sigmoid scoring where the routing bias affects the top-k selection only and the returned mixture weights are the unbiased sigmoids (normalized, then scaled byrouted_scaling_factor2.0). Routed experts are stored under the Mixtral convention (block_sparse_moe.experts.{0..127}.w1/w2/w3.weight, w1=gate_proj, w3=up_proj, w2=down_proj) and stacked through the sharedswitch_layersloader with detected proj names (falling back togate_proj/up_proj/down_proj). The single shared expert is a SEPARATE MLP (block_sparse_moe.shared_experts.{gate_proj,up_proj,down_proj}), loaded when present and added to the routed mixture; it is never packed into the switch tensors. The clamp-SwiGLU activation reusesmlxcel_core::compiled_gpt_oss_swiglu_activation, and the experts reuseswitch_layers::SwitchLinear.Block-sparse indexer (
minimax_m3_indexer.rs, MQA, modeled on the DSA lightning indexer):index_q_proj [512, 6144]produces 4 query heads ofsparse_index_dim128;index_k_proj [128, 6144]produces a single shared index-key stream that all query heads score against;index_q_norm/index_k_normare[128]Gemma norms over the index vectors, with RoPE on both. The single-head index key is cached alongside the regular K via a head-axis concat (guarded bysparse_index_dim == head_dim). Token scores are reduced to block scores by max oversparse_block_size, top-sparse_topk_blocksblocks are selected with forced initial and per-query local blocks, and the choice expands to an additive block mask on top of the causal mask. Decode stays dense while the selected window still covers at least half the live cache; only beyond that does the block mask apply. Load-time output-dim checks onindex_q_proj/index_k_projcatch an MHA-vs-MQA layout mismatch.Deviations from the issue spec, all reconciled against the real checkpoint: there is no standalone text-only repo, so the config parses a flat
text_configshape and is factored so the VL wrapper can construct it from the nested block;sparse_init_blockis 0 (the issue said 1); the layer plan is driven bymoe_layer_freq. The indexer is MQA (single shared index-key head), so the oldsparse_num_index_heads == num_key_value_headsconstraint is gone; the side cache instead requiressparse_index_dim == head_dimand disables (dense fallback) otherwise or when index weights are absent (dense-attention layers).What changed
src/models/minimax_m3.rs(model, decoder layer, load/from_weights, weight sanitizer,LanguageModel),minimax_m3_config.rs(serdeModelArgs/SparseAttentionConfigwith checkpoint-keyed defaults),minimax_m3_layers.rs(attention with head-axis index-key side cache and q/k/v shape checks, clamp-SwiGLU, dense MLP),minimax_m3_moe.rs(router,block_sparse_moeMixtral-named experts, separate shared expert),minimax_m3_indexer.rs(MQA block-sparse indexer, shape checks),minimax_m3_tests.rs(unit tests). Files split to stay under the 500-line limit.language_model.prefixes to the flatmodel.layout, drops MTP / next-N tensors, and dropsvision_tower./multi_modal_projector./patch_merge_mlp.tensors so a text-only load of the VL checkpoint ignores the vision front-end.src/models/detection.rs,src/model_metadata.rs,src/loaded_model.rs(enum + dispatch),src/distributed/tensor_parallel/inference.rs(arch string),src/models/mod.rs(ModelType,ALL_MODEL_TYPES,metadata, arch-coverage macro), anddocs/supported-models.md.Validation
The only public checkpoint is the 427B
MiniMaxAI/MiniMax-M3VL repo (top-levelmodel_type: "minimax_m3_vl"), which exceeds this machine's 128 GB unified memory, so it cannot be downloaded or loaded here. Validation is synthetic reduced-config unit tests plus a real-text_configparse test, with the tensor-layout assumptions cross-checked against the real safetensors index (verbatim tensor names and shapes). The acceptance items "real checkpoint loads/generates" and "server chat completion" carry over to a follow-up runtime validation when a fitting quantized conversion exists. The block-sparse path is realized as a masked dense attention (selecting top-k blocks via an additive mask); it reproduces the selection semantics but does not yet realize a native sparse kernel's compute/memory savings.Test plan
cargo test --release --features cuda --lib models::minimax_m3 -- --test-threads=1(9 passed): config parse and derived layer plan; sanitizer prefix rewrite plus vision/MTP drop on verbatim checkpoint keys;block_sparse_moeMixtral (w1/w2/w3) load with a separate shared expert; MQA indexer shapes (4 query heads, 1 key head); dense-layer index absence; router bias-affects-selection-only; partial RoPE leaving the non-rotary tail untouched; per-head Q/K norm independence; indexer degeneration (full block coverage is bit-identical to dense).cargo test --release --features cuda --lib -- models::metadata_tests models::detection minimax(80 passed): the arch-coverage test confirmsMiniMaxM3is wired through every registry.cargo fmt --all -- --checkclean.Closes #763