Skip to content

feat(models): add MiniMax-M3 text model (hybrid dense/MoE with block-sparse attention) #763

Description

@inureyes

Summary

Add support for the MiniMax-M3 text architecture (config model_type: "minimax_m3"). The existing minimax family in mlxcel is the MiniMax-M2 architecture (src/models/minimax.rs, 256 experts, top-8) and cannot load M3 checkpoints: M3 uses a different layer plan, router, activation, and a block-sparse attention mechanism, and src/models/detection.rs has no minimax_m3 arm.

Architecture specification

Decoder (60 layers, hidden_size 6144):

  • First 3 layers use a dense MLP (dense_intermediate_size 12288); the remaining layers are MoE, driven by the moe_layer_freq / mlp_layer_types config keys.
  • MoE: num_local_experts 128, num_experts_per_tok 4, sigmoid scoring with e_score_correction_bias (selection-only bias, config use_routing_bias), routed_scaling_factor 2.0, one shared expert (n_shared_experts 1, shared_intermediate_size 3072). When the shared expert width equals the routed expert width, checkpoints pack it as expert index 128 inside the switch tensors and it participates with a fixed score of 1.0; otherwise it is a separate MLP.
  • Activation: clamp-SwiGLU (hidden_act: "swigluoai"): x_glu * sigmoid(alpha * x_glu) * (x_linear + beta) with alpha 1.702, limit 7.0, beta 1.0. This is the same formula as gpt_oss_swiglu (src/models/gpt_oss.rs:397) and mlxcel_core::compiled_gpt_oss_swiglu_activation; reuse it.
  • Norms: Gemma-style RMSNorm (weight + 1) on all norms, including per-head Q/K norm (use_qk_norm, qk_norm_type: "per_head").
  • Attention: GQA with 64 query heads, 4 KV heads, head_dim 128, partial RoPE (partial_rotary_factor 0.5, rotary dim 64), rope_theta 5e6.
  • Sparse attention ("MSA"): a lightning-indexer variant that selects top-k blocks instead of top-k tokens. Config block sparse_attention_config with sparse_block_size 128, sparse_topk_blocks 16, sparse_num_index_heads 4, sparse_index_dim 128, sparse_init_block 1, sparse_local_block 1, score type "max". Separate index_q_proj / index_k_proj projections with Gemma RMSNorm on index Q/K and RoPE applied to the index projections; blocks are scored, the top-k blocks are selected, and a block mask is applied to the dense attention. The indexer needs a side cache for index keys next to the regular KV cache. Decode can stay dense while the selected sparse window covers at least half the cache; only switch to the sparse mask beyond that.
  • The config carries MTP metadata, but no MTP head needs to be implemented in this issue.

The existing DSA indexer shared by deepseek_v32 and glm_moe_dsa (src/models/deepseek_v32_indexer.rs) is the closest template: index projections, index norms, RoPE-on-index, and top-k selection to a mask are structurally the same; the deltas are block-level scoring/selection (instead of token-level) and the block mask construction.

Implementation plan

  1. New src/models/minimax_m3.rs with a serde config struct covering the keys above; from_weights reusing SwitchGLU/SwitchLinear (src/models/switch_layers.rs), the gpt_oss clamp-SwiGLU, Gemma RMSNorm helpers, and the M2 router pattern from src/models/minimax.rs (sigmoid plus bias for selection only).
  2. Implement the block-sparse indexer following deepseek_v32_indexer.rs, with an index-key side cache modeled on the existing indexer cache handling.
  3. Sanitize: handle checkpoint prefix rewrites (model.language_model. to language_model. and similar) plus packed shared-expert stacking; keep bf16 scales and biases untouched for quantized exports.
  4. Register: src/models/detection.rs ("minimax_m3"), src/model_metadata.rs, the exhaustive-match sites (model_metadata registration macro, TP arch string, generate summary), and docs/supported-models.md.
  5. Unit tests beside the code (minimax_m3_tests.rs): router selection with bias, packed shared expert, partial RoPE, per-head QK norm, and an indexer degeneration test (with top-k covering every block, output must match dense attention).

Acceptance criteria

  • mlxcel list reports the family; a real MiniMax-M3 checkpoint loads and generates coherent text via mlxcel generate.
  • Dense vs sparse-mask degeneration parity test passes.
  • Server chat completion works, including the batched decode path.
  • cargo clippy --all-targets -- -D warnings and cargo fmt --all -- --check are clean.

Validation

Checkpoints are published under the MiniMaxAI HuggingFace org (MiniMaxAI/MiniMax-M3, plus an MXFP8 export). The model is large; if no quantized conversion fits the test machine, validate with a reduced-layer synthetic config plus a real-weights load smoke test, and state the memory constraint in the PR.

Effort: high.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions