Skip to content

vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse

Classes:

Functions:

ROCMAiterMLASparseImpl

Bases: MLAAttentionImpl[ROCMAiterMLASparseMetadata], SharedTopkIndicesBuffer

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
 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
1154
1155
1156
class ROCMAiterMLASparseImpl(
    MLAAttentionImpl[ROCMAiterMLASparseMetadata], SharedTopkIndicesBuffer
):
    is_sparse = True
    supports_dense_mha_prefill = False
    supports_dcp = False
    use_aiter_sparse_mla = False

    def __init__(
        self,
        num_heads: int,
        head_size: int,
        scale: float,
        num_kv_heads: int,
        alibi_slopes: list[float] | None,
        sliding_window: int | None,
        kv_cache_dtype: str,
        logits_soft_cap: float | None,
        attn_type: str,
        kv_sharing_target_layer_name: str | None,
        # MLA Specific Arguments
        topk_indices_buffer: torch.Tensor | None = None,
        indexer: "Indexer | None" = None,
        **mla_args,
    ) -> None:
        AiterMLAHelper.check_num_heads_validity(num_heads)

        self.num_heads = num_heads
        self.head_size = head_size
        self.scale = float(scale)
        self.num_kv_heads = num_kv_heads
        sinks = mla_args.pop("sinks", None)
        if sinks is not None:
            if sinks.dtype != torch.float32:
                raise ValueError(
                    f"ROCm AITER MLA sinks must be float32, got {sinks.dtype}"
                )
            if sinks.ndim != 1 or sinks.numel() != num_heads:
                raise ValueError(
                    "ROCm AITER MLA sinks must have shape "
                    f"({num_heads},), got {tuple(sinks.shape)}"
                )
            if not sinks.is_contiguous():
                raise ValueError("ROCm AITER MLA sinks must be contiguous")
        self.sinks: torch.Tensor | None = sinks
        self.kv_cache_dtype = kv_cache_dtype
        self.kv_lora_rank: int = mla_args["kv_lora_rank"]
        self.softmax_scale = scale
        self.init_topk_indices_buffer(indexer, topk_indices_buffer)

        vllm_config = get_current_vllm_config()
        max_tokens = vllm_config.scheduler_config.max_num_batched_tokens
        q_concat_shape = (max_tokens, num_heads, head_size)
        (self.q_concat_buffer,) = current_workspace_manager().get_simultaneous(
            (q_concat_shape, vllm_config.model_config.dtype),
        )
        self.qk_rope_head_dim: int = mla_args["qk_rope_head_dim"]

        if rocm_aiter_ops.is_triton_sparse_mla_enabled():
            reason = self._aiter_sparse_mla_unsupported_reason(vllm_config)
            if reason is None:
                self.use_aiter_sparse_mla = True
            else:
                logger.warning_once(
                    "VLLM_ROCM_USE_AITER_TRITON_SPARSE_MLA is set, but %s; "
                    "using the default sparse MLA path instead.",
                    reason,
                )

    def _aiter_sparse_mla_unsupported_reason(
        self, vllm_config: VllmConfig
    ) -> str | None:
        model_dtype = vllm_config.model_config.dtype
        if model_dtype != torch.bfloat16:
            return f"the model dtype is {model_dtype}, not bfloat16"
        cache_dtype = self.kv_cache_dtype
        if cache_dtype not in ("auto", "bfloat16", "fp8", "fp8_e4m3"):
            return f"kv-cache dtype {cache_dtype!r} is neither bfloat16 nor e4m3 fp8"
        # Context parallelism merges each rank's partials with a per-token LSE.
        if self.dcp_world_size > 1 or self.pcp_world_size > 1:
            return "context parallelism is not supported"
        return None

    def record_logical_topk_ready(self) -> None:
        # This impl shares the top-k indices buffer via SharedTopkIndicesBuffer
        # but does not participate in sparse-MLA index groups.
        pass

    def _forward_mla(
        self,
        layer: AttentionLayer,
        q: torch.Tensor,  # [sq, heads, d_qk]
        kv_c_and_k_pe_cache: torch.Tensor,  # [blocks, heads, d_qk]
        attn_metadata: ROCMAiterMLASparseMetadata,
    ) -> tuple[torch.Tensor, torch.Tensor | None]:
        num_tokens = q.shape[0]
        base_mla_num_heads = AiterMLAHelper.get_actual_mla_num_heads(self.num_heads)
        mla_num_heads = base_mla_num_heads
        need_lse = self.sinks is not None
        from vllm.platforms.rocm import on_gfx942

        # Keep sink attention available for dtypes/head shapes without an
        # AITER return-LSE kernel. Sink layers never need persistent metadata.
        triton_sink_fallback = (
            need_lse
            and q.dtype == kv_c_and_k_pe_cache.dtype
            and (
                q.dtype == torch.float16
                or (
                    q.dtype == torch.bfloat16
                    and (
                        mla_num_heads > 128
                        or (on_gfx942() and 32 < mla_num_heads <= 64)
                    )
                )
            )
        )

        if triton_sink_fallback or _use_rocm_sparse_triton(
            kv_cache_dtype=self.kv_cache_dtype,
            head_size=q.shape[-1],
            kv_lora_rank=self.kv_lora_rank,
            num_prefills=attn_metadata.num_prefills,
            num_decodes=attn_metadata.num_decodes,
            num_decode_tokens=attn_metadata.num_decode_tokens,
            max_query_len=attn_metadata.max_query_len,
        ):
            output = torch.empty(
                [num_tokens, q.shape[1], self.kv_lora_rank],
                dtype=attn_metadata.attn_out_dtype,
                device=q.device,
            )
            triton_sinks = None
            if self.sinks is not None:
                triton_sinks = AiterMLAHelper.get_mla_padded_q(
                    self.num_heads,
                    self.sinks.reshape(1, self.num_heads, 1),
                    q.shape[1],
                ).reshape(-1)
            rocm_sparse_attn_prefill(
                q=q,
                kv=kv_c_and_k_pe_cache.view(-1, 1, q.shape[-1]),
                indices=None,
                topk_length=None,
                scale=self.scale,
                head_dim=q.shape[-1],
                nope_head_dim=self.kv_lora_rank,
                rope_head_dim=q.shape[-1] - self.kv_lora_rank,
                attn_sink=triton_sinks,
                output=output,
                ragged_indices=attn_metadata.paged_kv_indices,
                ragged_indptr=attn_metadata.paged_kv_indptr,
            )
            output = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output)
            return output, None

        # AITER's nonpersistent return-LSE dispatch has discrete head kernels.
        supported_head_buckets: tuple[int, ...] | None = None
        head_dtype_name = ""
        if need_lse:
            if (
                q.dtype == torch.bfloat16
                and kv_c_and_k_pe_cache.dtype == torch.bfloat16
            ):
                supported_head_buckets = (16, 32, 64, 128)
                head_dtype_name = "BF16"
            elif (
                q.dtype == current_platform.fp8_dtype()
                and kv_c_and_k_pe_cache.dtype == current_platform.fp8_dtype()
            ):
                supported_head_buckets = (16, 128)
                head_dtype_name = "FP8"
            else:
                raise ValueError(
                    "ROCm AITER MLA attention sinks require query and KV to "
                    "both use BF16 or both use FP8, got "
                    f"query={q.dtype}, KV={kv_c_and_k_pe_cache.dtype}"
                )

        if supported_head_buckets is not None:
            from vllm.platforms.rocm import on_mi3xx

            if on_mi3xx():
                supported_heads = next(
                    (
                        heads
                        for heads in supported_head_buckets
                        if heads >= mla_num_heads
                    ),
                    None,
                )
                if supported_heads is None:
                    raise ValueError(
                        "ROCm AITER MLA attention sinks support at most 128 "
                        f"padded local {head_dtype_name} heads; increase "
                        "tensor_parallel_size"
                    )
                if supported_heads != mla_num_heads:
                    q = AiterMLAHelper.get_mla_padded_q(
                        mla_num_heads, q, supported_heads
                    )
                    mla_num_heads = supported_heads
        output = torch.empty(
            [num_tokens, mla_num_heads, self.kv_lora_rank],
            dtype=attn_metadata.attn_out_dtype,
            device=q.device,
        )

        if need_lse:
            # gfx942 has no persistent MLA code object that writes final LSE.
            # The split-KV path consumes the same ragged indices and is exact.
            lse = rocm_aiter_ops.mla_decode_fwd_lse(
                q,
                kv_c_and_k_pe_cache,
                output,
                self.scale,
                attn_metadata.qo_indptr,
                1,
                attn_metadata.paged_kv_indptr,
                attn_metadata.paged_kv_indices,
                attn_metadata.paged_kv_last_page_len,
                q_scale=layer._q_scale,
                kv_scale=layer._k_scale,
            )
        else:
            # Preserve the persistent work-stealing fast path for models that
            # do not use sinks.
            mla_kwargs: dict = dict(
                q_scale=layer._q_scale,
                kv_scale=layer._k_scale,
            )
            if attn_metadata.work_meta_data is not None:
                mla_kwargs.update(
                    work_meta_data=attn_metadata.work_meta_data,
                    work_indptr=attn_metadata.work_indptr,
                    work_info_set=attn_metadata.work_info_set,
                    reduce_indptr=attn_metadata.reduce_indptr,
                    reduce_final_map=attn_metadata.reduce_final_map,
                    reduce_partial_map=attn_metadata.reduce_partial_map,
                )

            rocm_aiter_ops.mla_decode_fwd(
                q,
                kv_c_and_k_pe_cache,
                output,
                self.scale,
                attn_metadata.qo_indptr,
                1,
                attn_metadata.paged_kv_indptr,
                attn_metadata.paged_kv_indices,
                attn_metadata.paged_kv_last_page_len,
                **mla_kwargs,
            )
            lse = None

        if mla_num_heads != base_mla_num_heads:
            if mla_num_heads % base_mla_num_heads == 0:
                head_stride = mla_num_heads // base_mla_num_heads
                output = output[:, ::head_stride]
                if lse is not None:
                    lse = lse[:, ::head_stride]
            else:
                output = output[:, :base_mla_num_heads]
                if lse is not None:
                    lse = lse[:, :base_mla_num_heads]

        output = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output)
        if lse is not None:
            lse = AiterMLAHelper.get_mla_unpadded_o(
                self.num_heads, lse.unsqueeze(-1)
            ).squeeze(-1)

        if self.sinks is not None:
            assert lse is not None
            # Empty ragged rows have only sink mass and no value contribution.
            # AITER can return NaN output/LSE for those rows; do not multiply it
            # by a zero normalization factor and propagate the NaN.
            has_keys = (
                attn_metadata.paged_kv_indptr[1:] > attn_metadata.paged_kv_indptr[:-1]
            ).unsqueeze(-1)
            lse = torch.where(has_keys, lse, float("-inf"))
            sink_lse = torch.logaddexp(lse, self.sinks)
            sink_scale = torch.exp(lse - sink_lse)
            output = torch.where(
                has_keys.unsqueeze(-1),
                output.float() * sink_scale.unsqueeze(-1),
                0.0,
            ).to(output.dtype)
            lse = sink_lse

        return output, lse

    def _forward_mla_aiter(
        self,
        layer: AttentionLayer,
        q: torch.Tensor,  # [sq, heads, d_qk], not head-padded
        kv_c_and_k_pe_cache: torch.Tensor,
        attn_metadata: ROCMAiterMLASparseMetadata,
    ) -> torch.Tensor:
        """_forward_mla on aiter's Triton sparse MLA kernel.

        Reads the same index stream, but needs no q head padding and no
        persistent metadata, and its launch depends on shapes only, so it is
        CUDA-graph capturable.
        """
        num_actual_tokens = attn_metadata.num_actual_tokens
        output = torch.empty(
            [q.shape[0], self.num_heads, self.kv_lora_rank],
            dtype=attn_metadata.attn_out_dtype,
            device=q.device,
        )
        rocm_aiter_ops.triton_sparse_mla_fwd(
            q[:num_actual_tokens],
            kv_c_and_k_pe_cache.view(-1, 1, 1, q.shape[-1]),
            output[:num_actual_tokens],
            self.scale,
            attn_metadata.paged_kv_indptr,
            attn_metadata.paged_kv_indices,
            kv_lora_rank=self.kv_lora_rank,
            qk_rope_head_dim=self.qk_rope_head_dim,
            q_scale=layer._q_scale,
            kv_scale=layer._k_scale,
            attn_sink=self.sinks,
            # triton_convert_req_index_to_global_index writes 0, never -1,
            # for an invalid top-k entry, so no slot in the stream is negative.
            has_invalid=False,
        )
        return output

    def forward_mqa(
        self,
        q: torch.Tensor | tuple[torch.Tensor, torch.Tensor],
        kv_c_and_k_pe_cache: torch.Tensor,
        attn_metadata: ROCMAiterMLASparseMetadata,
        layer: AttentionLayer,
    ) -> tuple[torch.Tensor, torch.Tensor | None]:
        # NOTE(lucas): for the sparse FlashMLA kernels the kernels want to use
        # MQA 576/512 approach for both prefill and decode

        fp8_attention = self.kv_cache_dtype.startswith("fp8")
        if isinstance(q, tuple):
            ql_nope, q_pe = q
            if fp8_attention:
                q = layer._decode_concat_quant_fp8_op(  # type: ignore[attr-defined]
                    ql_nope, q_pe, layer._q_scale
                )
            else:
                q = self.q_concat_buffer[: ql_nope.shape[0]]
                if q_pe.shape[-1] == 0:
                    q.copy_(ql_nope)
                elif q.dtype == torch.float16:
                    torch.cat((ql_nope, q_pe), dim=-1, out=q)
                else:
                    ops.concat_mla_q(ql_nope, q_pe, q)

        num_actual_toks = attn_metadata.num_actual_tokens

        # Get topk indices
        assert self.topk_indices_buffer is not None
        topk_indices = fit_kpool_indices_to_aiter(
            self.topk_indices_buffer[:num_actual_toks], attn_metadata.topk_tokens
        )

        triton_convert_req_index_to_global_index(
            attn_metadata.req_id_per_token,
            attn_metadata.block_table,
            topk_indices,
            attn_metadata.paged_kv_indptr,
            attn_metadata.paged_kv_indices,
            BLOCK_SIZE=attn_metadata.block_size,
            NUM_TOPK_TOKENS=attn_metadata.topk_tokens,
        )

        # write the latent and rope to kv cache
        if fp8_attention:
            kv_c_and_k_pe_cache = kv_c_and_k_pe_cache.view(current_platform.fp8_dtype())
            if q.dtype != current_platform.fp8_dtype():
                original_q_shape = q.shape
                q, _ = ops.scaled_fp8_quant(q.view(q.shape[0], -1), layer._q_scale)
                q = q.view(original_q_shape)
        if self.use_aiter_sparse_mla:
            output = self._forward_mla_aiter(
                layer, q, kv_c_and_k_pe_cache, attn_metadata
            )
            return output, None
        mla_padded_q = AiterMLAHelper.get_mla_padded_q(self.num_heads, q)
        return self._forward_mla(
            layer, mla_padded_q, kv_c_and_k_pe_cache, attn_metadata
        )

