Skip to content

vllm.models.deepseek_v41.amd.rocm

Classes:

Functions:

DeepseekV41ROCMAiterMLAAttention

Bases: DeepseekV4Attention

ROCm sparse MLA attention layer for DeepSeek V4.1.

Source code in vllm/models/deepseek_v41/amd/rocm.py
 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
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
1381
1382
1383
1384
1385
1386
1387
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416
1417
1418
1419
1420
1421
1422
1423
1424
1425
1426
1427
1428
1429
1430
1431
1432
1433
1434
1435
1436
1437
1438
1439
1440
1441
1442
1443
1444
1445
1446
1447
1448
1449
1450
1451
1452
1453
1454
1455
1456
1457
1458
1459
1460
1461
1462
1463
1464
1465
1466
1467
1468
1469
1470
1471
1472
1473
1474
1475
1476
1477
1478
1479
1480
1481
1482
1483
1484
1485
1486
1487
1488
1489
1490
1491
1492
1493
1494
1495
1496
1497
1498
1499
1500
1501
1502
1503
1504
1505
1506
1507
1508
1509
1510
1511
1512
1513
1514
1515
1516
1517
1518
1519
1520
1521
1522
1523
1524
1525
1526
1527
1528
1529
1530
1531
1532
1533
1534
1535
1536
1537
1538
1539
1540
1541
1542
1543
1544
1545
1546
1547
1548
1549
1550
1551
1552
1553
1554
1555
1556
1557
1558
1559
1560
1561
1562
1563
1564
1565
1566
1567
1568
1569
1570
1571
1572
1573
1574
1575
1576
1577
1578
1579
1580
1581
1582
1583
1584
1585
1586
1587
1588
1589
1590
1591
1592
1593
1594
1595
1596
1597
1598
1599
1600
1601
1602
1603
1604
1605
1606
1607
1608
1609
1610
1611
1612
1613
1614
1615
1616
1617
1618
1619
1620
1621
1622
1623
1624
1625
1626
1627
1628
1629
1630
1631
1632
1633
1634
1635
1636
1637
1638
1639
1640
1641
1642
1643
1644
1645
1646
1647
1648
1649
class DeepseekV41ROCMAiterMLAAttention(DeepseekV4Attention):
    """ROCm sparse MLA attention layer for DeepSeek V4.1."""

    backend_cls = DeepseekV4ROCMAiterMLASparseBackend
    swa_backend_cls = DeepseekV41ROCMAiterSparseSWABackend
    _use_aiter_sparse_mla = False

    def _indexer_cls(
        self, k_cache: DeepseekV4IndexerCache | None
    ) -> type[DeepseekV4Indexer]:
        if k_cache is not None and k_cache.rocm_mxfp4:
            return DeepseekV41RocmMxfp4Indexer
        return super()._indexer_cls(k_cache)

    def __init__(self, *args, **kwargs):
        vllm_config = args[0] if args else kwargs["vllm_config"]
        super().__init__(*args, **kwargs)
        # CUDA executes WO_A with a quantized grouped-BMM kernel.  ROCm's
        # correctness path below dequantizes WO_A once and uses torch.einsum,
        # so retain the ordinary MXFP8 linear kernel for post-load processing
        # instead of asking for the CUDA-only BMM kernel.
        self.wo_a.is_bmm = False
        self._has_kv_transfer = vllm_config.kv_transfer_config is not None
        # Block scale for the preshuffled weight; None = not preshuffled.
        self._wqa_wkv_scale: torch.Tensor | None = None
        self._wo_b_scale: torch.Tensor | None = None
        self._fused_compressor_weight: torch.Tensor | None
        self.register_buffer("_fused_compressor_weight", None, persistent=False)
        self._fused_compressor_split_sizes: tuple[int, int] | None = None
        # Decode ragged topk metadata is a pure function of the indices its
        # index source published, so every consumer of a source rebuilds the
        # same thing. The source owns one cache per compress ratio, refreshed
        # when it runs; consumers below it read through.
        self._topk_ragged_cache: dict[int, _TopkRagged] = {}
        self._prefill_topk_ragged_cache: dict[
            str, tuple[torch.Tensor, torch.Tensor]
        ] = {}
        self._index_source_prefix: str | None = None
        if self.compress_ratio > 0:
            assert self.index_source_layer_id is not None
            self._index_source_prefix = _replace_layer_index(
                self.prefix, self.index_source_layer_id
            )
            if self._index_source_prefix not in self._static_forward_context:
                raise NotImplementedError(
                    f"Index source {self._index_source_prefix} not found on "
                    "this rank; PP splits inside a v4.1 index-sharing group "
                    "are not supported."
                )
        self._use_aiter_sparse_mla = _aiter_sparse_mla_enabled(self._has_kv_transfer)
        if self.compressor is None and self.indexer is None:
            # Dense layers have nothing to overlap; keep the base serial path.
            self.aux_stream_list = None

    @classmethod
    def get_padded_num_q_heads(cls, num_heads: int) -> int:
        return num_heads

    def prepare_attn_preshuffle(self) -> None:
        from vllm._aiter_ops import rocm_aiter_ops

        if not rocm_aiter_ops.is_enabled():
            return
        from vllm.model_executor.layers.quantization.utils.fp8_utils import (
            _upcast_e8m0_to_fp32,
        )
        from vllm.model_executor.utils import replace_parameter

        def _prep(linear) -> torch.Tensor | None:
            w = getattr(linear, "weight", None)
            if w is None or w.dim() != 2:
                return None
            # K % 128 (group-128 quant) and N % 16 (shuffle_weight) must hold.
            if w.shape[-1] % 128 != 0 or w.shape[0] % 16 != 0:
                return None
            ws = getattr(linear, "weight_scale_inv", None)  # per-block scale
            if ws is None:
                return None
            if ws.dtype == torch.float8_e8m0fnu:
                ws = _upcast_e8m0_to_fp32(ws).contiguous()
            # Shuffle the weight in place (single weight, no unshuffled copy).
            replace_parameter(
                linear,
                "weight",
                rocm_aiter_ops.shuffle_weight(w.data, layout=(16, 16)),
            )
            return ws

        self._wqa_wkv_scale = _prep(self.fused_wqa_wkv)
        self._wo_b_scale = _prep(self.wo_b)

    def prepare_compressor_gemm_fusion(self) -> bool:
        # V4.1 derives index keys from the source compressor's emitted latent
        # and has no nested ``indexer.compressor``.  Keep the projections
        # separate and use the shared linear/PyTorch correctness path.
        return False

    def _bpre_attn_gemm(
        self,
        weight: torch.Tensor,
        scale: torch.Tensor,
        x: torch.Tensor,
        reduce_tp: bool,
    ) -> torch.Tensor:
        from vllm._aiter_ops import rocm_aiter_ops

        x_fp8, x_scale = rocm_aiter_ops.group_fp8_quant(x, transpose_scale=True)
        out = rocm_aiter_ops.gemm_a8w8_blockscale_bpreshuffle(
            x_fp8, weight, x_scale, scale, output_dtype=x.dtype
        )
        if reduce_tp and get_tensor_model_parallel_world_size() > 1:
            out = tensor_model_parallel_all_reduce(out)
        return out

    def _fused_wqa_wkv_gemm(self, hidden_states: torch.Tensor) -> torch.Tensor:
        if self._wqa_wkv_scale is not None and hidden_states.dim() == 2:
            return self._bpre_attn_gemm(
                self.fused_wqa_wkv.weight, self._wqa_wkv_scale, hidden_states, False
            )
        return super()._fused_wqa_wkv_gemm(hidden_states)

    def _run_parallel_input_projections(
        self, hidden_states: torch.Tensor
    ) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
        if self.indexer is not None and self.compressor is None:
            # Serial: qr must sit on the default stream before the q-side
            # fork below can consume it (ROCm).
            aux_streams = self.aux_stream_list
            self.aux_stream_list = None
            try:
                return super()._run_parallel_input_projections(hidden_states)
            finally:
                self.aux_stream_list = aux_streams
        return super()._run_parallel_input_projections(hidden_states)

    def forward(
        self,
        positions: torch.Tensor,
        hidden_states: torch.Tensor,
        llama_4_scaling: torch.Tensor | None = None,
    ) -> torch.Tensor:
        if (
            self.indexer is None
            or self._prepare_and_attn_fn == self._prepare_and_attn_eager
            or not current_platform.enable_multi_stream_overlap(
                self.aux_stream_list, get_forward_context().attn_metadata
            )
        ):
            # Sequential fallback: no forks outside the capture region, where
            # HIP event sync is unreliable; MRV1 keeps the input prep in its
            # wide eager region (#51430).
            aux_streams = self.aux_stream_list
            self.aux_stream_list = None
            try:
                return super().forward(positions, hidden_states, llama_4_scaling)
            finally:
                self.aux_stream_list = aux_streams

        attn_out = self._alloc_attn_out(hidden_states.shape[0], hidden_states)
        if self.compressor is not None:
            q, kv, index_q, index_q_scale, index_weights = self._forward_csa2_full(
                hidden_states, positions
            )
        else:
            q, kv, index_q, index_q_scale, index_weights = self._forward_csa2_reindex(
                hidden_states, positions
            )
        self._sparse_indexer_and_attn(
            hidden_states,
            index_q,
            index_q_scale,
            index_weights,
            q,
            kv,
            positions,
            attn_out,
        )
        return self._o_proj(attn_out, positions)

    def _forward_csa2_full(
        self, hidden_states: torch.Tensor, positions: torch.Tensor
    ) -> tuple[
        torch.Tensor,
        torch.Tensor,
        torch.Tensor | None,
        torch.Tensor | None,
        torch.Tensor | None,
    ]:
        """CSA layers: one fork replaces the base pipeline's three.

        Each fork/join is a HIP event pair per side stream on ROCm, so the
        base three-fork pipeline is expensive there; the compressor chain and
        the indexer weights projection run fully in parallel with the default
        chain. The indexer K write and q-side stay serial after the join:
        ``wk`` consumes the aux-0 latent and the q-side's wq_b consumes the
        default-stream qr.
        """
        attn_metadata = get_forward_context().attn_metadata
        compressor = self.compressor
        indexer = self.indexer
        assert compressor is not None and indexer is not None
        aux_streams = self.aux_stream_list
        assert aux_streams is not None and len(aux_streams) >= 2

        def default_chain():
            qr_kv = self._fused_wqa_wkv_gemm(hidden_states)
            qr, qr_scale, kv = self._split_qkv_and_norm(qr_kv)
            q = self._wq_b_proj(qr, qr_scale).view(
                -1, self.n_local_heads, self.head_dim
            )
            return (
                self._fused_qnorm_rope_kv_insert(q, kv, positions, attn_metadata),
                qr,
                qr_scale,
                kv,
            )

        def compressor_chain():
            kv_score = torch.mm(
                hidden_states,
                compressor.fused_wkv_wgate.weight.T,
                out_dtype=torch.float32,
            )
            latent = compressor(kv_score, positions)
            compressor.insert_cache(latent, positions, self.rotary_emb)
            return latent

        def indexer_weights_chain():
            # ReplicatedLinear returns (output, bias); bias is None.
            weights, _ = indexer.weights_proj(hidden_states)
            return weights

        (q, qr, qr_scale, kv), (latent, indexer_weights) = execute_in_parallel(
            default_chain,
            [compressor_chain, indexer_weights_chain],
            self.ln_events[0],
            self.ln_events[1:3],
            aux_streams[:2],
            enable=True,
            default_first=True,
        )

        indexer._produce_k(latent, positions, self.indexer_rotary_emb)
        index_q, index_q_scale, index_weights_out = indexer.forward_q(
            qr, qr_scale, indexer_weights, positions, self.indexer_rotary_emb
        )
        return q, kv, index_q, index_q_scale, index_weights_out

    def _forward_csa2_reindex(
        self, hidden_states: torch.Tensor, positions: torch.Tensor
    ) -> tuple[
        torch.Tensor,
        torch.Tensor,
        torch.Tensor | None,
        torch.Tensor | None,
        torch.Tensor | None,
    ]:
        """Indexer-only layers: the q-side overlaps the SWA q path.

        The input projections run serial first: the q-side consumes qr, so
        it can only overlap the SWA q projection and KV insert.
        """
        indexer = self.indexer
        assert indexer is not None
        attn_metadata = get_forward_context().attn_metadata
        # Serial on these layers via the _run_parallel_input_projections
        # override: qr must sit on the default stream before the fork.
        qr_kv, _, indexer_weights = self._run_parallel_input_projections(hidden_states)
        assert indexer_weights is not None
        qr, qr_scale, kv = self._split_qkv_and_norm(qr_kv)

        def swa_q_chain():
            q = self._wq_b_proj(qr, qr_scale).view(
                -1, self.n_local_heads, self.head_dim
            )
            return self._fused_qnorm_rope_kv_insert(q, kv, positions, attn_metadata)

        def indexer_q_chain():
            return indexer.forward_q(
                qr, qr_scale, indexer_weights, positions, self.indexer_rotary_emb
            )

        aux_streams = self.aux_stream_list
        assert aux_streams is not None
        q, aux_results = execute_in_parallel(
            swa_q_chain,
            [indexer_q_chain],
            self.ln_events[0],
            self.ln_events[1:2],
            aux_streams[:1],
            enable=True,
            default_first=True,
        )
        index_q, index_q_scale, index_weights_out = aux_results[0]
        return q, kv, index_q, index_q_scale, index_weights_out

    @functools.cached_property
    def _wq_b_uses_aiter_block_scaled(self) -> bool:
        """True when both wq_b GEMMs run the aiter block-scaled fp8 kernel.

        Cached: the linear kernels and the aiter env gates are fixed once
        the model is built, so this is evaluated at the first forward
        only.

        The fused norm+quant path is only valid if the quant and GEMM it
        replaces are exactly the aiter ones; otherwise fall back to the
        shared path.
        """
        from vllm._aiter_ops import rocm_aiter_ops
        from vllm.model_executor.kernels.linear.scaled_mm import (
            Fp8BlockScaledMMLinearKernel,
        )

        if not rocm_aiter_ops.is_linear_fp8_enabled():
            return False

        linears = [self.wq_b]
        if self.indexer is not None:
            linears.append(self.indexer.wq_b)
        for linear in linears:
            kernel = getattr(getattr(linear, "quant_method", None), "fp8_linear", None)
            if not isinstance(kernel, Fp8BlockScaledMMLinearKernel):
                return False
        return True

    def _split_qkv_and_norm(
        self, qr_kv: torch.Tensor
    ) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor]:
        """Fuse q/kv RMSNorm + per-1x128 fp8 q quant into one aiter kernel.

        The shared path norms q and kv in one triton kernel and the wq_b
        linears then re-read the bf16 qr to quantize it. The aiter kernel
        computes both RMSNorms (fp32 accumulate) and the fp8 group quant
        in a single pass, writing fp8 qr + group scales directly; both
        wq_b GEMMs (attention and indexer) then consume that pair and
        skip their own input quant. kv stays bf16: the fused insert
        kernel RoPE/quantizes it itself. Falls back to the shared path
        when the aiter linear path is not active.
        """
        qr, kv = qr_kv.split([self.q_lora_rank, self.head_dim], dim=-1)
        if not (
            qr.dim() == 2
            and qr.shape[0] > 0
            and self.q_lora_rank % 128 == 0
            and self._wq_b_uses_aiter_block_scaled
        ):
            return super()._split_qkv_and_norm(qr_kv)

        from vllm._aiter_ops import rocm_aiter_ops

        return rocm_aiter_ops.fused_qk_rmsnorm_group_quant(
            q=qr,
            q_weight=self.q_norm.weight.data,
            q_epsilon=self.eps,
            kv=kv,
            kv_weight=self.kv_norm.weight.data,
            kv_epsilon=self.eps,
            group_size=128,
            transpose_scale=False,
        )

    def _alloc_attn_out(
        self, num_tokens: int, hidden_states: torch.Tensor
    ) -> torch.Tensor | QuantizedActivation:
        if not _ON_GFX950:
            return super()._alloc_attn_out(num_tokens, hidden_states)
        # wo_a's MXFP8 input: the decode reduce writes it directly, prefill
        # rows are rotated and quantized after their bf16 attention.
        width = self.n_local_heads * self.head_dim
        data = torch.empty(
            (num_tokens, width), dtype=torch.float8_e4m3fn, device=hidden_states.device
        )
        scale = torch.empty(
            (num_tokens, width // 32), dtype=torch.uint8, device=hidden_states.device
        )
        return QuantizedActivation(
            data=data,
            scale=scale,
            orig_dtype=hidden_states.dtype,
            orig_shape=data.shape,
            quant_key=kMxfp8Dynamic,
        )

    def _o_proj(
        self, attn_out: torch.Tensor | QuantizedActivation, positions: torch.Tensor
    ) -> torch.Tensor:
        if isinstance(attn_out, QuantizedActivation):
            z = rocm_mxfp8_wo_a_bmm(
                attn_out.data,
                attn_out.scale,
                self.wo_a,
                self.n_local_groups,
                self.o_lora_rank,
            )
            return self._wo_b_after_wo_a(z)
        o = attn_out[:, : self.n_local_heads, :]
        # ROCm BF16 reference wo_a path (inverse RoPE + einsum) + wo_b.
        z = rocm_inv_rope_einsum(
            self.rotary_emb,
            o,
            positions,
            self.rope_head_dim,
            self.n_local_groups,
            self.o_lora_rank,
            self.wo_a,
            inverse_rope=False,
        )
        return self._wo_b_after_wo_a(z.flatten(1))

    def _wo_b_after_wo_a(self, zf: torch.Tensor) -> torch.Tensor:
        if self._wo_b_scale is not None and zf.dim() == 2:
            return self._bpre_attn_gemm(self.wo_b.weight, self._wo_b_scale, zf, True)
        return self.wo_b(zf)

    def forward_mqa(
        self,
        q: torch.Tensor,
        kv: torch.Tensor,
        positions: torch.Tensor,
        output: torch.Tensor | QuantizedActivation,
    ) -> None:
        mxfp8_out = output if isinstance(output, QuantizedActivation) else None
        if mxfp8_out is None:
            assert isinstance(output, torch.Tensor)
            assert output.shape == q.shape, (
                f"output buffer shape {output.shape} must match q shape {q.shape}"
            )
            assert output.dtype == q.dtype, (
                f"output buffer dtype {output.dtype} must match q dtype {q.dtype}"
            )

        forward_context = get_forward_context()
        attn_metadata = forward_context.attn_metadata

        if attn_metadata is None:
            # Warmup dummy run: no real metadata. Reserve the workspace
            # _forward_prefill would; the dequantize / topk / sparse_fwd kernels
            # are skipped this step. The aiter prefill reads the caches in place,
            # so it asks only for the bf16 rows of the MXFP8 output.
            swa_only = self.compress_ratio == 0
            N = (
                0
                if swa_only
                else (self.max_model_len + self.compress_ratio - 1)
                // self.compress_ratio
            )
            M = N + self.window_size + self.max_num_batched_tokens
            shapes = self._prefill_workspace_shapes(
                M,
                self.max_num_batched_tokens,
                q,
                gather=not self._use_aiter_sparse_mla,
            )
            if shapes:
                current_workspace_manager().get_simultaneous(*shapes)
            if mxfp8_out is None:
                output.zero_()
            else:
                mxfp8_out.data.zero_()
                mxfp8_out.scale.zero_()
            return

        assert isinstance(attn_metadata, dict)
        rocm_metadata = cast(
            DeepseekV4FlashMLAMetadata | None,
            attn_metadata.get(self.compressed_cache_prefix)
            if self.compressed_cache_prefix is not None
            else None,
        )
        swa_metadata = cast(
            DeepseekV4ROCMAiterSparseSWAMetadata | None,
            attn_metadata.get(self.swa_cache_layer.prefix),
        )
        assert swa_metadata is not None

        swa_only = self.compress_ratio == 0
        self_kv_cache = None if swa_only else self._compressed_kv_cache()
        swa_kv_cache = self.swa_cache_layer.kv_cache

        num_decodes = swa_metadata.num_decodes
        num_prefills = swa_metadata.num_prefills
        num_decode_tokens = swa_metadata.num_decode_tokens

        if num_prefills > 0:
            self._forward_prefill(
                q=q[num_decode_tokens:],
                positions=positions[num_decode_tokens:],
                compressed_k_cache=self_kv_cache,
                swa_k_cache=swa_kv_cache,
                output=(
                    output[num_decode_tokens:]
                    if isinstance(output, torch.Tensor)
                    else None
                ),
                attn_metadata=rocm_metadata,
                swa_metadata=swa_metadata,
                mxfp8_out=(
                    None
                    if mxfp8_out is None
                    else (
                        mxfp8_out.data[num_decode_tokens:],
                        mxfp8_out.scale[num_decode_tokens:],
                    )
                ),
            )
        if mxfp8_out is not None:
            # The reduce epilogue rotates and quantizes every decode row and
            # the prefill path already did its own, so nothing is owed here.
            if num_decodes > 0:
                self._forward_decode(
                    q=q[:num_decode_tokens],
                    positions=positions[:num_decode_tokens],
                    kv_cache=self_kv_cache,
                    swa_metadata=swa_metadata,
                    attn_metadata=rocm_metadata,
                    swa_only=swa_only,
                    output=None,
                    output_mxfp8=(
                        mxfp8_out.data[:num_decode_tokens],
                        mxfp8_out.scale[:num_decode_tokens],
                    ),
                )
            return
        assert isinstance(output, torch.Tensor)
        rotated = 0
        if num_decodes > 0:
            rotated = self._forward_decode(
                q=q[:num_decode_tokens],
                positions=positions[:num_decode_tokens],
                kv_cache=self_kv_cache,
                swa_metadata=swa_metadata,
                attn_metadata=rocm_metadata,
                swa_only=swa_only,
                output=output[:num_decode_tokens],
            )
        # Only the decode reduce rotates its own rows, and only the leading
        # `rotated` of them; prefill rows and any decode path that did not
        # fuse still owe the standalone pass. Settle that here rather than in
        # _o_proj: the split is batch-dependent and _o_proj runs compiled,
        # where such a value freezes at its trace-time value.
        rocm_inverse_rope_rows_(
            output[rotated:, : self.n_local_heads, :],
            positions[rotated:],
            self.rotary_emb.cos_sin_cache,
            self.rope_head_dim,
        )

    def _decode_topk_ragged(
        self,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
        attn_metadata: DeepseekV4FlashMLAMetadata | None,
        num_decodes: int,
        num_decode_tokens: int,
    ) -> _TopkRagged:
        """Ragged form of the topk indices this layer's index source published.

        The packing depends only on the shared ``topk_indices_buffer`` and
        per-step metadata, apart from the layer's own compress ratio, so the
        source memoizes one result per ratio for all the layers below it.
        """
        assert attn_metadata is not None
        assert swa_metadata.is_valid_token is not None
        assert self.topk_indices_buffer is not None
        assert self._index_source_prefix is not None

        source = self._static_forward_context[self._index_source_prefix]
        if source is self:
            # Fresh indices as of this layer; drop what the last step cached.
            source._topk_ragged_cache = {}
        cached = source._topk_ragged_cache.get(self.compress_ratio)
        if cached is not None:
            return cached

        built = compute_global_topk_ragged_indices_and_indptr(
            self.topk_indices_buffer[:num_decode_tokens],
            swa_metadata.token_to_req_indices,
            attn_metadata.block_table[:num_decodes],
            attn_metadata.block_size // self.compress_ratio,
            swa_metadata.is_valid_token[:num_decode_tokens],
        )
        source._topk_ragged_cache[self.compress_ratio] = built
        return built

    def _forward_decode(
        self,
        q: torch.Tensor,
        positions: torch.Tensor,
        kv_cache: torch.Tensor | None,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
        attn_metadata: DeepseekV4FlashMLAMetadata | None,
        swa_only: bool,
        output: torch.Tensor | None,
        output_mxfp8: tuple[torch.Tensor, torch.Tensor] | None = None,
    ) -> int:
        """Returns how many leading rows the decode epilogue inverse-RoPE'd."""
        num_decodes = swa_metadata.num_decodes
        num_decode_tokens = swa_metadata.num_decode_tokens

        topk_lens = None
        topk_ragged_indices = None
        topk_ragged_indptr = None
        if not swa_only:
            (
                topk_ragged_indices,
                topk_ragged_indptr,
                topk_lens,
            ) = self._decode_topk_ragged(
                swa_metadata=swa_metadata,
                attn_metadata=attn_metadata,
                num_decodes=num_decodes,
                num_decode_tokens=num_decode_tokens,
            )

        if self._use_aiter_sparse_mla:
            assert swa_metadata.decode_swa_ragged_indices is not None
            assert swa_metadata.decode_swa_ragged_indptr is not None
            out = output
            if output_mxfp8 is not None:
                # The aiter kernel has no MXFP8 epilogue: write bf16 rows, then
                # rotate and quantize them as the decode reduce would have.
                out = self._mxfp8_bf16_rows(q.shape[0], q)
            assert out is not None
            self._aiter_sparse_mla(
                q,
                out,
                self.swa_cache_layer.kv_cache,
                swa_metadata.decode_swa_ragged_indices,
                swa_metadata.decode_swa_ragged_indptr,
                kv_cache,
                topk_ragged_indices,
                topk_ragged_indptr,
            )
            if output_mxfp8 is not None:
                rocm_inverse_rope_mxfp8_rows(
                    out,
                    positions,
                    self.rotary_emb.cos_sin_cache,
                    self.rope_head_dim,
                    output_mxfp8[0],
                    output_mxfp8[1],
                )
                return q.shape[0]
            # The aiter kernel leaves the inverse RoPE to the caller.
            return 0

        return rocm_sparse_attn_decode(
            q=q,
            kv_cache=kv_cache,
            swa_k_cache=self.swa_cache_layer.kv_cache,
            swa_only=swa_only,
            topk_indices=None,
            topk_lens=topk_lens,
            swa_indices=swa_metadata.decode_swa_indices,
            swa_lens=swa_metadata.decode_swa_lens,
            swa_ragged_indices=swa_metadata.decode_swa_ragged_indices,
            swa_ragged_indptr=swa_metadata.decode_swa_ragged_indptr,
            topk_ragged_indices=topk_ragged_indices,
            topk_ragged_indptr=topk_ragged_indptr,
            attn_sink=self.attn_sink,
            scale=self.scale,
            head_dim=self.head_dim,
            nope_head_dim=self.nope_head_dim,
            rope_head_dim=self.rope_head_dim,
            output=output,
            output_mxfp8=output_mxfp8,
            inv_rope_positions=positions,
            inv_rope_cos_sin_cache=self.rotary_emb.cos_sin_cache,
            extra_cache_nan_free=_trust_dsv4_extra_cache_nan_free(
                self.kv_cache_dtype,
                self._has_kv_transfer,
                not swa_only and kv_cache is not None,
            ),
        )

    def _forward_prefill(
        self,
        q: torch.Tensor,
        positions: torch.Tensor,
        compressed_k_cache: torch.Tensor | None,
        swa_k_cache: torch.Tensor,
        output: torch.Tensor | None,
        attn_metadata: DeepseekV4FlashMLAMetadata | None,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
        mxfp8_out: tuple[torch.Tensor, torch.Tensor] | None = None,
    ) -> None:
        """Sparse prefill into bf16 ``output``, or into ``mxfp8_out``.

        With ``mxfp8_out`` the rows go to a bf16 workspace first and are then
        inverse-RoPE'd and quantized into it in one pass.
        """
        if self._use_aiter_sparse_mla:
            self._forward_prefill_aiter(
                q=q,
                positions=positions,
                compressed_k_cache=compressed_k_cache,
                swa_k_cache=swa_k_cache,
                output=output,
                attn_metadata=attn_metadata,
                swa_metadata=swa_metadata,
                mxfp8_out=mxfp8_out,
            )
            return

        swa_only = attn_metadata is None

        num_prefill_tokens = swa_metadata.num_prefill_tokens
        num_decodes = swa_metadata.num_decodes
        num_decode_tokens = swa_metadata.num_decode_tokens

        seq_lens = swa_metadata.prefill_seq_lens
        gather_lens = swa_metadata.prefill_gather_lens
        assert seq_lens is not None
        assert gather_lens is not None

        query_start_loc_cpu = swa_metadata.query_start_loc_cpu
        query_start_loc = swa_metadata.query_start_loc
        assert query_start_loc_cpu is not None
        assert query_start_loc is not None
        prefill_token_base = query_start_loc_cpu[num_decodes]

        # Local indices filled by the index source; SWA-only layers pass
        # top_k=0 and never read them.
        assert self.topk_indices_buffer is not None
        topk_indices = self.topk_indices_buffer[num_decode_tokens:]
        topk_indices = topk_indices[:num_prefill_tokens]
        if not swa_only:
            top_k = topk_indices.shape[-1]
            N = (self.max_model_len + self.compress_ratio - 1) // self.compress_ratio
        else:
            top_k = 0
            N = 0

        M = N + self.window_size + self.max_num_batched_tokens
        chunk_plan = swa_metadata.get_prefill_chunk_plan(
            compress_ratio=self.compress_ratio,
            prefill_chunk_size=self.PREFILL_CHUNK_SIZE,
            has_compressed=not swa_only,
        )
        assert chunk_plan, "prefill chunk plan must be non-empty when num_prefills > 0"

        workspace = current_workspace_manager().get_simultaneous(
            *self._prefill_workspace_shapes(M, num_prefill_tokens, q)
        )
        kv_rows = workspace[0].view(-1, q.shape[-1])
        if mxfp8_out is not None:
            output = workspace[1]
        assert output is not None
        for chunk_start, chunk_end, chunk_N, chunk_M in chunk_plan:
            chunk_size = chunk_end - chunk_start
            kv = kv_rows[: chunk_size * chunk_M].view(chunk_size, chunk_M, -1)
            if not swa_only:
                assert attn_metadata is not None
                assert compressed_k_cache is not None
                block_table = attn_metadata.block_table[num_decodes:]
                # compressed_k_cache is OCP on every platform (Triton encoder).
                dequantize_and_gather_k_cache(
                    kv[:chunk_size],
                    compressed_k_cache,
                    seq_lens=seq_lens[chunk_start:chunk_end] // self.compress_ratio,
                    gather_lens=None,
                    block_table=block_table[chunk_start:chunk_end],
                    block_size=attn_metadata.block_size // self.compress_ratio,
                    offset=0,
                    use_fnuz=False,
                )

            swa_block_table = swa_metadata.block_table[num_decodes:]
            dequantize_and_gather_k_cache(
                kv[:chunk_size],
                swa_k_cache,
                seq_lens=seq_lens[chunk_start:chunk_end],
                gather_lens=gather_lens[chunk_start:chunk_end],
                block_table=swa_block_table[chunk_start:chunk_end],
                block_size=swa_metadata.block_size,
                offset=chunk_N,
                use_fnuz=current_platform.is_fp8_fnuz(),
            )

            query_start = (
                query_start_loc_cpu[num_decodes + chunk_start] - prefill_token_base
            )
            query_end = (
                query_start_loc_cpu[num_decodes + chunk_end] - prefill_token_base
            )

            combined_indices, combined_lens = combine_topk_swa_indices(
                topk_indices[query_start:query_end],
                query_start_loc[
                    num_decodes + chunk_start : num_decodes + chunk_end + 1
                ],
                seq_lens[chunk_start:chunk_end],
                gather_lens[chunk_start:chunk_end],
                self.window_size,
                self.compress_ratio,
                top_k,
                chunk_M,
                chunk_N,
            )
            rocm_sparse_attn_prefill(
                q=q[query_start:query_end],
                kv=kv.view(-1, 1, q.shape[-1]),
                indices=combined_indices,
                topk_length=combined_lens,
                scale=self.scale,
                head_dim=self.head_dim,
                nope_head_dim=self.nope_head_dim,
                rope_head_dim=self.rope_head_dim,
                attn_sink=self.attn_sink,
                output=output[query_start:query_end],
            )
        if mxfp8_out is not None:
            rocm_inverse_rope_mxfp8_rows(
                output,
                positions[:num_prefill_tokens],
                self.rotary_emb.cos_sin_cache,
                self.rope_head_dim,
                mxfp8_out[0][:num_prefill_tokens],
                mxfp8_out[1][:num_prefill_tokens],
            )

    def _prefill_workspace_shapes(
        self, M: int, num_prefill_tokens: int, q: torch.Tensor, gather: bool = True
    ) -> list[tuple[tuple[int, ...], torch.dtype]]:
        """The prefill gather buffer, plus bf16 rows for the MXFP8 output.

        One ``get_simultaneous`` call: separate calls alias the same memory.
        ``gather=False`` leaves out the gather buffer, which the aiter kernel,
        reading the caches in place, never needs.
        """
        shapes: list[tuple[tuple[int, ...], torch.dtype]] = []
        if gather:
            shapes.append(((self.PREFILL_CHUNK_SIZE, M, q.shape[-1]), torch.bfloat16))
        if _ON_GFX950:
            shapes.append(
                (
                    (num_prefill_tokens, self.n_local_heads, self.head_dim),
                    torch.bfloat16,
                )
            )
        return shapes

    def _prefill_topk_ragged(
        self,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
        attn_metadata: DeepseekV4FlashMLAMetadata,
        compressed_k_cache: torch.Tensor,
    ) -> tuple[torch.Tensor, torch.Tensor]:
        """Prefill counterpart of ``_decode_topk_ragged``: layers that share an
        index source and a compressed cache reuse one build per step."""
        assert swa_metadata.token_to_req_indices is not None
        assert swa_metadata.query_start_loc is not None
        assert swa_metadata.seq_lens is not None
        assert swa_metadata.is_valid_token is not None
        assert self.topk_indices_buffer is not None
        assert self._index_source_prefix is not None
        assert self.compressed_cache_prefix is not None

        source = self._static_forward_context[self._index_source_prefix]
        if source is self:
            source._prefill_topk_ragged_cache = {}
        cached = source._prefill_topk_ragged_cache.get(self.compressed_cache_prefix)
        if cached is not None:
            return cached

        num_decode_tokens = swa_metadata.num_decode_tokens
        num_prefill_tokens = swa_metadata.num_prefill_tokens
        built = build_prefill_topk_ragged_indices(
            self.topk_indices_buffer[
                num_decode_tokens : num_decode_tokens + num_prefill_tokens
            ],
            swa_metadata.token_to_req_indices,
            swa_metadata.query_start_loc,
            swa_metadata.seq_lens,
            swa_metadata.is_valid_token,
            attn_metadata.block_table,
            block_size=attn_metadata.block_size // self.compress_ratio,
            compress_ratio=self.compress_ratio,
            num_compressed=(self.max_model_len + self.compress_ratio - 1)
            // self.compress_ratio,
            token_offset=num_decode_tokens,
            num_rows=compressed_k_cache.shape[0] * compressed_k_cache.shape[1],
        )
        source._prefill_topk_ragged_cache[self.compressed_cache_prefix] = built
        return built

    def _forward_prefill_aiter(
        self,
        q: torch.Tensor,
        positions: torch.Tensor,
        compressed_k_cache: torch.Tensor | None,
        swa_k_cache: torch.Tensor,
        output: torch.Tensor | None,
        attn_metadata: DeepseekV4FlashMLAMetadata | None,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
        mxfp8_out: tuple[torch.Tensor, torch.Tensor] | None = None,
    ) -> None:
        """Prefill straight off the paged caches, with no bf16 gather.

        With ``mxfp8_out`` the kernel writes bf16 rows to the workspace, which
        are then inverse-RoPE'd and quantized into it as ``_forward_prefill``
        does.
        """
        num_prefill_tokens = swa_metadata.num_prefill_tokens
        assert swa_metadata.prefill_swa_ragged_indices is not None
        assert swa_metadata.prefill_swa_ragged_indptr is not None

        topk_indices = None
        topk_indptr = None
        if attn_metadata is None:
            compressed_k_cache = None
        else:
            assert compressed_k_cache is not None
            topk_indices, topk_indptr = self._prefill_topk_ragged(
                swa_metadata, attn_metadata, compressed_k_cache
            )

        if mxfp8_out is not None:
            output = self._mxfp8_bf16_rows(num_prefill_tokens, q)
        assert output is not None
        self._aiter_sparse_mla(
            q[:num_prefill_tokens],
            output[:num_prefill_tokens],
            swa_k_cache,
            swa_metadata.prefill_swa_ragged_indices,
            swa_metadata.prefill_swa_ragged_indptr,
            compressed_k_cache,
            topk_indices,
            topk_indptr,
        )
        if mxfp8_out is not None:
            rocm_inverse_rope_mxfp8_rows(
                output[:num_prefill_tokens],
                positions[:num_prefill_tokens],
                self.rotary_emb.cos_sin_cache,
                self.rope_head_dim,
                mxfp8_out[0][:num_prefill_tokens],
                mxfp8_out[1][:num_prefill_tokens],
            )

    def _mxfp8_bf16_rows(self, num_tokens: int, q: torch.Tensor) -> torch.Tensor:
        """bf16 attention rows for the aiter kernel to write when the layer's
        output is MXFP8, taken from the workspace the warmup reserved."""
        shapes = self._prefill_workspace_shapes(0, num_tokens, q, gather=False)
        assert shapes, "an MXFP8 attention output is gfx950-only"
        return current_workspace_manager().get_simultaneous(*shapes)[0]

    def _aiter_sparse_mla(
        self,
        q: torch.Tensor,
        output: torch.Tensor,
        swa_k_cache: torch.Tensor,
        swa_indices: torch.Tensor,
        swa_indptr: torch.Tensor,
        compressed_k_cache: torch.Tensor | None,
        topk_indices: torch.Tensor | None,
        topk_indptr: torch.Tensor | None,
    ) -> None:
        from vllm._aiter_ops import rocm_aiter_ops

        # Both fp8 caches are read in place: the SWA window as the main segment,
        # the top-k compressed tokens as the extra one.
        rocm_aiter_ops.triton_sparse_mla_fwd(
            q,
            swa_k_cache,
            output,
            self.scale,
            swa_indptr,
            swa_indices,
            attn_sink=self.attn_sink[: q.shape[1]],
            extra_kv_buffer=compressed_k_cache,
            extra_kv_indptr=topk_indptr,
            extra_kv_indices=topk_indices,
        )

_wq_b_uses_aiter_block_scaled cached property

True when both wq_b GEMMs run the aiter block-scaled fp8 kernel.

Cached: the linear kernels and the aiter env gates are fixed once the model is built, so this is evaluated at the first forward only.

The fused norm+quant path is only valid if the quant and GEMM it replaces are exactly the aiter ones; otherwise fall back to the shared path.

_decode_topk_ragged(swa_metadata, attn_metadata, num_decodes, num_decode_tokens)

Ragged form of the topk indices this layer's index source published.

The packing depends only on the shared topk_indices_buffer and per-step metadata, apart from the layer's own compress ratio, so the source memoizes one result per ratio for all the layers below it.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _decode_topk_ragged(
    self,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    attn_metadata: DeepseekV4FlashMLAMetadata | None,
    num_decodes: int,
    num_decode_tokens: int,
) -> _TopkRagged:
    """Ragged form of the topk indices this layer's index source published.

    The packing depends only on the shared ``topk_indices_buffer`` and
    per-step metadata, apart from the layer's own compress ratio, so the
    source memoizes one result per ratio for all the layers below it.
    """
    assert attn_metadata is not None
    assert swa_metadata.is_valid_token is not None
    assert self.topk_indices_buffer is not None
    assert self._index_source_prefix is not None

    source = self._static_forward_context[self._index_source_prefix]
    if source is self:
        # Fresh indices as of this layer; drop what the last step cached.
        source._topk_ragged_cache = {}
    cached = source._topk_ragged_cache.get(self.compress_ratio)
    if cached is not None:
        return cached

    built = compute_global_topk_ragged_indices_and_indptr(
        self.topk_indices_buffer[:num_decode_tokens],
        swa_metadata.token_to_req_indices,
        attn_metadata.block_table[:num_decodes],
        attn_metadata.block_size // self.compress_ratio,
        swa_metadata.is_valid_token[:num_decode_tokens],
    )
    source._topk_ragged_cache[self.compress_ratio] = built
    return built

_forward_csa2_full(hidden_states, positions)

CSA layers: one fork replaces the base pipeline's three.

Each fork/join is a HIP event pair per side stream on ROCm, so the base three-fork pipeline is expensive there; the compressor chain and the indexer weights projection run fully in parallel with the default chain. The indexer K write and q-side stay serial after the join: wk consumes the aux-0 latent and the q-side's wq_b consumes the default-stream qr.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _forward_csa2_full(
    self, hidden_states: torch.Tensor, positions: torch.Tensor
) -> tuple[
    torch.Tensor,
    torch.Tensor,
    torch.Tensor | None,
    torch.Tensor | None,
    torch.Tensor | None,
]:
    """CSA layers: one fork replaces the base pipeline's three.

    Each fork/join is a HIP event pair per side stream on ROCm, so the
    base three-fork pipeline is expensive there; the compressor chain and
    the indexer weights projection run fully in parallel with the default
    chain. The indexer K write and q-side stay serial after the join:
    ``wk`` consumes the aux-0 latent and the q-side's wq_b consumes the
    default-stream qr.
    """
    attn_metadata = get_forward_context().attn_metadata
    compressor = self.compressor
    indexer = self.indexer
    assert compressor is not None and indexer is not None
    aux_streams = self.aux_stream_list
    assert aux_streams is not None and len(aux_streams) >= 2

    def default_chain():
        qr_kv = self._fused_wqa_wkv_gemm(hidden_states)
        qr, qr_scale, kv = self._split_qkv_and_norm(qr_kv)
        q = self._wq_b_proj(qr, qr_scale).view(
            -1, self.n_local_heads, self.head_dim
        )
        return (
            self._fused_qnorm_rope_kv_insert(q, kv, positions, attn_metadata),
            qr,
            qr_scale,
            kv,
        )

    def compressor_chain():
        kv_score = torch.mm(
            hidden_states,
            compressor.fused_wkv_wgate.weight.T,
            out_dtype=torch.float32,
        )
        latent = compressor(kv_score, positions)
        compressor.insert_cache(latent, positions, self.rotary_emb)
        return latent

    def indexer_weights_chain():
        # ReplicatedLinear returns (output, bias); bias is None.
        weights, _ = indexer.weights_proj(hidden_states)
        return weights

    (q, qr, qr_scale, kv), (latent, indexer_weights) = execute_in_parallel(
        default_chain,
        [compressor_chain, indexer_weights_chain],
        self.ln_events[0],
        self.ln_events[1:3],
        aux_streams[:2],
        enable=True,
        default_first=True,
    )

    indexer._produce_k(latent, positions, self.indexer_rotary_emb)
    index_q, index_q_scale, index_weights_out = indexer.forward_q(
        qr, qr_scale, indexer_weights, positions, self.indexer_rotary_emb
    )
    return q, kv, index_q, index_q_scale, index_weights_out

_forward_csa2_reindex(hidden_states, positions)

Indexer-only layers: the q-side overlaps the SWA q path.

The input projections run serial first: the q-side consumes qr, so it can only overlap the SWA q projection and KV insert.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _forward_csa2_reindex(
    self, hidden_states: torch.Tensor, positions: torch.Tensor
) -> tuple[
    torch.Tensor,
    torch.Tensor,
    torch.Tensor | None,
    torch.Tensor | None,
    torch.Tensor | None,
]:
    """Indexer-only layers: the q-side overlaps the SWA q path.

    The input projections run serial first: the q-side consumes qr, so
    it can only overlap the SWA q projection and KV insert.
    """
    indexer = self.indexer
    assert indexer is not None
    attn_metadata = get_forward_context().attn_metadata
    # Serial on these layers via the _run_parallel_input_projections
    # override: qr must sit on the default stream before the fork.
    qr_kv, _, indexer_weights = self._run_parallel_input_projections(hidden_states)
    assert indexer_weights is not None
    qr, qr_scale, kv = self._split_qkv_and_norm(qr_kv)

    def swa_q_chain():
        q = self._wq_b_proj(qr, qr_scale).view(
            -1, self.n_local_heads, self.head_dim
        )
        return self._fused_qnorm_rope_kv_insert(q, kv, positions, attn_metadata)

    def indexer_q_chain():
        return indexer.forward_q(
            qr, qr_scale, indexer_weights, positions, self.indexer_rotary_emb
        )

    aux_streams = self.aux_stream_list
    assert aux_streams is not None
    q, aux_results = execute_in_parallel(
        swa_q_chain,
        [indexer_q_chain],
        self.ln_events[0],
        self.ln_events[1:2],
        aux_streams[:1],
        enable=True,
        default_first=True,
    )
    index_q, index_q_scale, index_weights_out = aux_results[0]
    return q, kv, index_q, index_q_scale, index_weights_out

_forward_decode(q, positions, kv_cache, swa_metadata, attn_metadata, swa_only, output, output_mxfp8=None)

Returns how many leading rows the decode epilogue inverse-RoPE'd.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _forward_decode(
    self,
    q: torch.Tensor,
    positions: torch.Tensor,
    kv_cache: torch.Tensor | None,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    attn_metadata: DeepseekV4FlashMLAMetadata | None,
    swa_only: bool,
    output: torch.Tensor | None,
    output_mxfp8: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> int:
    """Returns how many leading rows the decode epilogue inverse-RoPE'd."""
    num_decodes = swa_metadata.num_decodes
    num_decode_tokens = swa_metadata.num_decode_tokens

    topk_lens = None
    topk_ragged_indices = None
    topk_ragged_indptr = None
    if not swa_only:
        (
            topk_ragged_indices,
            topk_ragged_indptr,
            topk_lens,
        ) = self._decode_topk_ragged(
            swa_metadata=swa_metadata,
            attn_metadata=attn_metadata,
            num_decodes=num_decodes,
            num_decode_tokens=num_decode_tokens,
        )

    if self._use_aiter_sparse_mla:
        assert swa_metadata.decode_swa_ragged_indices is not None
        assert swa_metadata.decode_swa_ragged_indptr is not None
        out = output
        if output_mxfp8 is not None:
            # The aiter kernel has no MXFP8 epilogue: write bf16 rows, then
            # rotate and quantize them as the decode reduce would have.
            out = self._mxfp8_bf16_rows(q.shape[0], q)
        assert out is not None
        self._aiter_sparse_mla(
            q,
            out,
            self.swa_cache_layer.kv_cache,
            swa_metadata.decode_swa_ragged_indices,
            swa_metadata.decode_swa_ragged_indptr,
            kv_cache,
            topk_ragged_indices,
            topk_ragged_indptr,
        )
        if output_mxfp8 is not None:
            rocm_inverse_rope_mxfp8_rows(
                out,
                positions,
                self.rotary_emb.cos_sin_cache,
                self.rope_head_dim,
                output_mxfp8[0],
                output_mxfp8[1],
            )
            return q.shape[0]
        # The aiter kernel leaves the inverse RoPE to the caller.
        return 0

    return rocm_sparse_attn_decode(
        q=q,
        kv_cache=kv_cache,
        swa_k_cache=self.swa_cache_layer.kv_cache,
        swa_only=swa_only,
        topk_indices=None,
        topk_lens=topk_lens,
        swa_indices=swa_metadata.decode_swa_indices,
        swa_lens=swa_metadata.decode_swa_lens,
        swa_ragged_indices=swa_metadata.decode_swa_ragged_indices,
        swa_ragged_indptr=swa_metadata.decode_swa_ragged_indptr,
        topk_ragged_indices=topk_ragged_indices,
        topk_ragged_indptr=topk_ragged_indptr,
        attn_sink=self.attn_sink,
        scale=self.scale,
        head_dim=self.head_dim,
        nope_head_dim=self.nope_head_dim,
        rope_head_dim=self.rope_head_dim,
        output=output,
        output_mxfp8=output_mxfp8,
        inv_rope_positions=positions,
        inv_rope_cos_sin_cache=self.rotary_emb.cos_sin_cache,
        extra_cache_nan_free=_trust_dsv4_extra_cache_nan_free(
            self.kv_cache_dtype,
            self._has_kv_transfer,
            not swa_only and kv_cache is not None,
        ),
    )

_forward_prefill(q, positions, compressed_k_cache, swa_k_cache, output, attn_metadata, swa_metadata, mxfp8_out=None)

Sparse prefill into bf16 output, or into mxfp8_out.

With mxfp8_out the rows go to a bf16 workspace first and are then inverse-RoPE'd and quantized into it in one pass.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _forward_prefill(
    self,
    q: torch.Tensor,
    positions: torch.Tensor,
    compressed_k_cache: torch.Tensor | None,
    swa_k_cache: torch.Tensor,
    output: torch.Tensor | None,
    attn_metadata: DeepseekV4FlashMLAMetadata | None,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    mxfp8_out: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> None:
    """Sparse prefill into bf16 ``output``, or into ``mxfp8_out``.

    With ``mxfp8_out`` the rows go to a bf16 workspace first and are then
    inverse-RoPE'd and quantized into it in one pass.
    """
    if self._use_aiter_sparse_mla:
        self._forward_prefill_aiter(
            q=q,
            positions=positions,
            compressed_k_cache=compressed_k_cache,
            swa_k_cache=swa_k_cache,
            output=output,
            attn_metadata=attn_metadata,
            swa_metadata=swa_metadata,
            mxfp8_out=mxfp8_out,
        )
        return

    swa_only = attn_metadata is None

    num_prefill_tokens = swa_metadata.num_prefill_tokens
    num_decodes = swa_metadata.num_decodes
    num_decode_tokens = swa_metadata.num_decode_tokens

    seq_lens = swa_metadata.prefill_seq_lens
    gather_lens = swa_metadata.prefill_gather_lens
    assert seq_lens is not None
    assert gather_lens is not None

    query_start_loc_cpu = swa_metadata.query_start_loc_cpu
    query_start_loc = swa_metadata.query_start_loc
    assert query_start_loc_cpu is not None
    assert query_start_loc is not None
    prefill_token_base = query_start_loc_cpu[num_decodes]

    # Local indices filled by the index source; SWA-only layers pass
    # top_k=0 and never read them.
    assert self.topk_indices_buffer is not None
    topk_indices = self.topk_indices_buffer[num_decode_tokens:]
    topk_indices = topk_indices[:num_prefill_tokens]
    if not swa_only:
        top_k = topk_indices.shape[-1]
        N = (self.max_model_len + self.compress_ratio - 1) // self.compress_ratio
    else:
        top_k = 0
        N = 0

    M = N + self.window_size + self.max_num_batched_tokens
    chunk_plan = swa_metadata.get_prefill_chunk_plan(
        compress_ratio=self.compress_ratio,
        prefill_chunk_size=self.PREFILL_CHUNK_SIZE,
        has_compressed=not swa_only,
    )
    assert chunk_plan, "prefill chunk plan must be non-empty when num_prefills > 0"

    workspace = current_workspace_manager().get_simultaneous(
        *self._prefill_workspace_shapes(M, num_prefill_tokens, q)
    )
    kv_rows = workspace[0].view(-1, q.shape[-1])
    if mxfp8_out is not None:
        output = workspace[1]
    assert output is not None
    for chunk_start, chunk_end, chunk_N, chunk_M in chunk_plan:
        chunk_size = chunk_end - chunk_start
        kv = kv_rows[: chunk_size * chunk_M].view(chunk_size, chunk_M, -1)
        if not swa_only:
            assert attn_metadata is not None
            assert compressed_k_cache is not None
            block_table = attn_metadata.block_table[num_decodes:]
            # compressed_k_cache is OCP on every platform (Triton encoder).
            dequantize_and_gather_k_cache(
                kv[:chunk_size],
                compressed_k_cache,
                seq_lens=seq_lens[chunk_start:chunk_end] // self.compress_ratio,
                gather_lens=None,
                block_table=block_table[chunk_start:chunk_end],
                block_size=attn_metadata.block_size // self.compress_ratio,
                offset=0,
                use_fnuz=False,
            )

        swa_block_table = swa_metadata.block_table[num_decodes:]
        dequantize_and_gather_k_cache(
            kv[:chunk_size],
            swa_k_cache,
            seq_lens=seq_lens[chunk_start:chunk_end],
            gather_lens=gather_lens[chunk_start:chunk_end],
            block_table=swa_block_table[chunk_start:chunk_end],
            block_size=swa_metadata.block_size,
            offset=chunk_N,
            use_fnuz=current_platform.is_fp8_fnuz(),
        )

        query_start = (
            query_start_loc_cpu[num_decodes + chunk_start] - prefill_token_base
        )
        query_end = (
            query_start_loc_cpu[num_decodes + chunk_end] - prefill_token_base
        )

        combined_indices, combined_lens = combine_topk_swa_indices(
            topk_indices[query_start:query_end],
            query_start_loc[
                num_decodes + chunk_start : num_decodes + chunk_end + 1
            ],
            seq_lens[chunk_start:chunk_end],
            gather_lens[chunk_start:chunk_end],
            self.window_size,
            self.compress_ratio,
            top_k,
            chunk_M,
            chunk_N,
        )
        rocm_sparse_attn_prefill(
            q=q[query_start:query_end],
            kv=kv.view(-1, 1, q.shape[-1]),
            indices=combined_indices,
            topk_length=combined_lens,
            scale=self.scale,
            head_dim=self.head_dim,
            nope_head_dim=self.nope_head_dim,
            rope_head_dim=self.rope_head_dim,
            attn_sink=self.attn_sink,
            output=output[query_start:query_end],
        )
    if mxfp8_out is not None:
        rocm_inverse_rope_mxfp8_rows(
            output,
            positions[:num_prefill_tokens],
            self.rotary_emb.cos_sin_cache,
            self.rope_head_dim,
            mxfp8_out[0][:num_prefill_tokens],
            mxfp8_out[1][:num_prefill_tokens],
        )

