Skip to content

vllm.models.minimax_m3.amd.model

Inference-only MiniMax M3 (text backbone) model — AMD ROCm implementation.

Self-contained per-platform impl (mirrors deepseek_v4/amd). It is identical to ../nvidia/model.py except for RMS normalization: FlashInfer's Gemma RMSNorm kernels are CUDA-only, so MiniMAXGemmaRMSNorm here uses a native (FlashInfer-free) implementation.

The MiniMax-M3-preview config selects a single set of branches
  • qk_norm_type == "per_head"
  • hidden_act == "swigluoai"
  • use_gemma_norm == True -> Gemma-style RMSNorm everywhere
  • attention_output_gate == False
  • scoring_func == "sigmoid" with a routing-bias correction term
  • sparse_attention_config present -> a subset of layers run the extra "index" attention branch.

Classes:

MiniMAXGemmaRMSNorm

Bases: Module

Gemma-style RMS normalization (native ROCm implementation).

Normalizes in fp32 and scales by (1 + weight) — numerically equivalent to the FlashInfer gemma_rmsnorm / gemma_fused_add_rmsnorm kernels used in the NVIDIA path, which are unavailable on ROCm. When residual is given, the fused add + norm returns the updated (normed, residual) pair.

The fp32 normalize + scale + (optional) residual-add run in a single fused Triton pass (amd.ops.gemma_rmsnorm / gemma_fused_add_rmsnorm) instead of a chain of elementwise PyTorch kernels.

Source code in vllm/models/minimax_m3/amd/model.py
class MiniMAXGemmaRMSNorm(nn.Module):
    """Gemma-style RMS normalization (native ROCm implementation).

    Normalizes in fp32 and scales by ``(1 + weight)`` — numerically equivalent
    to the FlashInfer ``gemma_rmsnorm`` / ``gemma_fused_add_rmsnorm`` kernels
    used in the NVIDIA path, which are unavailable on ROCm. When ``residual`` is
    given, the fused add + norm returns the updated ``(normed, residual)`` pair.

    The fp32 normalize + scale + (optional) residual-add run in a single fused
    Triton pass (``amd.ops.gemma_rmsnorm`` / ``gemma_fused_add_rmsnorm``) instead
    of a chain of elementwise PyTorch kernels.
    """

    def __init__(
        self,
        hidden_size: int,
        eps: float = 1e-6,
    ) -> None:
        super().__init__()
        self.weight = nn.Parameter(torch.zeros(hidden_size))
        self.variance_epsilon = eps

    def forward(
        self,
        x: torch.Tensor,
        residual: torch.Tensor | None = None,
    ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
        if residual is None:
            return gemma_rmsnorm(x, self.weight, self.variance_epsilon)
        return gemma_fused_add_rmsnorm(x, residual, self.weight, self.variance_epsilon)

MiniMaxM3Attention

Bases: Module

Dense attention with per-head QK norm and partial RoPE.

Source code in vllm/models/minimax_m3/amd/model.py
class MiniMaxM3Attention(nn.Module):
    """Dense attention with per-head QK norm and partial RoPE."""

    def __init__(
        self,
        config: PretrainedConfig,
        layer_id: int,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
        cache_config: CacheConfig | None = None,
    ) -> None:
        super().__init__()
        self.hidden_size = config.hidden_size
        tp_size = get_tensor_model_parallel_world_size()

        self.total_num_heads = config.num_attention_heads
        assert self.total_num_heads % tp_size == 0
        self.num_heads = self.total_num_heads // tp_size
        self.total_num_kv_heads = config.num_key_value_heads
        if self.total_num_kv_heads >= tp_size:
            assert self.total_num_kv_heads % tp_size == 0
        else:
            assert tp_size % self.total_num_kv_heads == 0
        self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
        self.head_dim = config.head_dim
        self.q_size = self.num_heads * self.head_dim
        self.kv_size = self.num_kv_heads * self.head_dim
        self.scaling = self.head_dim**-0.5

        self.qkv_proj = QKVParallelLinear(
            self.hidden_size,
            self.head_dim,
            self.total_num_heads,
            self.total_num_kv_heads,
            bias=False,
            quant_config=quant_config,
            prefix=f"{prefix}.qkv_proj",
        )
        # reduce_results=False: the attention all-reduce is fused with the
        # following post_attention_layernorm (GemmaRMSNorm) in the decoder layer
        # via fused_allreduce_gemma_rms_norm.
        self.o_proj = RowParallelLinear(
            self.total_num_heads * self.head_dim,
            self.hidden_size,
            bias=False,
            reduce_results=False,
            quant_config=quant_config,
            prefix=f"{prefix}.o_proj",
        )

        # Per-head QK norm (qk_norm_type == "per_head", use_gemma_norm == True).
        self.q_norm = MiniMAXGemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps)
        self.k_norm = MiniMAXGemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps)

        # Partial RoPE: rotary_dim == head_dim * partial_rotary_factor. Honors
        # config.rope_scaling (e.g. YaRN) so long-context positions are covered.
        self.rotary_emb = _build_rotary_emb(config, self.head_dim)

        self.attn = Attention(
            self.num_heads,
            self.head_dim,
            self.scaling,
            num_kv_heads=self.num_kv_heads,
            cache_config=cache_config,
            quant_config=quant_config,
            prefix=f"{prefix}.attn",
        )

    def forward(
        self,
        positions: torch.Tensor,
        hidden_states: torch.Tensor,
    ) -> torch.Tensor:
        qkv, _ = self.qkv_proj(hidden_states)
        # Fused per-head Gemma QK-norm + partial NeoX RoPE on q/k, in place (dense
        # mode: no index branch, no KV-cache insert). Matches nvidia/model.py and
        # replaces the unfused split -> q_norm/k_norm -> rotary_emb chain; verified
        # bit-equivalent on ROCm (q/k rel ~2e-3 bf16 noise, v untouched).
        ops.fused_minimax_m3_qknorm_rope_kv_insert(
            qkv,
            self.q_norm.weight,
            self.k_norm.weight,
            self.rotary_emb.cos_sin_cache,
            positions,
            self.num_heads,
            self.num_kv_heads,
            self.rotary_emb.rotary_dim,
            self.q_norm.variance_epsilon,
            kv_cache_dtype="auto",
        )
        q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
        attn_output = self.attn(q, k, v)
        output, _ = self.o_proj(attn_output)
        return output

MiniMaxM3MLP

Bases: Module

Dense SwiGLU-OAI MLP (used by the leading dense layers).

Source code in vllm/models/minimax_m3/amd/model.py
class MiniMaxM3MLP(nn.Module):
    """Dense SwiGLU-OAI MLP (used by the leading dense layers)."""

    def __init__(
        self,
        config: PretrainedConfig,
        intermediate_size: int,
        quant_config: QuantizationConfig | None = None,
        reduce_results: bool = True,
        prefix: str = "",
    ) -> None:
        super().__init__()
        self.gate_up_proj = MergedColumnParallelLinear(
            config.hidden_size,
            [intermediate_size] * 2,
            bias=False,
            quant_config=quant_config,
            prefix=f"{prefix}.gate_up_proj",
        )
        self.down_proj = RowParallelLinear(
            intermediate_size,
            config.hidden_size,
            bias=False,
            quant_config=quant_config,
            reduce_results=reduce_results,
            prefix=f"{prefix}.down_proj",
        )
        if config.hidden_act != "swigluoai":
            raise ValueError(
                f"Unsupported activation: {config.hidden_act}. "
                "Only swigluoai is supported."
            )
        # gate * sigmoid(alpha * gate) * (up + beta), with both halves clamped.
        # Kept as our fp32 Triton kernel (not the #22 SWIGLUOAI_UNINTERLEAVE op
        # ``silu_and_mul_with_clamp``): that op IS built on ROCm but rounds
        # intermediates to bf16 (rel ~3e-3 vs our fp32 ~1e-6), which costs gsm8k
        # accuracy since this activation feeds the MXFP8 quant + MoE.
        self.swiglu_alpha = config.swiglu_alpha
        self.swiglu_beta = config.swiglu_beta
        self.swiglu_limit = config.swiglu_limit

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        gate_up, _ = self.gate_up_proj(x)
        x = swiglu_oai_split(
            gate_up,
            alpha=self.swiglu_alpha,
            beta=self.swiglu_beta,
            limit=self.swiglu_limit,
        )
        x, _ = self.down_proj(x)
        return x

MiniMaxM3MoE

Bases: Module

Sigmoid-routed MoE block with a routing-bias correction and a shared expert.

Source code in vllm/models/minimax_m3/amd/model.py
class MiniMaxM3MoE(nn.Module):
    """Sigmoid-routed MoE block with a routing-bias correction and a shared
    expert."""

    def __init__(
        self,
        config: PretrainedConfig,
        layer_id: int,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
    ) -> None:
        super().__init__()
        self.tp_size = get_tensor_model_parallel_world_size()
        if self.tp_size > config.num_local_experts:
            raise ValueError(
                f"Tensor parallel size {self.tp_size} is greater than "
                f"the number of experts {config.num_local_experts}."
            )

        self.routed_scaling_factor = getattr(config, "routed_scaling_factor", 1.0)
        self.n_shared_experts = getattr(config, "n_shared_experts", None)

        # Sigmoid routing uses a per-expert score-correction bias for selection.
        self.use_routing_bias = getattr(config, "use_routing_bias", False)
        if self.use_routing_bias:
            self.e_score_correction_bias = nn.Parameter(
                torch.empty(config.num_local_experts, dtype=torch.float32)
            )
            self.e_score_correction_bias.weight_loader = (
                MiniMaxM3MoE.ebias_weight_loader
            )
        else:
            self.e_score_correction_bias = None

        # Router weights are stored in fp32; GateLinear upcasts the bf16
        # activations and computes the gate in fp32 (fp32 router logits).
        self.gate = GateLinear(
            config.hidden_size,
            config.num_local_experts,
            bias=False,
            params_dtype=torch.float32,
            out_dtype=torch.float32,
            prefix=f"{prefix}.gate",
        )

        # Shared-expert fusion (opt-in via VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS,
        # off under expert parallelism) folds the shared expert into the routed
        # MoE call as the last expert slot, so we don't build a separate module.
        # On gfx950 with aiter MoE the append is fused inside aiter's grouped
        # top-k kernel; otherwise it goes through the vLLM top-k bias router.
        # It is disabled under expert parallelism (the shared slot is appended to
        # the routed top-k, which the EP expert-mapping path does not handle).
        # TODO: Historically, only `VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS=1`
        # is checked to enable FSE for MiniMax-M3, despite AITER not being used.
        # This should be cleaned up and use `resolve_layer_fused_shared_expert`.
        fse_requested = envs.VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS

        self.is_fused_shared_expert_enabled = False
        if (
            fse_requested
            and bool(getattr(config, "n_shared_experts", None))
            and not get_current_vllm_config().parallel_config.enable_expert_parallel
        ):
            fse_compatible, fse_reason = is_shared_expert_quant_fse_compatible(
                quant_config,
                f"{prefix}.experts",
                f"{prefix}.shared_experts",
            )
            if not fse_compatible:
                logger.warning(
                    "VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS is enabled but "
                    "cannot be enabled: %s.",
                    fse_reason,
                )
            else:
                self.is_fused_shared_expert_enabled = True

        # When additionally on gfx950 with an active aiter MoE backend, the shared
        # expert is appended inside aiter's an active aiter MoE backend, the shared
        # expert is appended inside aiter's biased grouped top-k kernel
        # (``num_fused_shared_experts``) instead of the vLLM router's torch concat.
        # Otherwise FSE still runs via the vLLM top-k bias router.
        # TODO: `on_gfx950()` check here should not be MiniMax-M3 specific, and
        # the check should be done on resolved MOE backend directly.
        self.use_aiter_moe_fse = _aiter_moe_fused_shared_experts_enabled(
            self.is_fused_shared_expert_enabled
        )

        self.shared_experts: MiniMaxM3MLP | None = None
        if self.n_shared_experts and not self.is_fused_shared_expert_enabled:
            self.shared_experts = MiniMaxM3MLP(
                config=config,
                intermediate_size=config.intermediate_size * self.n_shared_experts,
                quant_config=quant_config,
                reduce_results=False,
                prefix=f"{prefix}.shared_experts",
            )

        # The aiter MoE fused path goes through aiter's biased grouped top-k
        # (GroupedTopKRouter, as in DeepSeek-V4): M3 is not group-routed, so a
        # trivial single group (num_expert_group=topk_group=1) reduces to plain
        # top-k while applying the sigmoid + bias correction and appending the
        # always-on shared expert; aiter applies the routed scaling internally.
        # Every other path (vLLM top-k bias router, or no fusion) applies the
        # routed scaling to the MoE output here.
        self.experts = FusedMoEFactory(
            num_experts=config.num_local_experts,
            top_k=config.num_experts_per_tok,
            hidden_size=config.hidden_size,
            intermediate_size=config.intermediate_size,
            intermediate_pad=0,
            scoring_func=config.scoring_func,
            e_score_correction_bias=self.e_score_correction_bias,
            renormalize=True,
            use_grouped_topk=self.use_aiter_moe_fse,
            num_expert_group=1 if self.use_aiter_moe_fse else None,
            topk_group=1 if self.use_aiter_moe_fse else None,
            activation="swigluoai_uninterleave",
            swiglu_limit=config.swiglu_limit,
            swiglu_alpha=config.swiglu_alpha,
            swiglu_beta=config.swiglu_beta,
            routed_scaling_factor=self.routed_scaling_factor,
            apply_routed_scale_to_output=not self.use_aiter_moe_fse,
            router_logits_dtype=self.gate.out_dtype,
            shared_experts=self.shared_experts,
            n_shared_experts=(
                self.n_shared_experts if self.is_fused_shared_expert_enabled else None
            ),
            fuse_shared_experts=self.is_fused_shared_expert_enabled,
            quant_config=quant_config,
            prefix=f"{prefix}.experts",
        )

    @staticmethod
    def ebias_weight_loader(param: nn.Parameter, loaded_weight: torch.Tensor) -> None:
        assert param.size() == loaded_weight.size()
        param.data.copy_(loaded_weight.to(torch.float32))

    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
        num_tokens, hidden_dim = hidden_states.shape
        hidden_states = hidden_states.view(-1, hidden_dim)

        # router_logits: (num_tokens, n_experts); GateLinear casts to fp32.
        router_logits, _ = self.gate(hidden_states)
        final_hidden_states = self.experts(
            hidden_states=hidden_states, router_logits=router_logits
        )

        return final_hidden_states.view(num_tokens, hidden_dim)

MiniMaxM3SparseAttention

Bases: Module, AttentionLayerBase

Block-sparse attention layer with the lightning-indexer branch.