_forward_mla_aiter(layer, q, kv_c_and_k_pe_cache, attn_metadata)

_forward_mla on aiter's Triton sparse MLA kernel.

Reads the same index stream, but needs no q head padding and no persistent metadata, and its launch depends on shapes only, so it is CUDA-graph capturable.

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
def _forward_mla_aiter(
    self,
    layer: AttentionLayer,
    q: torch.Tensor,  # [sq, heads, d_qk], not head-padded
    kv_c_and_k_pe_cache: torch.Tensor,
    attn_metadata: ROCMAiterMLASparseMetadata,
) -> torch.Tensor:
    """_forward_mla on aiter's Triton sparse MLA kernel.

    Reads the same index stream, but needs no q head padding and no
    persistent metadata, and its launch depends on shapes only, so it is
    CUDA-graph capturable.
    """
    num_actual_tokens = attn_metadata.num_actual_tokens
    output = torch.empty(
        [q.shape[0], self.num_heads, self.kv_lora_rank],
        dtype=attn_metadata.attn_out_dtype,
        device=q.device,
    )
    rocm_aiter_ops.triton_sparse_mla_fwd(
        q[:num_actual_tokens],
        kv_c_and_k_pe_cache.view(-1, 1, 1, q.shape[-1]),
        output[:num_actual_tokens],
        self.scale,
        attn_metadata.paged_kv_indptr,
        attn_metadata.paged_kv_indices,
        kv_lora_rank=self.kv_lora_rank,
        qk_rope_head_dim=self.qk_rope_head_dim,
        q_scale=layer._q_scale,
        kv_scale=layer._k_scale,
        attn_sink=self.sinks,
        # triton_convert_req_index_to_global_index writes 0, never -1,
        # for an invalid top-k entry, so no slot in the stream is negative.
        has_invalid=False,
    )
    return output