_forward_prefill_aiter(q, positions, compressed_k_cache, swa_k_cache, output, attn_metadata, swa_metadata, mxfp8_out=None)

Prefill straight off the paged caches, with no bf16 gather.

With mxfp8_out the kernel writes bf16 rows to the workspace, which are then inverse-RoPE'd and quantized into it as _forward_prefill does.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _forward_prefill_aiter(
    self,
    q: torch.Tensor,
    positions: torch.Tensor,
    compressed_k_cache: torch.Tensor | None,
    swa_k_cache: torch.Tensor,
    output: torch.Tensor | None,
    attn_metadata: DeepseekV4FlashMLAMetadata | None,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    mxfp8_out: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> None:
    """Prefill straight off the paged caches, with no bf16 gather.

    With ``mxfp8_out`` the kernel writes bf16 rows to the workspace, which
    are then inverse-RoPE'd and quantized into it as ``_forward_prefill``
    does.
    """
    num_prefill_tokens = swa_metadata.num_prefill_tokens
    assert swa_metadata.prefill_swa_ragged_indices is not None
    assert swa_metadata.prefill_swa_ragged_indptr is not None

    topk_indices = None
    topk_indptr = None
    if attn_metadata is None:
        compressed_k_cache = None
    else:
        assert compressed_k_cache is not None
        topk_indices, topk_indptr = self._prefill_topk_ragged(
            swa_metadata, attn_metadata, compressed_k_cache
        )

    if mxfp8_out is not None:
        output = self._mxfp8_bf16_rows(num_prefill_tokens, q)
    assert output is not None
    self._aiter_sparse_mla(
        q[:num_prefill_tokens],
        output[:num_prefill_tokens],
        swa_k_cache,
        swa_metadata.prefill_swa_ragged_indices,
        swa_metadata.prefill_swa_ragged_indptr,
        compressed_k_cache,
        topk_indices,
        topk_indptr,
    )
    if mxfp8_out is not None:
        rocm_inverse_rope_mxfp8_rows(
            output[:num_prefill_tokens],
            positions[:num_prefill_tokens],
            self.rotary_emb.cos_sin_cache,
            self.rope_head_dim,
            mxfp8_out[0][:num_prefill_tokens],
            mxfp8_out[1][:num_prefill_tokens],
        )

_mxfp8_bf16_rows(num_tokens, q)