This is a merged attention layer: it owns the projections (qkv + index q/k), per-head QK norms and RoPE, and the attention-backend wiring that a generic Attention layer would normally provide — it binds the MiniMaxM3SparseBackend + main impl, registers the main paged K/V cache, and owns the lightning indexer (MiniMaxM3Indexer), which holds the index-key side cache.

The index branch (index_{q,k}proj + index_norm) feeds the sparse top-k block selection. M3 always disables the index value/output projections (sparse_disable_index_value set for every sparse layer), so index_{v,o}_proj are never created.

Source code in vllm/models/minimax_m3/amd/model.py
 582
 583
 584
 585
 586
 587
 588
 589
 590
 591
 592
 593
 594
 595
 596
 597
 598
 599
 600
 601
 602
 603
 604
 605
 606
 607
 608
 609
 610
 611
 612
 613
 614
 615
 616
 617
 618
 619
 620
 621
 622
 623
 624
 625
 626
 627
 628
 629
 630
 631
 632
 633
 634
 635
 636
 637
 638
 639
 640
 641
 642
 643
 644
 645
 646
 647
 648
 649
 650
 651
 652
 653
 654
 655
 656
 657
 658
 659
 660
 661
 662
 663
 664
 665
 666
 667
 668
 669
 670
 671
 672
 673
 674
 675
 676
 677
 678
 679
 680
 681
 682
 683
 684
 685
 686
 687
 688
 689
 690
 691
 692
 693
 694
 695
 696
 697
 698
 699
 700
 701
 702
 703
 704
 705
 706
 707
 708
 709
 710
 711
 712
 713
 714
 715
 716
 717
 718
 719
 720
 721
 722
 723
 724
 725
 726
 727
 728
 729
 730
 731
 732
 733
 734
 735
 736
 737
 738
 739
 740
 741
 742
 743
 744
 745
 746
 747
 748
 749
 750
 751
 752
 753
 754
 755
 756
 757
 758
 759
 760
 761
 762
 763
 764
 765
 766
 767
 768
 769
 770
 771
 772
 773
 774
 775
 776
 777
 778
 779
 780
 781
 782
 783
 784
 785
 786
 787
 788
 789
 790
 791
 792
 793
 794
 795
 796
 797
 798
 799
 800
 801
 802
 803
 804
 805
 806
 807
 808
 809
 810
 811
 812
 813
 814
 815
 816
 817
 818
 819
 820
 821
 822
 823
 824
 825
 826
 827
 828
 829
 830
 831
 832
 833
 834
 835
 836
 837
 838
 839
 840
 841
 842
 843
 844
 845
 846
 847
 848
 849
 850
 851
 852
 853
 854
 855
 856
 857
 858
 859
 860
 861
 862
 863
 864
 865
 866
 867
 868
 869
 870
 871
 872
 873
 874
 875
 876
 877
 878
 879
 880
 881
 882
 883
 884
 885
 886
 887
 888
 889
 890
 891
 892
 893
 894
 895
 896
 897
 898
 899
 900
 901
 902
 903
 904
 905
 906
 907
 908
 909
 910
 911
 912
 913
 914
 915
 916
 917
 918
 919
 920
 921
 922
 923
 924
 925
 926
 927
 928
 929
 930
 931
 932
 933
 934
 935
 936
 937
 938
 939
 940
 941
 942
 943
 944
 945
 946
 947
 948
 949
 950
 951
 952
 953
 954
 955
 956
 957
 958
 959
 960
 961
 962
 963
 964
 965
 966
 967
 968
 969
 970
 971
 972
 973
 974
 975
 976
 977
 978
 979
 980
 981
 982
 983
 984
 985
 986
 987
 988
 989
 990
 991
 992
 993
 994
 995
 996
 997
 998
 999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