ROCMAiterMLASparseMetadataBuilder dataclass

Bases: AttentionMetadataBuilder[ROCMAiterMLASparseMetadata]

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
@dataclass
class ROCMAiterMLASparseMetadataBuilder(
    AttentionMetadataBuilder[ROCMAiterMLASparseMetadata]
):
    _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH

    def __init__(
        self,
        kv_cache_spec: AttentionSpec,
        layer_names: list[str],
        vllm_config: VllmConfig,
        device: torch.device,
    ):
        self.kv_cache_spec = kv_cache_spec
        self.model_config = vllm_config.model_config
        self.model_dtype = vllm_config.model_config.dtype
        self.kv_cache_dtype = vllm_config.cache_config.cache_dtype
        parallel_config = vllm_config.parallel_config
        self.device = device
        max_num_batched_tokens = vllm_config.scheduler_config.max_num_batched_tokens

        self.vllm_config = vllm_config
        self._init_reorder_batch_threshold(1, supports_spec_as_decode=True)

        self.num_heads = self.model_config.get_num_attention_heads(parallel_config)
        self.mla_dims = get_mla_dims(self.model_config)
        self.topk_tokens = vllm_config.model_config.hf_text_config.index_topk
        attention_context = vllm_config.compilation_config.static_forward_context
        # Sink decode must use AITER's nonpersistent path. In particular,
        # gfx942 has no persistent+LSE kernel, and its metadata heuristic
        # terminates for HY-V4's TP1 H64 shape. The Triton sparse MLA kernel
        # never reads the persistent metadata either.
        self._use_persistent_metadata = all(
            getattr(attention_context[name].impl, "sinks", None) is None
            and not getattr(attention_context[name].impl, "use_aiter_sparse_mla", False)
            for name in layer_names
        )
        # Bounds the KV-split heuristic (see `_sparse_decode_max_split`).
        self._num_compute_units = current_platform.num_compute_units()
        self.max_model_len_tensor = torch.tensor(
            [self.model_config.max_model_len], device=device, dtype=torch.int32
        )
        # this is ignored by `flash_mla_with_kvcache` if indices not None
        self.dummy_block_table = torch.empty(
            (1, 1), dtype=torch.int32, device=self.device
        )

        self.req_id_per_token_buffer = torch.zeros(
            (vllm_config.scheduler_config.max_num_batched_tokens,),
            dtype=torch.int32,
            device=device,
        )
        self.qo_indptr = torch.arange(
            0, max_num_batched_tokens + 1, dtype=torch.int32, device=device
        )
        self.paged_kv_last_page_len = torch.ones(
            max_num_batched_tokens, dtype=torch.int32, device=device
        )

        # These two needs to be calculated in runtime,
        # but we still needs to prepare the buffer
        self.paged_kv_indices = torch.zeros(
            [max_num_batched_tokens * self.topk_tokens],
            dtype=torch.int32,
            device=device,
        )
        self.paged_kv_indptr = torch.zeros(
            [max_num_batched_tokens + 1], dtype=torch.int32, device=device
        )

        # ----- Persistent MLA metadata buffers -----
        # The aiter sparse decode kernel supports a "persistent" path that
        # uses precomputed work-splitting metadata for better load balancing
        # across CUs. Mirrors the approach used in rocm_aiter_mla.py.
        #
        # In the sparse case each query token is its own "batch" entry in the
        # qo_indptr (qo_indptr = [0, 1, 2, ..., num_tokens]) and max_qo_len=1.
        # We pad get_mla_metadata_info_v1's batch_size to max_num_batched_tokens
        # so the buffers are large enough for any decode shape we might see.
        from aiter import dtypes, get_mla_metadata_info_v1

        # Keep metadata sizing consistent with the padded tensor shape passed
        # to the sparse decode kernel.
        self._num_attention_heads = AiterMLAHelper.get_actual_mla_num_heads(
            self.num_heads
        )

        q_dtype = self.model_dtype
        kv_cache_dtype_str = getattr(vllm_config.cache_config, "cache_dtype", "auto")
        if kv_cache_dtype_str in ("fp8", "fp8_e4m3", "fp8_e5m2"):
            kv_cache_dtype_str = "fp8"
        else:
            kv_cache_dtype_str = "bf16"
        kv_dtype = dtypes.d_dtypes.get(kv_cache_dtype_str, dtypes.bf16)

        (
            (work_meta_data_size, work_meta_data_type),
            (work_indptr_size, work_indptr_type),
            (work_info_set_size, work_info_set_type),
            (reduce_indptr_size, reduce_indptr_type),
            (reduce_final_map_size, reduce_final_map_type),
            (reduce_partial_map_size, reduce_partial_map_type),
        ) = get_mla_metadata_info_v1(
            max_num_batched_tokens,
            1,
            self._num_attention_heads,
            q_dtype,
            kv_dtype,
            is_sparse=True,
            fast_mode=True,
        )
        self._mla_work_meta_data = torch.empty(
            work_meta_data_size, dtype=work_meta_data_type, device=device
        )
        self._mla_work_indptr = torch.empty(
            work_indptr_size, dtype=work_indptr_type, device=device
        )
        self._mla_work_info_set = torch.empty(
            work_info_set_size, dtype=work_info_set_type, device=device
        )
        self._mla_reduce_indptr = torch.empty(
            reduce_indptr_size, dtype=reduce_indptr_type, device=device
        )
        self._mla_reduce_final_map = torch.empty(
            reduce_final_map_size, dtype=reduce_final_map_type, device=device
        )
        self._mla_reduce_partial_map = torch.empty(
            reduce_partial_map_size,
            dtype=reduce_partial_map_type,
            device=device,
        )

        self._prev_req_extent: int = 0
        self._prev_indices_extent: int = 0
        self._prev_metadata_key: tuple | None = None

    def _sparse_decode_max_split(self, max_seq_len: int) -> int:
        """Cap ``max_split_per_batch`` for the aiter sparse-MLA decode reduce.

        The reduce only covers the selected tokens per row (``<= topk_tokens``),
        so aiter's default (``-1`` => split across every CU) over-fragments it.
        Mirror ``triton_mla.py``: aim for a minimum work per split, round to a
        power of two, and cap by the CU count. Numerics are unchanged.
        """
        effective_len = min(max_seq_len, self.topk_tokens)
        min_work_per_split = 128
        ideal_splits = triton.next_power_of_2(
            max(1, effective_len // min_work_per_split)
        )
        return min(ideal_splits, self._num_compute_units)

    def build(
        self,
        common_prefix_len: int,
        common_attn_metadata: CommonAttentionMetadata,
        fast_build: bool = False,
    ) -> ROCMAiterMLASparseMetadata:
        num_tokens = common_attn_metadata.num_actual_tokens
        (num_decodes, num_prefills, num_decode_tokens, _) = split_decodes_and_prefills(
            common_attn_metadata,
            decode_threshold=self.reorder_batch_threshold or 1,
        )
        starts = np.asarray(common_attn_metadata.query_start_loc_cpu, dtype=np.int32)
        seg_lengths = np.diff(starts)
        req_id_per_token = np.repeat(
            np.arange(seg_lengths.shape[0], dtype=np.int32), seg_lengths
        )
        # Only re-zero the shrink-tail. paged_kv_indptr is fully rewritten
        # by the cumsum below. paged_kv_indices entries past new_indices_extent
        # are never read (the attention kernel only touches the ranges
        # defined by paged_kv_indptr).
        new_req_extent = int(req_id_per_token.shape[0])
        new_indices_extent = num_tokens * self.topk_tokens
        if self._prev_req_extent > new_req_extent:
            self.req_id_per_token_buffer[new_req_extent : self._prev_req_extent].fill_(
                0
            )
        if self._prev_indices_extent > new_indices_extent:
            self.paged_kv_indices[new_indices_extent : self._prev_indices_extent].fill_(
                0
            )
        self._prev_req_extent = new_req_extent
        self._prev_indices_extent = new_indices_extent
        self.req_id_per_token_buffer[:new_req_extent].copy_(
            np_to_pinned_tensor(req_id_per_token), non_blocking=True
        )
        query_lens = (
            common_attn_metadata.query_start_loc[1:]
            - common_attn_metadata.query_start_loc[:-1]
        )
        seq_lens = common_attn_metadata.seq_lens
        sparse_seqlen = generate_sparse_seqlen_triton(
            query_lens,
            seq_lens,
            common_attn_metadata.query_start_loc,
            self.topk_tokens,
            num_tokens,
            common_attn_metadata.max_query_len,
        )

        torch.cumsum(sparse_seqlen, dim=0, out=self.paged_kv_indptr[1 : num_tokens + 1])
        self.paged_kv_indptr[num_tokens + 1 :].fill_(self.paged_kv_indptr[num_tokens])

        req_id_per_token = self.req_id_per_token_buffer[:num_tokens]
        qo_indptr = self.qo_indptr[: num_tokens + 1]
        paged_kv_last_page_len = self.paged_kv_last_page_len[:num_tokens]
        paged_kv_indptr = self.paged_kv_indptr[: num_tokens + 1]
        paged_kv_indices = self.paged_kv_indices[: num_tokens * self.topk_tokens]

        # ----- Compute persistent MLA metadata -----
        # The AITER sparse decode kernel uses qseqlen=1 (each query token is
        # treated as its own batch entry). Build its persistent work metadata
        # only when AITER is selected and no layer needs the nonpersistent LSE
        # path for attention sinks. The output is a deterministic function of
        # the per-request query and context lengths (both clamped to
        # topk_tokens, past which per-token KV length saturates) and num_heads;
        # fingerprint those CPU-side and skip the launch when nothing changed.
        head_size = self.mla_dims.kv_lora_rank + self.mla_dims.qk_rope_head_dim
        use_triton_sparse = _use_rocm_sparse_triton(
            kv_cache_dtype=self.kv_cache_dtype,
            head_size=head_size,
            kv_lora_rank=self.mla_dims.kv_lora_rank,
            num_prefills=num_prefills,
            num_decodes=num_decodes,
            num_decode_tokens=num_decode_tokens,
            max_query_len=common_attn_metadata.max_query_len,
        )
        work_meta_data = None
        work_indptr = None
        work_info_set = None
        reduce_indptr = None
        reduce_final_map = None
        reduce_partial_map = None
        if self._use_persistent_metadata and not use_triton_sparse:
            num_reqs = common_attn_metadata.num_reqs
            with gpu_sync_allowed():
                seq_lens_cpu = common_attn_metadata.seq_lens[:num_reqs].cpu().numpy()
            clamped_seq_lens = np.minimum(
                seq_lens_cpu,
                self.topk_tokens,
            )
            clamped_context_lens = np.minimum(
                seq_lens_cpu - seg_lengths,
                self.topk_tokens,
            )
            metadata_key = (
                num_tokens,
                int(common_attn_metadata.max_query_len),
                self._num_attention_heads,
                clamped_seq_lens.tobytes(),
                clamped_context_lens.tobytes(),
                seg_lengths.tobytes(),
            )
            if metadata_key != self._prev_metadata_key:
                from aiter import get_mla_metadata_v1

                max_split_per_batch = self._sparse_decode_max_split(
                    int(common_attn_metadata.max_seq_len)
                )
                get_mla_metadata_v1(
                    qo_indptr,
                    paged_kv_indptr,
                    paged_kv_last_page_len,
                    self._num_attention_heads,
                    1,
                    True,
                    self._mla_work_meta_data,
                    self._mla_work_info_set,
                    self._mla_work_indptr,
                    self._mla_reduce_indptr,
                    self._mla_reduce_final_map,
                    self._mla_reduce_partial_map,
                    page_size=1,
                    kv_granularity=16,
                    max_seqlen_qo=1,
                    uni_seqlen_qo=1,
                    fast_mode=True,
                    max_split_per_batch=max_split_per_batch,
                )
                torch.cuda.current_stream(self.device).synchronize()
                self._prev_metadata_key = metadata_key
            work_meta_data = self._mla_work_meta_data
            work_indptr = self._mla_work_indptr
            work_info_set = self._mla_work_info_set
            reduce_indptr = self._mla_reduce_indptr
            reduce_final_map = self._mla_reduce_final_map
            reduce_partial_map = self._mla_reduce_partial_map

        metadata = ROCMAiterMLASparseMetadata(
            num_reqs=common_attn_metadata.num_reqs,
            max_query_len=common_attn_metadata.max_query_len,
            max_seq_len=common_attn_metadata.max_seq_len,
            num_actual_tokens=common_attn_metadata.num_actual_tokens,
            query_start_loc=common_attn_metadata.query_start_loc,
            slot_mapping=common_attn_metadata.slot_mapping,
            block_table=common_attn_metadata.block_table_tensor,
            req_id_per_token=req_id_per_token,
            block_size=self.kv_cache_spec.block_size,
            attn_out_dtype=self.model_dtype,
            topk_tokens=self.topk_tokens,
            num_decodes=num_decodes,
            num_prefills=num_prefills,
            num_decode_tokens=num_decode_tokens,
            qo_indptr=qo_indptr,
            paged_kv_last_page_len=paged_kv_last_page_len,
            paged_kv_indices=paged_kv_indices,
            paged_kv_indptr=paged_kv_indptr,
            work_meta_data=work_meta_data,
            work_indptr=work_indptr,
            work_info_set=work_info_set,
            reduce_indptr=reduce_indptr,
            reduce_final_map=reduce_final_map,
            reduce_partial_map=reduce_partial_map,
        )
        return metadata

_sparse_decode_max_split(max_seq_len)

Cap max_split_per_batch for the aiter sparse-MLA decode reduce.

The reduce only covers the selected tokens per row (<= topk_tokens), so aiter's default (-1 => split across every CU) over-fragments it. Mirror triton_mla.py: aim for a minimum work per split, round to a power of two, and cap by the CU count. Numerics are unchanged.

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
def _sparse_decode_max_split(self, max_seq_len: int) -> int:
    """Cap ``max_split_per_batch`` for the aiter sparse-MLA decode reduce.

    The reduce only covers the selected tokens per row (``<= topk_tokens``),
    so aiter's default (``-1`` => split across every CU) over-fragments it.
    Mirror ``triton_mla.py``: aim for a minimum work per split, round to a
    power of two, and cap by the CU count. Numerics are unchanged.
    """
    effective_len = min(max_seq_len, self.topk_tokens)
    min_work_per_split = 128
    ideal_splits = triton.next_power_of_2(
        max(1, effective_len // min_work_per_split)
    )
    return min(ideal_splits, self._num_compute_units)

_use_rocm_sparse_triton(*, kv_cache_dtype, head_size, kv_lora_rank, num_prefills, num_decodes, num_decode_tokens, max_query_len)

Select the rope-free BF16 path not supported by AITER sparse MLA.

The ragged Triton kernel indexes metadata per query token, so multi-token speculative verification rows have the same capability requirements as plain decode rows.

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
def _use_rocm_sparse_triton(
    *,
    kv_cache_dtype: str,
    head_size: int,
    kv_lora_rank: int,
    num_prefills: int,
    num_decodes: int,
    num_decode_tokens: int,
    max_query_len: int,
) -> bool:
    """Select the rope-free BF16 path not supported by AITER sparse MLA.

    The ragged Triton kernel indexes metadata per query token, so multi-token
    speculative verification rows have the same capability requirements as
    plain decode rows.
    """
    return (
        not kv_cache_dtype.startswith("fp8")
        and head_size == kv_lora_rank
        and (num_prefills > 0 or num_decodes > 0)
    )

fit_kpool_indices_to_aiter(token_indices, topk_tokens)

Keep the live kpool tail while fitting AITER's fixed top-k width.

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
def fit_kpool_indices_to_aiter(
    token_indices: torch.Tensor, topk_tokens: int
) -> torch.Tensor:
    """Keep the live kpool tail while fitting AITER's fixed top-k width."""
    if token_indices.shape[1] < topk_tokens:
        raise ValueError("token_indices width must be at least topk_tokens")
    if token_indices.shape[1] == topk_tokens:
        return token_indices

    num_tokens, width = token_indices.shape
    tail_width = width - topk_tokens
    output = torch.empty(
        (num_tokens, topk_tokens),
        dtype=token_indices.dtype,
        device=token_indices.device,
    )
    if num_tokens == 0:
        return output

    _fit_kpool_indices_kernel[(num_tokens,)](
        token_indices,
        output,
        token_indices.stride(0),
        token_indices.stride(1),
        output.stride(0),
        output.stride(1),
        NUM_TOPK_TOKENS=topk_tokens,
        TAIL_WIDTH=tail_width,
        BLOCK_T=triton.next_power_of_2(topk_tokens),
        BLOCK_TAIL=triton.next_power_of_2(tail_width),
    )
    return output

triton_convert_req_index_to_global_index(req_id, block_table, token_indices, cu_seqlens, paged_kv_indices, BLOCK_SIZE=64, NUM_TOPK_TOKENS=2048, BLOCK_N=128)

out[token_id, indice_id] = block_table[req_id[token_id], token_indices[token_id, indice_id] // BLOCK_SIZE] * BLOCK_SIZE + token_indices[token_id, indice_id] % BLOCK_SIZE

Only when token_indices[token_id, indice_id] == -1 do we output -1. For safety, we also output -1 if the derived block_id would be out-of-bounds.

Source code in vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py
def triton_convert_req_index_to_global_index(
    req_id: torch.Tensor,  # int32 [num_tokens]
    block_table: torch.Tensor,  # int32 [num_requests, max_num_blocks_per_req]
    token_indices: torch.Tensor,  # int32 [num_tokens, NUM_TOPK_TOKENS]
    cu_seqlens: torch.Tensor,  # int32 [num_tokens + 1]
    paged_kv_indices: torch.Tensor,  # int32 [num_tokens * topk] out_buffer
    BLOCK_SIZE: int = 64,
    NUM_TOPK_TOKENS: int = 2048,
    BLOCK_N: int = 128,  # tile width along columns
):
    """out[token_id, indice_id] =
        block_table[req_id[token_id],
            token_indices[token_id, indice_id] // BLOCK_SIZE] * BLOCK_SIZE
        + token_indices[token_id, indice_id] % BLOCK_SIZE

    Only when token_indices[token_id, indice_id] == -1 do we output -1.
    For safety, we also output -1 if the derived block_id would be
        out-of-bounds.
    """
    assert req_id.dtype == torch.int32
    assert block_table.dtype == torch.int32
    assert token_indices.dtype == torch.int32
    assert token_indices.shape[1] == NUM_TOPK_TOKENS
    assert NUM_TOPK_TOKENS % BLOCK_N == 0, (
        f"NUM_TOPK_TOKENS ({NUM_TOPK_TOKENS}) must be divisible byBLOCK_N ({BLOCK_N})"
    )
    # print("req_id: ", req_id, flush=True)
    num_tokens = req_id.shape[0]
    _, max_num_blocks_per_req = block_table.shape
    tiles_per_row = NUM_TOPK_TOKENS // BLOCK_N

    # Ensure contiguous tensors on the same device
    req_id_c = req_id.contiguous()
    block_table_c = block_table.contiguous()
    token_indices_c = token_indices.contiguous()

    # Strides in elements
    bt_stride0, bt_stride1 = block_table_c.stride()
    ti_stride0, ti_stride1 = token_indices_c.stride()

    # Exact 2D grid: tokens × column tiles
    grid = (num_tokens, tiles_per_row)

    _convert_req_index_to_global_index_kernel[grid](
        req_id_c,
        block_table_c,
        token_indices_c,
        cu_seqlens,
        paged_kv_indices,
        # shapes / constexprs
        max_num_blocks_per_req,
        BLOCK_SIZE,
        BLOCK_N,
        # strides
        bt_stride0,
        bt_stride1,
        ti_stride0,
        ti_stride1,
    )
    return