bf16 attention rows for the aiter kernel to write when the layer's output is MXFP8, taken from the workspace the warmup reserved.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _mxfp8_bf16_rows(self, num_tokens: int, q: torch.Tensor) -> torch.Tensor:
    """bf16 attention rows for the aiter kernel to write when the layer's
    output is MXFP8, taken from the workspace the warmup reserved."""
    shapes = self._prefill_workspace_shapes(0, num_tokens, q, gather=False)
    assert shapes, "an MXFP8 attention output is gfx950-only"
    return current_workspace_manager().get_simultaneous(*shapes)[0]

_prefill_topk_ragged(swa_metadata, attn_metadata, compressed_k_cache)

Prefill counterpart of _decode_topk_ragged: layers that share an index source and a compressed cache reuse one build per step.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _prefill_topk_ragged(
    self,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    attn_metadata: DeepseekV4FlashMLAMetadata,
    compressed_k_cache: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Prefill counterpart of ``_decode_topk_ragged``: layers that share an
    index source and a compressed cache reuse one build per step."""
    assert swa_metadata.token_to_req_indices is not None
    assert swa_metadata.query_start_loc is not None
    assert swa_metadata.seq_lens is not None
    assert swa_metadata.is_valid_token is not None
    assert self.topk_indices_buffer is not None
    assert self._index_source_prefix is not None
    assert self.compressed_cache_prefix is not None

    source = self._static_forward_context[self._index_source_prefix]
    if source is self:
        source._prefill_topk_ragged_cache = {}
    cached = source._prefill_topk_ragged_cache.get(self.compressed_cache_prefix)
    if cached is not None:
        return cached

    num_decode_tokens = swa_metadata.num_decode_tokens
    num_prefill_tokens = swa_metadata.num_prefill_tokens
    built = build_prefill_topk_ragged_indices(
        self.topk_indices_buffer[
            num_decode_tokens : num_decode_tokens + num_prefill_tokens
        ],
        swa_metadata.token_to_req_indices,
        swa_metadata.query_start_loc,
        swa_metadata.seq_lens,
        swa_metadata.is_valid_token,
        attn_metadata.block_table,
        block_size=attn_metadata.block_size // self.compress_ratio,
        compress_ratio=self.compress_ratio,
        num_compressed=(self.max_model_len + self.compress_ratio - 1)
        // self.compress_ratio,
        token_offset=num_decode_tokens,
        num_rows=compressed_k_cache.shape[0] * compressed_k_cache.shape[1],
    )
    source._prefill_topk_ragged_cache[self.compressed_cache_prefix] = built
    return built

_prefill_workspace_shapes(M, num_prefill_tokens, q, gather=True)

The prefill gather buffer, plus bf16 rows for the MXFP8 output.

One get_simultaneous call: separate calls alias the same memory. gather=False leaves out the gather buffer, which the aiter kernel, reading the caches in place, never needs.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _prefill_workspace_shapes(
    self, M: int, num_prefill_tokens: int, q: torch.Tensor, gather: bool = True
) -> list[tuple[tuple[int, ...], torch.dtype]]:
    """The prefill gather buffer, plus bf16 rows for the MXFP8 output.

    One ``get_simultaneous`` call: separate calls alias the same memory.
    ``gather=False`` leaves out the gather buffer, which the aiter kernel,
    reading the caches in place, never needs.
    """
    shapes: list[tuple[tuple[int, ...], torch.dtype]] = []
    if gather:
        shapes.append(((self.PREFILL_CHUNK_SIZE, M, q.shape[-1]), torch.bfloat16))
    if _ON_GFX950:
        shapes.append(
            (
                (num_prefill_tokens, self.n_local_heads, self.head_dim),
                torch.bfloat16,
            )
        )
    return shapes

_split_qkv_and_norm(qr_kv)

Fuse q/kv RMSNorm + per-1x128 fp8 q quant into one aiter kernel.

The shared path norms q and kv in one triton kernel and the wq_b linears then re-read the bf16 qr to quantize it. The aiter kernel computes both RMSNorms (fp32 accumulate) and the fp8 group quant in a single pass, writing fp8 qr + group scales directly; both wq_b GEMMs (attention and indexer) then consume that pair and skip their own input quant. kv stays bf16: the fused insert kernel RoPE/quantizes it itself. Falls back to the shared path when the aiter linear path is not active.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _split_qkv_and_norm(
    self, qr_kv: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor]:
    """Fuse q/kv RMSNorm + per-1x128 fp8 q quant into one aiter kernel.

    The shared path norms q and kv in one triton kernel and the wq_b
    linears then re-read the bf16 qr to quantize it. The aiter kernel
    computes both RMSNorms (fp32 accumulate) and the fp8 group quant
    in a single pass, writing fp8 qr + group scales directly; both
    wq_b GEMMs (attention and indexer) then consume that pair and
    skip their own input quant. kv stays bf16: the fused insert
    kernel RoPE/quantizes it itself. Falls back to the shared path
    when the aiter linear path is not active.
    """
    qr, kv = qr_kv.split([self.q_lora_rank, self.head_dim], dim=-1)
    if not (
        qr.dim() == 2
        and qr.shape[0] > 0
        and self.q_lora_rank % 128 == 0
        and self._wq_b_uses_aiter_block_scaled
    ):
        return super()._split_qkv_and_norm(qr_kv)

    from vllm._aiter_ops import rocm_aiter_ops

    return rocm_aiter_ops.fused_qk_rmsnorm_group_quant(
        q=qr,
        q_weight=self.q_norm.weight.data,
        q_epsilon=self.eps,
        kv=kv,
        kv_weight=self.kv_norm.weight.data,
        kv_epsilon=self.eps,
        group_size=128,
        transpose_scale=False,
    )

DeepseekV41RocmMxfp4Indexer

Bases: DeepseekV4Indexer

The indexer on aiter's paged MXFP4 cache: its K store and Q quant write in the order aiter's MQA-logits kernel reads, and its layers score with it. _produce_k and forward_q mirror DeepseekV4Indexer's except for those two calls.

Source code in vllm/models/deepseek_v41/amd/rocm.py
class DeepseekV41RocmMxfp4Indexer(DeepseekV4Indexer):
    """The indexer on aiter's paged MXFP4 cache: its K store and Q quant write
    in the order aiter's MQA-logits kernel reads, and its layers score with it.
    ``_produce_k`` and ``forward_q`` mirror ``DeepseekV4Indexer``'s except for
    those two calls."""

    mqa_cls = RocmSparseMQAIndexer
    attn_cls = RocmSparseAttnIndexer

    def _produce_k(
        self,
        latent: torch.Tensor | None,
        positions: torch.Tensor,
        rotary_emb: torch.nn.Module,
    ) -> None:
        attn_metadata = get_forward_context().attn_metadata
        if not isinstance(attn_metadata, dict) or latent is None:
            return
        assert self.owns_k
        indexer_metadata = cast(Any, attn_metadata[self.k_cache.prefix])
        k_pre, _ = self.wk(latent)
        rocm_mxfp4_indexer_k_store(
            k_pre,
            positions,
            rotary_emb.cos_sin_cache,
            self.k_norm.weight,
            self.k_norm.variance_epsilon,
            self.k_cache.kv_cache,
            indexer_metadata.slot_mapping,
            self.compress_ratio,
            self.use_fp4_kv,
            num_heads=self.n_head,
        )

    def forward_q(
        self,
        qr: torch.Tensor | QuantizedActivation,
        qr_scale: torch.Tensor | None,
        indexer_weights: torch.Tensor,
        positions: torch.Tensor,
        rotary_emb: torch.nn.Module,
    ) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]:
        q = self._wq_b_proj(qr, qr_scale).view(-1, self.n_head, self.head_dim)
        (q, q_scale), weights = rocm_mxfp4_indexer_q_quant(
            positions,
            q,
            rotary_emb.cos_sin_cache,
            indexer_weights,
            self.softmax_scale,
            self.n_head**-0.5,
            use_fp4=self.use_fp4_kv,
            weights_out_dtype=self.indexer_weights_dtype,
        )
        return q, q_scale, weights

_aiter_indexer_cache_ops() cached

The aiter indexer key writer and query quantizer.

Source code in vllm/models/deepseek_v41/amd/rocm.py
@functools.cache
def _aiter_indexer_cache_ops() -> tuple[Callable[..., None], Callable[..., tuple]]:
    """The aiter indexer key writer and query quantizer."""
    from aiter.ops.triton.fusions.k_norm_rope_mxfp4_cache import (
        k_norm_rope_mxfp4_cache,
    )
    from aiter.ops.triton.rope.q_rope_mxfp4_quant import q_rope_mxfp4_quant

    return k_norm_rope_mxfp4_cache, q_rope_mxfp4_quant

_copy_ragged_to_graph_buffers(ragged_indices, ragged_indptr, ragged_indices_buffer, ragged_indptr_buffer, num_rows, max_entries_per_row)

Copy dynamic ragged metadata into persistent CUDA graph buffers.

FULL decode graphs capture kernel argument addresses. Keep the returned tensors backed by stable storage, while indptr continues to bound reads.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def _copy_ragged_to_graph_buffers(
    ragged_indices: torch.Tensor,
    ragged_indptr: torch.Tensor,
    ragged_indices_buffer: torch.Tensor,
    ragged_indptr_buffer: torch.Tensor,
    num_rows: int,
    max_entries_per_row: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Copy dynamic ragged metadata into persistent CUDA graph buffers.

    FULL decode graphs capture kernel argument addresses. Keep the returned
    tensors backed by stable storage, while indptr continues to bound reads.
    """
    indptr_out = ragged_indptr_buffer[: num_rows + 1]
    indptr_out.copy_(ragged_indptr, non_blocking=True)

    max_entries = max(num_rows * max_entries_per_row, 1)
    ragged_out = ragged_indices_buffer[:max_entries]
    source_entries = ragged_indices.numel()
    if source_entries > 0:
        ragged_out[:source_entries].copy_(ragged_indices, non_blocking=True)
    if _ON_GFX950:
        # Preserve the graph-stable base pointer while exposing source capacity
        # to the sync-free split selector; indptr still carries the true NNZ.
        ragged_out = ragged_out[: max(source_entries, 1)]
    return ragged_out, indptr_out

apply_pre_quantized_block_scaled_mm(linear, x_fp8, x_scale)

Block-scaled fp8 GEMM on pre-quantized activations.

The fused q/kv norm kernel writes fp8 qr + per-1x128 scales; this drives the linear's block-scaled GEMM directly with them, bypassing apply_weights which would re-quantize the fp8 input. Only valid for the wq_b-style column/replicated linears: their output is the local TP shard, so no all-reduce is needed.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def apply_pre_quantized_block_scaled_mm(
    linear: torch.nn.Module,
    x_fp8: torch.Tensor,
    x_scale: torch.Tensor,
) -> torch.Tensor:
    """Block-scaled fp8 GEMM on pre-quantized activations.

    The fused q/kv norm kernel writes fp8 qr + per-1x128 scales; this
    drives the linear's block-scaled GEMM directly with them, bypassing
    apply_weights which would re-quantize the fp8 input. Only valid for
    the wq_b-style column/replicated linears: their output is the local
    TP shard, so no all-reduce is needed.
    """
    from vllm.model_executor.kernels.linear.scaled_mm.BlockScaledMMLinearKernel import (
        FP8BlockParams,
    )

    params = FP8BlockParams.from_layer(linear)
    weight_scale = (
        params.weight_scale
        if params.weight_scale_inv is None
        else params.weight_scale_inv
    )
    kernel = linear.quant_method.fp8_linear
    out = kernel.apply_block_scaled_mm(
        A=x_fp8, B=params.weight, As=x_scale, Bs=weight_scale
    )
    return out.to(dtype=kernel.config.out_dtype)