class MiniMaxM3SparseAttention(nn.Module, AttentionLayerBase):
    """Block-sparse attention layer with the lightning-indexer branch.

    This is a merged attention layer: it owns the projections (qkv + index
    q/k), per-head QK norms and RoPE, *and* the attention-backend wiring that a
    generic ``Attention`` layer would normally provide — it binds the
    ``MiniMaxM3SparseBackend`` + main impl, registers the main paged K/V cache,
    and owns the lightning indexer (``MiniMaxM3Indexer``), which holds the
    index-key side cache.

    The index branch (index_{q,k}_proj + index_{q,k}_norm) feeds the sparse
    top-k block selection. M3 always disables the index value/output
    projections (``sparse_disable_index_value`` set for every sparse layer), so
    ``index_{v,o}_proj`` are never created.
    """

    def __init__(
        self,
        config: PretrainedConfig,
        layer_id: int,
        quant_config: QuantizationConfig | None = None,
        prefix: str = "",
        cache_config: CacheConfig | None = None,
        topk_indices_buffer: torch.Tensor | None = None,
        sparse_table_buffers: tuple[torch.Tensor, torch.Tensor] | None = None,
    ) -> None:
        super().__init__()
        self.hidden_size = config.hidden_size
        tp_size = get_tensor_model_parallel_world_size()

        self.total_num_heads = config.num_attention_heads
        assert self.total_num_heads % tp_size == 0
        self.num_heads = self.total_num_heads // tp_size
        self.total_num_kv_heads = config.num_key_value_heads
        if self.total_num_kv_heads >= tp_size:
            assert self.total_num_kv_heads % tp_size == 0
        else:
            assert tp_size % self.total_num_kv_heads == 0
        self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
        self.head_dim = config.head_dim
        self.q_size = self.num_heads * self.head_dim
        self.kv_size = self.num_kv_heads * self.head_dim
        self.scaling = self.head_dim**-0.5

        # Cross-layer index sharing (ATOM index_topk_freq): when True this sparse
        # layer reuses the previous compute layer's top-k block selection from the
        # shared topk_indices_buffer instead of recomputing it. Static per layer
        # -> cudagraph-capture-safe.
        self.skip_index_topk = _should_skip_index_topk(config, layer_id)

        # Sparse "index" branch dims. index_q has the same head count as the KV
        # heads (sparse_num_index_heads == num_key_value_heads), so it shards
        # identically -- including replication when tp_size > num_key_value_heads.
        sparse_cfg = config.sparse_attention_config
        self.total_idx_heads = sparse_cfg["sparse_num_index_heads"]
        self.num_idx_heads = self.num_kv_heads
        self.idx_head_dim = sparse_cfg["sparse_index_dim"]
        self.index_q_size = self.num_idx_heads * self.idx_head_dim

        # Single fused projection: q, k, v, index_q, index_k in one GEMM.
        self.qkv_proj = MinimaxM3QKVParallelLinearWithIndexer(
            self.hidden_size,
            self.head_dim,
            self.total_num_heads,
            self.total_num_kv_heads,
            self.total_idx_heads,
            self.idx_head_dim,
            bias=False,
            quant_config=quant_config,
            prefix=f"{prefix}.qkv_proj",
        )
        # reduce_results=False: the attention all-reduce is fused with the
        # following post_attention_layernorm (GemmaRMSNorm) in the decoder layer
        # via fused_allreduce_gemma_rms_norm.
        self.o_proj = RowParallelLinear(
            self.total_num_heads * self.head_dim,
            self.hidden_size,
            bias=False,
            reduce_results=False,
            quant_config=quant_config,
            prefix=f"{prefix}.o_proj",
        )

        # Per-head QK norm (qk_norm_type == "per_head", use_gemma_norm == True).
        self.q_norm = MiniMAXGemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps)
        self.k_norm = MiniMAXGemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps)

        # Partial RoPE: rotary_dim == head_dim * partial_rotary_factor. Honors
        # config.rope_scaling (e.g. YaRN) so long-context positions are covered.
        self.rotary_emb = _build_rotary_emb(config, self.head_dim)

        self.index_q_norm = MiniMAXGemmaRMSNorm(
            self.idx_head_dim, eps=config.rms_norm_eps
        )
        self.index_k_norm = MiniMAXGemmaRMSNorm(
            self.idx_head_dim, eps=config.rms_norm_eps
        )
        self.index_rotary_emb = self.rotary_emb

        # Attention-backend wiring.
        vllm_config = get_current_vllm_config()
        self.layer_name = f"{prefix}.attn"
        self.kv_cache_dtype = (
            cache_config.cache_dtype if cache_config is not None else "auto"
        )
        self.kv_cache_torch_dtype = kv_cache_dtype_str_to_dtype(
            self.kv_cache_dtype, vllm_config.model_config
        )
        # MiniMax-M3 sparse attention owns its KV-cache insert/read path instead
        # of wrapping the generic Attention module. Keep the same runtime scale
        # attributes so FP8 KV reads can honor vLLM's per-layer descale contract.
        set_default_quant_scales(self, register_buffer=True)

        # Indexer side-cache dtype, mirroring --kv-cache-dtype for the main cache
        # (--attention-config '{"indexer_kv_dtype": ...}'). fp8 e4m3 is what the
        # AITER indexer needs; bf16 keeps the Triton indexer.
        self.indexer_kv_dtype = vllm_config.attention_config.resolve_indexer_kv_dtype(
            "bf16"
        )

        # Shared top-k buffer: the indexer writes the selected blocks into it and
        # the attend impl reads them back (no Python value crosses the break).
        self.topk_indices_buffer = topk_indices_buffer
        # Indexer and main attention are separate impls. The AITER indexer is
        # ROCm-only, so it is selected here rather than inside the neutral
        # MiniMaxM3Indexer, which knows nothing about it and would pick Triton.
        indexer_impl_cls = select_aiter_indexer_impl_cls(
            topk_blocks=sparse_cfg["sparse_topk_blocks"],
            sparse_block_size=sparse_cfg["sparse_block_size"],
            num_index_heads=self.num_idx_heads,
            index_head_dim=self.idx_head_dim,
            indexer_kv_dtype=self.indexer_kv_dtype,
            score_type=sparse_cfg.get("sparse_score_type", "max"),
        )
        # The attend gate below needs this: an emitted table is what lets it
        # serve more than one KV head per rank. Buffers to write the table into
        # only arrive when the model already resolved the attend to the AITER
        # path the table addresses, so that is not rechecked here.
        self.indexer_emits_table = (
            indexer_impl_cls is not None and sparse_table_buffers is not None
        )
        # The backend names the metadata builder, which for the AITER path is
        # the one that rebases the block table the indexer's top-k resolves
        # through, so both have to come from the same selection.
        self.attn_backend, main_impl_cls = select_main_backend_and_impl_cls(
            topk_blocks=sparse_cfg["sparse_topk_blocks"],
            kv_cache_dtype=self.kv_cache_dtype,
            num_kv_heads=self.num_kv_heads,
            emits_sparse_block_table=self.indexer_emits_table,
        )
        # impl is AttentionImplBase (broader than AttentionLayerBase's annotation).
        self.impl: MiniMaxM3SparseImpl = main_impl_cls(  # type: ignore[assignment]
            self.num_heads,
            self.head_dim,
            self.scaling,
            self.num_kv_heads,
            kv_cache_dtype=self.kv_cache_dtype,
            topk_blocks=sparse_cfg["sparse_topk_blocks"],
            sparse_block_size=sparse_cfg["sparse_block_size"],
        )
        self.use_aiter_sparse_pa = minimax_m3_use_aiter_sparse_pa(
            self.num_kv_heads, emits_sparse_block_table=self.indexer_emits_table
        )
        self.kv_cache_k = torch.tensor([])
        self.kv_cache_v = torch.tensor([])
        self._aiter_sparse_pa_cache_data_ptr = 0
        self._aiter_sparse_pa_block_page_stride = 0
        indexer_kwargs = dict(
            num_kv_heads=self.num_kv_heads,
            scale=self.scaling,
            topk_blocks=sparse_cfg["sparse_topk_blocks"],
            sparse_block_size=sparse_cfg["sparse_block_size"],
            num_index_heads=self.num_idx_heads,
            index_head_dim=self.idx_head_dim,
            prefix=self.layer_name,
            init_blocks=sparse_cfg.get("sparse_init_block", 0),
            local_blocks=sparse_cfg.get("sparse_local_block", 0),
            score_type=sparse_cfg.get("sparse_score_type", "max"),
            cache_config=cache_config,
            indexer_kv_dtype=self.indexer_kv_dtype,
            topk_indices_buffer=topk_indices_buffer,
        )
        self.sparse_bt_buffer: torch.Tensor | None = None
        self.sparse_ctx_buffer: torch.Tensor | None = None
        if indexer_impl_cls is None:
            # Self-contained nn.Module: owns its side cache, selects its impl.
            self.indexer = MiniMaxM3Indexer(**indexer_kwargs)
        else:
            if self.indexer_emits_table:
                assert sparse_table_buffers is not None
                self.sparse_bt_buffer, self.sparse_ctx_buffer = sparse_table_buffers
            self.indexer = MiniMaxM3AiterIndexer(
                impl_cls=indexer_impl_cls,
                sparse_bt_buffer=self.sparse_bt_buffer,
                sparse_ctx_buffer=self.sparse_ctx_buffer,
                **indexer_kwargs,
            )

        # Register the main K/V cache so the KV-cache manager allocates it.
        compilation_config = vllm_config.compilation_config
        if self.layer_name in compilation_config.static_forward_context:
            raise ValueError(f"Duplicate layer name: {self.layer_name}")
        compilation_config.static_forward_context[self.layer_name] = self
        self.kv_cache = torch.tensor([])  # replaced by bind_kv_cache

    def get_attn_backend(self) -> type[MiniMaxM3SparseBackend]:
        return self.attn_backend

    def get_kv_cache_spec(self, vllm_config: VllmConfig) -> KVCacheSpec | None:
        # Main GQA K/V cache. Block size may change after load, refresh it.
        # AITER sparse PA wants K and V as two separate head groups
        # (H folded into the content dim).
        sparse_pa = self.use_aiter_sparse_pa
        kv_bytes = (
            self.num_kv_heads * self.head_dim * self.kv_cache_torch_dtype.itemsize
        )
        return FullAttentionSpec(
            block_size=vllm_config.cache_config.block_size,
            num_kv_heads=self.num_kv_heads,
            head_size=self.head_dim,
            head_size_v=self.head_dim,
            dtype=self.kv_cache_torch_dtype,
            kv_quant_mode=get_kv_quant_mode(self.kv_cache_dtype),
            num_head_slots=2 if sparse_pa else None,
            state_content_bytes=kv_bytes if sparse_pa else None,
        )

    def _ensure_aiter_sparse_pa_kv_cache(self) -> None:
        if self.kv_cache.numel() == 0:
            return
        if self._aiter_sparse_pa_cache_data_ptr == self.kv_cache.data_ptr():
            return

        kv_cache = self.kv_cache
        if is_quantized_kv_cache(self.kv_cache_dtype):
            kv_cache = kv_cache.view(self.impl.kv_cache_fp8_dtype)
        key_cache, value_cache = kv_cache.unbind(1)

        x = 16 // kv_cache.element_size()
        if self.head_dim % x != 0:
            raise RuntimeError(
                "MiniMax-M3 AITER sparse PA requires head_dim divisible by "
                f"16 / dtype_size, got head_dim={self.head_dim}, x={x}"
            )

        num_blocks, _, block_size, _ = kv_cache.shape
        pages_in_side = block_size // 16
        if key_cache.is_contiguous() and value_cache.is_contiguous():
            k_src, v_src = key_cache, value_cache
            block_pages, v_page_offset = pages_in_side, 0
        elif kv_cache.is_contiguous():
            k_src = v_src = kv_cache
            block_pages, v_page_offset = 2 * pages_in_side, pages_in_side
        else:
            raise RuntimeError(
                "MiniMax-M3 AITER sparse PA needs each K/V head slot stored as "
                "whole 16-token pages, but the resolved KV cache layout gives "
                "neither dense planes nor block-contiguous pages."
            )

        num_phys16 = num_blocks * block_pages
        self.kv_cache_k = k_src.view(
            num_phys16,
            self.num_kv_heads,
            self.head_dim // x,
            16,
            x,
        )
        # Offsetting V by half a block lets both sides share one page id.
        self.kv_cache_v = v_src.view(
            num_phys16,
            self.num_kv_heads,
            16 // x,
            self.head_dim,
            x,
        )[v_page_offset:]
        self._aiter_sparse_pa_block_page_stride = minimax_m3_sparse_block_page_stride(
            self.kv_cache_k, self.kv_cache_v
        )
        self._aiter_sparse_pa_cache_data_ptr = self.kv_cache.data_ptr()

    def get_aiter_sparse_pa_kv_cache(self) -> tuple[torch.Tensor, torch.Tensor]:
        self._ensure_aiter_sparse_pa_kv_cache()
        return self.kv_cache_k, self.kv_cache_v

    def _insert_aiter_sparse_pa_kv(
        self,
        k: torch.Tensor,
        v: torch.Tensor,
        index_k: torch.Tensor | None,
        slot_mapping: torch.Tensor,
        index_slot_mapping: torch.Tensor | None,
    ) -> None:
        if self.kv_cache.numel() == 0:
            return
        from aiter import reshape_and_cache

        from vllm.models.minimax_m3.amd.ops.sparse_pa import (
            minimax_m3_insert_index_cache,
        )

        key_cache, value_cache = self.get_aiter_sparse_pa_kv_cache()
        if value_cache.shape[0] != key_cache.shape[0]:
            attn_metadata = get_forward_context().attn_metadata
            page16_slot_mapping = None
            if isinstance(attn_metadata, dict):
                main_md = attn_metadata[self.layer_name]
                assert isinstance(main_md, MiniMaxM3SparseMetadata)
                page16_slot_mapping = main_md.page16_slot_mapping
            if (
                page16_slot_mapping is None
                or page16_slot_mapping.shape != slot_mapping.shape
            ):
                page16_slot_mapping = minimax_m3_rebase_slots_to_page16(
                    slot_mapping, self.kv_cache.shape[2]
                )
            slot_mapping = page16_slot_mapping
        kv_cache_dtype = (
            self.kv_cache_dtype
            if is_quantized_kv_cache(self.kv_cache_dtype)
            else "auto"
        )
        reshape_and_cache(
            k.contiguous(),
            v.contiguous(),
            key_cache,
            value_cache,
            slot_mapping,
            kv_cache_dtype=kv_cache_dtype,
            k_scale=getattr(self, "_k_scale", None),
            v_scale=getattr(self, "_v_scale", None),
            asm_layout=True,
        )
        if index_k is None or index_slot_mapping is None:
            return
        index_cache = self.indexer.index_cache.kv_cache
        if index_cache.numel() == 0:
            return
        minimax_m3_insert_index_cache(index_k, index_cache, index_slot_mapping)

    def forward(
        self,
        positions: torch.Tensor,
        hidden_states: torch.Tensor,
    ) -> torch.Tensor:
        # Single fused projection emitting [q | k | v | index_q | index_k].
        qkv, _ = self.qkv_proj(hidden_states)

        # Horizontally-fused per-head Gemma QK-norm + partial NeoX RoPE on the
        # main (q/k) and index (index_q/index_k) branches, all read straight out
        # of the single fused ``qkv`` tensor. Once the paged caches are bound the
        # kernel also inserts k/v and the index key into them (each with its own
        # slot_mapping); the memory-profiling run (caches unbound, no slot_mapping)
        # short-circuits to zeros below. The main and index slot mappings are read
        # from the forward context's slot_mapping dict, matching the
        # breakable-cudagraph path -- see nvidia/model.py.
        cos_sin_cache = self.rotary_emb.cos_sin_cache
        rotary_dim = self.rotary_emb.rotary_dim
        eps = self.q_norm.variance_epsilon
        num_tokens = qkv.shape[0]

        fwd_slot_mapping = get_forward_context().slot_mapping
        if (
            not isinstance(fwd_slot_mapping, dict)
            or self.layer_name not in fwd_slot_mapping
        ):
            # Memory-profiling run: caches not yet bound, slot_mapping is empty.
            return qkv.new_zeros((num_tokens, self.hidden_size))

        main_slot_mapping = fwd_slot_mapping[self.layer_name]
        q = qkv.new_empty((num_tokens, self.q_size))
        if self.skip_index_topk:
            index_q = None
            if self.use_aiter_sparse_pa:
                ops.fused_minimax_m3_qknorm_rope_kv_insert(
                    qkv,
                    self.q_norm.weight,
                    self.k_norm.weight,
                    cos_sin_cache,
                    positions,
                    self.num_heads,
                    self.num_kv_heads,
                    rotary_dim,
                    eps,
                    num_index_heads=self.num_idx_heads,
                    q_out=q,
                    skip_index_branch=True,
                )
                k_start = self.q_size
                v_start = k_start + self.kv_size
                k = qkv[:, k_start:v_start].view(
                    num_tokens, self.num_kv_heads, self.head_dim
                )
                v = qkv[:, v_start : v_start + self.kv_size].view(
                    num_tokens, self.num_kv_heads, self.head_dim
                )
                self._insert_aiter_sparse_pa_kv(
                    k,
                    v,
                    None,
                    main_slot_mapping,
                    None,
                )
            else:
                ops.fused_minimax_m3_qknorm_rope_kv_insert(
                    qkv,
                    self.q_norm.weight,
                    self.k_norm.weight,
                    cos_sin_cache,
                    positions,
                    self.num_heads,
                    self.num_kv_heads,
                    rotary_dim,
                    eps,
                    num_index_heads=self.num_idx_heads,
                    slot_mapping=main_slot_mapping,
                    kv_cache=self.kv_cache,
                    block_size=self.kv_cache.size(2),  # paged-cache block size
                    q_out=q,
                    kv_cache_dtype=self.kv_cache_dtype,
                    skip_index_branch=True,
                )
        else:
            index_slot_mapping = fwd_slot_mapping[self.indexer.index_cache.prefix]
            # index_q matches the index-K cache dtype (e4m3 for the AITER fp8
            # score path); the fused kernel emits fp8 directly when this buffer
            # is e4m3.
            index_q = qkv.new_empty(
                (num_tokens, self.index_q_size),
                dtype=self.indexer.index_cache.dtype,
            )
            if self.use_aiter_sparse_pa:
                ops.fused_minimax_m3_qknorm_rope_kv_insert(
                    qkv,
                    self.q_norm.weight,
                    self.k_norm.weight,
                    cos_sin_cache,
                    positions,
                    self.num_heads,
                    self.num_kv_heads,
                    rotary_dim,
                    eps,
                    self.index_q_norm.weight,
                    self.index_k_norm.weight,
                    self.num_idx_heads,
                    q_out=q,
                    index_q_out=index_q,
                    kv_cache_dtype=self.kv_cache_dtype,
                )
                k_start = self.q_size
                v_start = k_start + self.kv_size
                index_k_start = v_start + self.kv_size + self.index_q_size
                k = qkv[:, k_start:v_start].view(
                    num_tokens, self.num_kv_heads, self.head_dim
                )
                v = qkv[:, v_start : v_start + self.kv_size].view(
                    num_tokens, self.num_kv_heads, self.head_dim
                )
                index_k = qkv[
                    :, index_k_start : index_k_start + self.idx_head_dim
                ].view(num_tokens, self.idx_head_dim)
                self._insert_aiter_sparse_pa_kv(
                    k,
                    v,
                    index_k,
                    main_slot_mapping,
                    index_slot_mapping,
                )
            else:
                ops.fused_minimax_m3_qknorm_rope_kv_insert(
                    qkv,
                    self.q_norm.weight,
                    self.k_norm.weight,
                    cos_sin_cache,
                    positions,
                    self.num_heads,
                    self.num_kv_heads,
                    rotary_dim,
                    eps,
                    self.index_q_norm.weight,
                    self.index_k_norm.weight,
                    self.num_idx_heads,
                    main_slot_mapping,
                    index_slot_mapping,
                    self.kv_cache,
                    self.indexer.index_cache.kv_cache,
                    self.kv_cache.size(2),  # paged-cache block size
                    q,
                    index_q,
                    self.kv_cache_dtype,
                )

        output = torch.empty_like(q)
        attn_output = self._run_attention(q, index_q, output)
        output, _ = self.o_proj(attn_output)
        return output

    @eager_break_during_capture
    def _run_attention(
        self,
        query: torch.Tensor,
        index_query: torch.Tensor | None,
        output: torch.Tensor,
    ) -> torch.Tensor:
        # Single eager break around both: their split-K kernels read per-request
        # metadata and can't be captured into a cudagraph. The indexer writes its
        # top-k into the shared ``topk_indices_buffer``; the attend reads it back.
        # When skip_index_topk is set (ATOM index_topk_freq), reuse the selection
        # the preceding compute layer wrote into the shared buffer this forward.
        decode_sparse_table: tuple[torch.Tensor, torch.Tensor] | None = None
        if not self.skip_index_topk:
            assert index_query is not None
            # The AITER indexer emits the attend table into its persistent
            # buffers, resolving its selection through the page-16 rebase of the
            # attend's block table. The Triton indexer can instead fuse decode
            # top-k with main's per-forward sparse-table allocation.
            if self.indexer_emits_table:
                attn_metadata = get_forward_context().attn_metadata
                decode_page16 = prefill_page16 = None
                if isinstance(attn_metadata, dict):
                    main_md = attn_metadata[self.layer_name]
                    assert isinstance(main_md, MiniMaxM3SparseMetadata)
                    # Rebased once per step by the AITER attend's builder, which
                    # this path selected along with the impl.
                    d, p = main_md.decode, main_md.prefill
                    if d is not None:
                        assert isinstance(d, MiniMaxM3SparseAiterPADecodeMetadata)
                        decode_page16 = d.page16_block_table
                    if p is not None:
                        assert isinstance(p, MiniMaxM3SparseAiterPAPrefillMetadata)
                        prefill_page16 = p.page16_block_table
                self.indexer(
                    index_query,
                    decode_page16_block_table=decode_page16,
                    prefill_page16_block_table=prefill_page16,
                )
            else:
                attn_metadata = get_forward_context().attn_metadata
                if self.use_aiter_sparse_pa and isinstance(attn_metadata, dict):
                    main_md = attn_metadata[self.layer_name]
                    assert isinstance(main_md, MiniMaxM3SparseMetadata)
                    if main_md.num_decodes > 0:
                        d = main_md.decode
                        assert d is not None
                        topk = self.topk_indices_buffer
                        assert topk is not None
                        decode_sparse_table = minimax_m3_alloc_sparse_block_table(
                            topk[:, : main_md.num_decode_tokens, :]
                        )
                        block_page_stride = self._aiter_sparse_pa_block_page_stride
                        assert block_page_stride > 0
                        self.indexer(
                            index_query,
                            attention_block_table=d.block_table,
                            sparse_block_table_out=decode_sparse_table[0],
                            sparse_context_lens_out=decode_sparse_table[1],
                            block_page_stride=block_page_stride,
                        )
                    else:
                        self.indexer(index_query)
                else:
                    self.indexer(index_query)
        if self.use_aiter_sparse_pa:
            assert isinstance(self.impl, MiniMaxM3SparseAiterPAImpl)
            return self.impl.forward(
                self,
                query,
                self.kv_cache,
                output,
                decode_sparse_table=decode_sparse_table,
            )
        return self.impl.forward(self, query, self.kv_cache, output)

MiniMaxM3SparseForCausalLM

Bases: Module, SupportsPP, SupportsEagle3

MiniMax M3 (sparse/dense backbone) for causal language modeling.

Source code in vllm/models/minimax_m3/amd/model.py
class MiniMaxM3SparseForCausalLM(nn.Module, SupportsPP, SupportsEagle3):
    """MiniMax M3 (sparse/dense backbone) for causal language modeling."""

    packed_modules_mapping = {
        "qkv_proj": ["q_proj", "k_proj", "v_proj"],
        "gate_up_proj": ["gate_proj", "up_proj"],
    }

    def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
        super().__init__()
        config = vllm_config.model_config.hf_text_config
        quant_config = vllm_config.quant_config
        self.config = config
        self.quant_config = quant_config
        self.model = MiniMaxM3Model(
            vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
        )
        if get_pp_group().is_last_rank:
            self.lm_head = ParallelLMHead(
                config.vocab_size,
                config.hidden_size,
                quant_config=quant_config,
                prefix=maybe_prefix(prefix, "lm_head"),
            )
        else:
            self.lm_head = PPMissingLayer()
        self.logits_processor = LogitsProcessor(config.vocab_size)
        self.make_empty_intermediate_tensors = (  # type: ignore[method-assign]
            self.model.make_empty_intermediate_tensors
        )

    def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
        return self.model.embed_input_ids(input_ids)

    def forward(
        self,
        input_ids: torch.Tensor | None,
        positions: torch.Tensor,
        intermediate_tensors: IntermediateTensors | None = None,
        inputs_embeds: torch.Tensor | None = None,
        **kwargs,
    ) -> torch.Tensor | IntermediateTensors:
        return self.model(input_ids, positions, intermediate_tensors, inputs_embeds)

    def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor | None:
        return self.logits_processor(self.lm_head, hidden_states)

    def get_expert_mapping(self) -> list[tuple[str, str, int, str]]:
        return self.model.get_expert_mapping()

    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
        loader = AutoWeightsLoader(self)
        return loader.load_weights(weights)