combine_topk_swa_indices(topk_indices, query_start_loc, seq_lens, gather_lens, window_size, compress_ratio, topk, M, N)

Combine compressed-attention and sliding-window indices with Torch.

The Triton implementation inherited from DeepSeek V4 launches a two-dimensional grid with 128 workers per request. On gfx950 it can issue an out-of-bounds access for V4.1's mixed prefill metadata (including the synthetic mixed-token warmup). This path is prefill-only and the tensors are small, so use ordinary Torch indexing until a gfx950-safe fused kernel is available.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def combine_topk_swa_indices(
    topk_indices: torch.Tensor,
    query_start_loc: torch.Tensor,
    seq_lens: torch.Tensor,
    gather_lens: torch.Tensor,
    window_size: int,
    compress_ratio: int,
    topk: int,
    M: int,
    N: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Combine compressed-attention and sliding-window indices with Torch.

    The Triton implementation inherited from DeepSeek V4 launches a
    two-dimensional grid with 128 workers per request.  On gfx950 it can issue
    an out-of-bounds access for V4.1's mixed prefill metadata (including the
    synthetic mixed-token warmup).  This path is prefill-only and the tensors
    are small, so use ordinary Torch indexing until a gfx950-safe fused kernel
    is available.
    """
    topk_indices = topk_indices.reshape(topk_indices.shape[0], -1).contiguous()
    num_tokens = topk_indices.shape[0]
    combined_topk = (
        (topk + window_size + _SPARSE_PREFILL_TOPK_ALIGNMENT - 1)
        // _SPARSE_PREFILL_TOPK_ALIGNMENT
        * _SPARSE_PREFILL_TOPK_ALIGNMENT
    )
    combined_indices = torch.full(
        (num_tokens, combined_topk),
        fill_value=-1,
        dtype=torch.int32,
        device=topk_indices.device,
    )
    combined_lens = torch.empty(
        num_tokens, dtype=torch.int32, device=topk_indices.device
    )

    # query_start_loc may have a non-zero base for a narrowed mixed batch.
    query_lens = query_start_loc[1:] - query_start_loc[:-1]
    req_ids = torch.repeat_interleave(
        torch.arange(seq_lens.shape[0], device=seq_lens.device), query_lens
    )
    query_starts = query_start_loc[:-1] - query_start_loc[0]
    token_offsets = torch.arange(num_tokens, device=seq_lens.device) - (
        torch.repeat_interleave(query_starts, query_lens)
    )
    positions = seq_lens[req_ids] - query_lens[req_ids] + token_offsets

    logical_topk_width = min(topk, topk_indices.shape[1])
    topk_lens = torch.minimum(
        (positions + 1) // compress_ratio,
        torch.full_like(positions, logical_topk_width),
    ).clamp_min(0)
    topk_offsets = torch.arange(logical_topk_width, device=seq_lens.device)
    topk_mask = topk_offsets[None, :] < topk_lens[:, None]
    topk_values = topk_indices[:, :logical_topk_width].to(torch.int32)
    topk_valid = topk_mask & (topk_values >= 0) & (topk_values < N)
    combined_indices[:, :logical_topk_width] = torch.where(
        topk_valid,
        topk_values + (M * req_ids).to(torch.int32)[:, None],
        -1,
    )

    swa_lens = torch.minimum(
        positions + 1, torch.full_like(positions, window_size)
    ).clamp_min(0)
    swa_offsets = torch.arange(window_size, device=seq_lens.device)
    swa_mask = swa_offsets[None, :] < swa_lens[:, None]
    swa_columns = topk_lens[:, None] + swa_offsets[None, :]
    gather_starts = seq_lens - gather_lens
    swa_values = (
        M * req_ids[:, None]
        + N
        + swa_offsets[None, :]
        + positions[:, None]
        - swa_lens[:, None]
        + 1
        - gather_starts[req_ids, None]
    ).to(torch.int32)
    rows = torch.arange(num_tokens, device=seq_lens.device)[:, None].expand_as(
        swa_columns
    )
    flat_dst = rows[swa_mask] * combined_topk + swa_columns[swa_mask]
    combined_indices.view(-1).index_copy_(0, flat_dst, swa_values[swa_mask])
    combined_lens.copy_((topk_lens + swa_lens).to(torch.int32))
    return combined_indices, combined_lens

rocm_mxfp4_indexer_k_store(k_pre, positions, cos_sin_cache, rms_norm_weight, rms_norm_eps, k_cache, kv_slot_mapping, compress_ratio, use_fp4_cache, *, num_heads)

indexer_k_norm_rope_store for the ROCm MXFP4 cache: aiter's cache op writes the key in the order its MQA-logits kernel reads with num_heads query heads.

Source code in vllm/models/deepseek_v41/amd/rocm.py
def rocm_mxfp4_indexer_k_store(
    k_pre: torch.Tensor,
    positions: torch.Tensor,
    cos_sin_cache: torch.Tensor,
    rms_norm_weight: torch.Tensor,
    rms_norm_eps: float,
    k_cache: torch.Tensor,
    kv_slot_mapping: torch.Tensor,
    compress_ratio: int,
    use_fp4_cache: bool,
    *,
    num_heads: int,
) -> None:
    """`indexer_k_norm_rope_store` for the ROCm MXFP4 cache: aiter's cache op
    writes the key in the order its MQA-logits kernel reads with ``num_heads``
    query heads."""
    assert use_fp4_cache, "the ROCm indexer cache op writes MXFP4 only"
    layout = rocm_paged_mxfp4_cache_layout(num_heads, k_pre.shape[1], k_cache.shape[1])
    k_norm_rope_mxfp4_cache, _ = _aiter_indexer_cache_ops()
    k_norm_rope_mxfp4_cache(
        k_pre,
        positions,
        cos_sin_cache,
        rms_norm_weight,
        rms_norm_eps,
        k_cache,
        kv_slot_mapping,
        compress_ratio,
        shuffle=layout,
    )

rocm_mxfp4_indexer_q_quant(positions, index_q, index_q_cos_sin_cache, index_weights, index_weights_softmax_scale, index_weights_head_scale, use_fp4=True, weights_out_dtype=torch.float32)

fused_indexer_q_rope_quant for the ROCm MXFP4 indexer, on aiter's op: ((packed [T, H, D // 2], e8m0 as one int32 per head [T, H]), fp32 weights).

Source code in vllm/models/deepseek_v41/amd/rocm.py
def rocm_mxfp4_indexer_q_quant(
    positions: torch.Tensor,
    index_q: torch.Tensor,
    index_q_cos_sin_cache: torch.Tensor,
    index_weights: torch.Tensor,
    index_weights_softmax_scale: float,
    index_weights_head_scale: float,
    use_fp4: bool = True,
    weights_out_dtype: torch.dtype = torch.float32,
) -> tuple[tuple[torch.Tensor, torch.Tensor], torch.Tensor]:
    """`fused_indexer_q_rope_quant` for the ROCm MXFP4 indexer, on aiter's op:
    ((packed [T, H, D // 2], e8m0 as one int32 per head [T, H]), fp32
    weights)."""
    assert use_fp4 and weights_out_dtype == torch.float32, (
        "the ROCm indexer quantizes Q to MXFP4 and scores with fp32 weights"
    )
    _, q_rope_mxfp4_quant = _aiter_indexer_cache_ops()
    q_packed, q_scale, weights_out = q_rope_mxfp4_quant(
        index_q,
        positions,
        index_q_cos_sin_cache,
        index_weights,
        index_weights_softmax_scale * index_weights_head_scale,
    )
    return (q_packed, q_scale.view(torch.int32).squeeze(-1)), weights_out