MiniMaxM3SparseForConditionalGeneration

Bases: Module, SupportsMultiModal, SupportsPP, SupportsEagle3

Top-level (VL) entry point for MiniMax M3.

Owns the shared MiniMax-M3 vision tower on ROCm and delegates text generation to the AMD language-model path.

Source code in vllm/models/minimax_m3/amd/model.py
@MULTIMODAL_REGISTRY.register_processor(
    MiniMaxM3VLMultiModalProcessor,
    info=MiniMaxM3VLProcessingInfo,
    dummy_inputs=MiniMaxM3VLDummyInputsBuilder,
)
class MiniMaxM3SparseForConditionalGeneration(
    nn.Module, SupportsMultiModal, SupportsPP, SupportsEagle3
):
    """Top-level (VL) entry point for MiniMax M3.

    Owns the shared MiniMax-M3 vision tower on ROCm and delegates text
    generation to the AMD language-model path.
    """

    # The vision tower runs replicated per rank under ``--mm-encoder-tp-mode
    # data``; ``run_dp_sharded_mrope_vision_model`` shards the work across
    # ranks (see ``_process_image_input`` / ``_process_video_input``).
    supports_encoder_tp_data = True

    packed_modules_mapping = {
        "qkv_proj": ["q_proj", "k_proj", "v_proj"],
        "gate_up_proj": ["gate_proj", "up_proj"],
    }

    hf_to_vllm_mapper = WeightsMapper(
        orig_to_new_prefix={
            "multi_modal_projector.": "vision_tower.multi_modal_projector.",
            "patch_merge_mlp.": "vision_tower.patch_merge_mlp.",
        },
        orig_to_new_substr={
            ".mlp.fc1": ".fc1",
            ".mlp.fc2": ".fc2",
        },
    )

    @classmethod
    def get_placeholder_str(cls, modality: str, i: int) -> str | None:
        if modality == "image":
            return MiniMaxM3VLProcessingInfo.IMAGE_TOKEN
        if modality == "video":
            return MiniMaxM3VLProcessingInfo.VIDEO_TOKEN
        raise ValueError(f"Unsupported modality: {modality!r}")

    def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
        super().__init__()
        config = vllm_config.model_config.hf_config
        self.config = config
        self.quant_config = vllm_config.quant_config
        self.multimodal_config = vllm_config.model_config.multimodal_config
        assert self.multimodal_config is not None
        self.use_data_parallel = self.multimodal_config.mm_encoder_tp_mode == "data"

        text_hidden_size = getattr(config.text_config, "hidden_size", None)
        assert text_hidden_size is not None, "text_config.hidden_size is required"
        projector_hidden_size = getattr(config, "projector_hidden_size", None)

        with self._mark_tower_model(vllm_config, {"image", "video"}):
            vision_config = config.vision_config
            self.vision_tower = MiniMaxVLVisionModel(
                config=PretrainedConfig.from_dict(vision_config),
                text_hidden_size=text_hidden_size,
                projector_hidden_size=projector_hidden_size,
                quant_config=self.quant_config,
                prefix=maybe_prefix(prefix, "vision_tower"),
            )

        self.language_model = init_vllm_registered_model(
            vllm_config=vllm_config,
            hf_config=config.text_config,
            prefix=maybe_prefix(prefix, "language_model"),
            architectures=["MiniMaxM3SparseForCausalLM"],
        )
        self.make_empty_intermediate_tensors = (  # type: ignore[method-assign]
            self.language_model.make_empty_intermediate_tensors
        )

    def _parse_and_validate_image_input(self, **kwargs: object) -> dict | None:
        pixel_values = kwargs.pop("pixel_values", None)
        image_grid_thw = kwargs.pop("image_grid_thw", None)
        if pixel_values is None:
            return None
        return {"pixel_values": pixel_values, "image_grid_thw": image_grid_thw}

    def _parse_and_validate_video_input(self, **kwargs: object) -> dict | None:
        pixel_values_videos = kwargs.pop("pixel_values_videos", None)
        video_grid_thw = kwargs.pop("video_grid_thw", None)
        if pixel_values_videos is None:
            return None
        return {
            "pixel_values_videos": pixel_values_videos,
            "video_grid_thw": video_grid_thw,
        }

    def _process_image_input(self, image_input: dict) -> tuple[torch.Tensor, ...]:
        pixel_values: torch.Tensor = image_input["pixel_values"].type(
            self.vision_tower.dtype
        )
        grid_thw: torch.Tensor = image_input["image_grid_thw"]
        assert grid_thw.ndim == 2

        if self.use_data_parallel:
            # Already returns a per-item tuple of embeddings.
            return run_dp_sharded_mrope_vision_model(
                self.vision_tower,
                pixel_values,
                grid_thw.tolist(),
                rope_type="rope_3d",
            )

        image_embeds = self.vision_tower(
            pixel_values=pixel_values,
            grid_thw=grid_thw.tolist(),
        )

        # Split the concatenated output into one tensor per image item.
        merge_size = self.vision_tower.spatial_merge_size
        sizes = (grid_thw.prod(-1) // (merge_size * merge_size)).tolist()
        return image_embeds.split(sizes)

    def _process_video_input(self, video_input: dict) -> tuple[torch.Tensor, ...]:
        pixel_values: torch.Tensor = video_input["pixel_values_videos"].type(
            self.vision_tower.dtype
        )
        grid_thw: torch.Tensor = video_input["video_grid_thw"]
        assert grid_thw.ndim == 2

        if self.use_data_parallel:
            # Already returns a per-item tuple of embeddings.
            return run_dp_sharded_mrope_vision_model(
                self.vision_tower,
                pixel_values,
                grid_thw.tolist(),
                rope_type="rope_3d",
            )

        video_embeds = self.vision_tower(
            pixel_values=pixel_values,
            grid_thw=grid_thw.tolist(),
        )

        # Split the concatenated output into one tensor per video item.
        merge_size = self.vision_tower.spatial_merge_size
        sizes = (grid_thw.prod(-1) // (merge_size * merge_size)).tolist()
        return video_embeds.split(sizes)

    def _parse_and_validate_multimodal_inputs(
        self, **kwargs: object
    ) -> dict[str, dict]:
        mm_input_by_modality: dict[str, dict] = {}
        for input_key in kwargs:
            if input_key == "pixel_values" and "image" not in mm_input_by_modality:
                image_input = self._parse_and_validate_image_input(**kwargs)
                if image_input is not None:
                    mm_input_by_modality["image"] = image_input
            if (
                input_key == "pixel_values_videos"
                and "video" not in mm_input_by_modality
            ):
                video_input = self._parse_and_validate_video_input(**kwargs)
                if video_input is not None:
                    mm_input_by_modality["video"] = video_input
        return mm_input_by_modality

    def embed_multimodal(self, **kwargs: object) -> MultiModalEmbeddings:
        mm_input_by_modality = self._parse_and_validate_multimodal_inputs(**kwargs)
        if not mm_input_by_modality:
            return []

        multimodal_embeddings: list[torch.Tensor] = []
        for modality in mm_input_by_modality:
            multimodal_input = mm_input_by_modality[modality]
            if modality == "image":
                image_embeddings = self._process_image_input(multimodal_input)
                multimodal_embeddings.extend(image_embeddings)
            if modality == "video":
                video_embeddings = self._process_video_input(multimodal_input)
                multimodal_embeddings.extend(video_embeddings)

        return tuple(multimodal_embeddings)

    def forward(
        self,
        input_ids: torch.Tensor | None,
        positions: torch.Tensor,
        intermediate_tensors: IntermediateTensors | None = None,
        inputs_embeds: torch.Tensor | None = None,
        **kwargs,
    ) -> torch.Tensor:
        return self.language_model(
            input_ids, positions, intermediate_tensors, inputs_embeds
        )

    def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor | None:
        return self.language_model.compute_logits(hidden_states)

    def get_expert_mapping(self) -> list[tuple[str, str, int, str]]:
        return self.language_model.get_expert_mapping()

    def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
        loader = AutoWeightsLoader(self)
        return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper)

_aiter_moe_fused_shared_experts_enabled(is_fused_shared_expert_enabled)

Whether the fused shared expert routes through aiter's grouped top-k MoE.

A strict sub-case of is_fused_shared_expert_enabled: shared-expert fusion must already be opted in (VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS) and allowed (not under expert parallelism). When additionally on gfx950 with an active aiter MoE backend, the shared expert is appended inside aiter's biased grouped top-k kernel (num_fused_shared_experts) instead of the vLLM router's torch concat. Otherwise FSE still runs via the vLLM top-k bias router.

Source code in vllm/models/minimax_m3/amd/model.py
def _aiter_moe_fused_shared_experts_enabled(
    is_fused_shared_expert_enabled: bool,
) -> bool:
    """Whether the fused shared expert routes through aiter's grouped top-k MoE.

    A strict sub-case of `is_fused_shared_expert_enabled`: shared-expert
    fusion must already be opted in (``VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS``)
    and allowed (not under expert parallelism). When additionally on gfx950 with
    an active aiter MoE backend, the shared expert is appended inside aiter's
    biased grouped top-k kernel (``num_fused_shared_experts``) instead of the
    vLLM router's torch concat. Otherwise FSE still runs via the vLLM top-k bias
    router.
    """
    from vllm.platforms.rocm import on_gfx950

    return (
        on_gfx950()
        and is_fused_shared_expert_enabled
        and rocm_aiter_ops.is_fusion_moe_shared_experts_enabled()
    )

_build_rotary_emb(config, head_dim)

Build the (partial NeoX) RoPE, honoring an optional rope_scaling config.

Without scaling the cos/sin cache is sized to max_position_embeddings (524288 native); a request whose positions exceed that reads the cache out of bounds and the worker hard-crashes (no Python traceback). When rope_scaling is set (e.g. YaRN factor: 2 to reach 1M), thread it into get_rope so the proper scaled embedding is built and its cache covers original_max_position_embeddings * factor positions. Default behavior (no scaling) is unchanged. Shared by the dense and sparse attention layers, and the index branch reuses the returned module.

Note: for the VL checkpoint, set rope_scaling on the text config (--hf-overrides '{"text_config":{"rope_scaling":{...}}}') -- that is the config the decoder reads here; a top-level override does not reach it.

Source code in vllm/models/minimax_m3/amd/model.py
def _build_rotary_emb(config: PretrainedConfig, head_dim: int):
    """Build the (partial NeoX) RoPE, honoring an optional ``rope_scaling`` config.

    Without scaling the cos/sin cache is sized to ``max_position_embeddings``
    (524288 native); a request whose positions exceed that reads the cache out of
    bounds and the worker hard-crashes (no Python traceback). When ``rope_scaling``
    is set (e.g. YaRN ``factor: 2`` to reach 1M), thread it into ``get_rope`` so the
    proper scaled embedding is built and its cache covers
    ``original_max_position_embeddings * factor`` positions. Default behavior
    (no scaling) is unchanged. Shared by the dense and sparse attention layers, and
    the index branch reuses the returned module.

    Note: for the VL checkpoint, set ``rope_scaling`` on the *text* config
    (``--hf-overrides '{"text_config":{"rope_scaling":{...}}}'``) -- that is the
    config the decoder reads here; a top-level override does not reach it.
    """
    rope_parameters = {
        "rope_theta": config.rope_theta,
        "partial_rotary_factor": config.partial_rotary_factor,
    }
    max_position = config.max_position_embeddings
    rope_scaling = getattr(config, "rope_scaling", None)
    if rope_scaling:
        rope_parameters.update(rope_scaling)
        # HF uses "rope_type" (older configs: "type"); get_rope reads "rope_type".
        if "rope_type" not in rope_parameters and "type" in rope_scaling:
            rope_parameters["rope_type"] = rope_scaling["type"]
        rope_parameters.setdefault(
            "original_max_position_embeddings", config.max_position_embeddings
        )
        factor = float(rope_scaling.get("factor", 1.0))
        # Cover the extended range (informational for get_rope's default branch;
        # the YaRN embedding sizes its own cache from original * factor).
        max_position = int(rope_parameters["original_max_position_embeddings"] * factor)
    return get_rope(
        head_dim,
        max_position=max_position,
        rope_parameters=rope_parameters,
    )

_is_moe_layer(config, layer_id)

Whether this layer's MLP is a sparse MoE block (vs a dense MLP).

Source code in vllm/models/minimax_m3/amd/model.py
def _is_moe_layer(config: PretrainedConfig, layer_id: int) -> bool:
    """Whether this layer's MLP is a sparse MoE block (vs a dense MLP)."""
    moe_layer_freq = getattr(config, "moe_layer_freq", None)
    if moe_layer_freq is None:
        return True
    return moe_layer_freq[layer_id] != 0

_should_skip_index_topk(config, layer_id)

ATOM index_topk_freq (cross-layer index sharing).

Only 1 of every index_topk_freq sparse-attention layers recomputes the lightning-indexer top-k block selection; the rest reuse the selection the preceding compute layer wrote into the shared topk_indices_buffer this same forward pass. This cuts the indexer score + top-k cost ~freqx with negligible accuracy impact (adjacent sparse layers pick nearly the same blocks; ATOM validated GSM8K with freq=4). Gated by use_index_cache; enable via --hf-overrides '{"use_index_cache": true, "index_topk_freq": 4}'.

Source code in vllm/models/minimax_m3/amd/model.py
def _should_skip_index_topk(config: PretrainedConfig, layer_id: int) -> bool:
    """ATOM ``index_topk_freq`` (cross-layer index sharing).

    Only 1 of every ``index_topk_freq`` sparse-attention layers recomputes the
    lightning-indexer top-k block selection; the rest reuse the selection the
    preceding compute layer wrote into the shared ``topk_indices_buffer`` this
    same forward pass. This cuts the indexer score + top-k cost ~``freq``x with
    negligible accuracy impact (adjacent sparse layers pick nearly the same
    blocks; ATOM validated GSM8K with freq=4). Gated by ``use_index_cache``;
    enable via ``--hf-overrides '{"use_index_cache": true, "index_topk_freq": 4}'``.
    """
    if not getattr(config, "use_index_cache", False):
        return False
    freq = int(getattr(config, "index_topk_freq", 1) or 1)
    if freq <= 1:
        return False
    ordinal = _sparse_attention_layer_ordinals(config).get(layer_id)
    if ordinal is None:
        return False
    offset = int(getattr(config, "index_skip_topk_offset", 0) or 0)
    return max(ordinal - offset, 0) % freq != 0

_sparse_attention_layer_ids(config)

Layer ids whose attention runs the extra sparse "index" branch.

Source code in vllm/models/minimax_m3/amd/model.py
def _sparse_attention_layer_ids(config: PretrainedConfig) -> set[int]:
    """Layer ids whose attention runs the extra sparse "index" branch."""
    cfg = getattr(config, "sparse_attention_config", None)
    if not cfg:
        return set()
    freq = cfg.get("sparse_attention_freq")
    if freq is None:
        return set()
    return {i for i, f in enumerate(freq) if f != 0}

_sparse_attention_layer_ordinals(config)

Map each sparse-attention layer id to its ordinal among sparse layers.

Source code in vllm/models/minimax_m3/amd/model.py
def _sparse_attention_layer_ordinals(config: PretrainedConfig) -> dict[int, int]:
    """Map each sparse-attention layer id to its ordinal among sparse layers."""
    return {
        lid: ordinal
        for ordinal, lid in enumerate(sorted(_sparse_attention_layer_ids(config)))
    }