Skip to content

vllm.models.deepseek_v4.amd.rocm

Classes:

Functions:

DeepseekV4ROCMAiterMLAAttention

Bases: DeepseekV4Attention

ROCm sparse MLA attention layer for DeepSeek V4.

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

    backend_cls = DeepseekV4ROCMAiterMLASparseBackend
    _use_aiter_sparse_mla = False

    def __init__(self, *args, **kwargs):
        vllm_config = args[0] if args else kwargs["vllm_config"]
        super().__init__(*args, **kwargs)
        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._wo_a_fp8_weight: torch.Tensor | None = None
        self._wo_a_e8m0_scale: torch.Tensor | None = None
        self._wo_a_cos_cache: torch.Tensor | None = None
        self._wo_a_sin_cache: 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

        if self.indexer is None:
            # Dense layers have no compressor work to overlap; HCA layers
            # (compressor, no indexer) keep the streams for the dual-stream
            # fork below.
            if self.compressor is None:
                self.aux_stream_list = None
        else:
            # Disable indexer inner overlap.
            self.indexer.aux_stream = None
        self._use_aiter_sparse_mla = _aiter_sparse_mla_enabled(self._has_kv_transfer)

    def _run_sequential_pipeline(
        self,
        hidden_states: torch.Tensor,
        positions: torch.Tensor,
        o_padded: torch.Tensor,
    ) -> None:
        """Disable ROCm streams when the current execution region cannot overlap."""
        aux_streams = self.aux_stream_list
        self.aux_stream_list = None
        try:
            qr_kv, kv_score, indexer_kv_score, indexer_weights = (
                self._run_parallel_input_projections(hidden_states)
            )
            qr, qr_scale, kv = self._split_qkv_and_norm(qr_kv)
            self._prepare_and_attn_fn(
                hidden_states,
                qr,
                kv,
                qr_scale,
                kv_score,
                indexer_kv_score,
                indexer_weights,
                positions,
                o_padded,
            )
        finally:
            self.aux_stream_list = aux_streams

    def forward(
        self,
        positions: torch.Tensor,
        hidden_states: torch.Tensor,
        llama_4_scaling: torch.Tensor | None = None,
    ) -> torch.Tensor:
        # Pre-allocate attention output with FlashMLA-padded head count.
        # The op writes into `o_padded`; we slice to n_local_heads after.
        num_tokens = hidden_states.shape[0]
        o_padded = torch.empty(
            (num_tokens, self.padded_heads, self.head_dim),
            dtype=hidden_states.dtype,
            device=hidden_states.device,
        )

        if current_platform.enable_multi_stream_overlap(
            self.aux_stream_list, get_forward_context().attn_metadata
        ):
            # The ROCm override consumes these sentinels inside the capture
            # boundary, moving the stream fan-out ahead of the projections.
            self._prepare_and_attn_fn(
                hidden_states,
                None,
                None,
                None,
                None,
                None,
                None,
                positions,
                o_padded,
            )
        else:
            self._run_sequential_pipeline(hidden_states, positions, o_padded)

        o = o_padded[:, : self.n_local_heads, :]

        # Inverse-RoPE + wo_a + wo_b output projection (platform-specific).
        return self._o_proj(o, positions)

    def _prepare_and_attn(
        self,
        hidden_states: torch.Tensor,
        qr: torch.Tensor | None,
        kv: torch.Tensor | None,
        qr_scale: torch.Tensor | None,
        kv_score: torch.Tensor | None,
        indexer_kv_score: torch.Tensor | None,
        indexer_weights: torch.Tensor | None,
        positions: torch.Tensor,
        o_padded: torch.Tensor,
    ) -> None:
        """Run the ROCm fork/join (HCA or CSA) inside the capture boundary."""
        aux_streams = self.aux_stream_list
        # The sequential pipeline disables aux_stream_list before calling
        # back with real projection inputs; aux_streams is None ends that
        # recursion here.
        if aux_streams is None:
            saved_streams = self.aux_stream_list
            self.aux_stream_list = None
            try:
                super()._prepare_and_attn(
                    hidden_states,
                    cast(torch.Tensor, qr),
                    cast(torch.Tensor, kv),
                    qr_scale,
                    cast(torch.Tensor, kv_score),
                    cast(torch.Tensor, indexer_kv_score),
                    cast(torch.Tensor, indexer_weights),
                    positions,
                    o_padded,
                )
            finally:
                self.aux_stream_list = saved_streams
            return

        # Re-check: forward's gate ran inside a captured segment that
        # _prepare_and_attn_eager (MRV1) then broke, making this region eager.
        if not current_platform.enable_multi_stream_overlap(
            self.aux_stream_list, get_forward_context().attn_metadata
        ):
            self._run_sequential_pipeline(hidden_states, positions, o_padded)
            return

        indexer = self.indexer
        compressor = self.compressor
        assert compressor is not None

        def default_chain():
            qr_kv = self._fused_wqa_wkv_gemm(hidden_states)
            qr_out, qr_scale_out, kv_out = self._split_qkv_and_norm(qr_kv)
            q = self._wq_b_proj(qr_out, qr_scale_out).view(
                -1, self.n_local_heads, self.head_dim
            )
            attn_metadata = get_forward_context().attn_metadata
            q = self._fused_qnorm_rope_kv_insert(q, kv_out, positions, attn_metadata)
            return q, qr_out, qr_scale_out, kv_out

        def main_compressor_chain() -> None:
            score = torch.mm(
                hidden_states,
                compressor.fused_wkv_wgate.weight.T,
                out_dtype=torch.float32,
            )
            compressor(score, positions, self.rotary_emb)

        if indexer is None:
            # HCA dual-stream: the main compressor runs on aux stream 0 while
            # the default stream produces q and inserts KV into the SWA cache.
            # Both branches only read hidden_states, so the join merely has to
            # precede the sparse attention that consumes the compressed KV.
            (q, _qr_out, _qr_scale_out, kv_out), _ = execute_in_parallel(
                default_chain,
                [main_compressor_chain],
                self.ln_events[0],
                [self.ln_events[1]],
                aux_streams[:1],
                enable=True,
            )
            self._sparse_indexer_and_attn(
                hidden_states, None, None, None, q, kv_out, positions, o_padded
            )
            return

        def indexer_compressor_chain() -> None:
            score = torch.mm(
                hidden_states,
                indexer.compressor.fused_wkv_wgate.weight.T,
                out_dtype=torch.float32,
            )
            indexer.compressor(score, positions, self.indexer_rotary_emb)

        # CSA three-stream: the main and indexer compressors run on aux
        # streams 0 and 1 while the default stream produces q and inserts KV
        # into the SWA cache. Every branch only reads hidden_states plus its
        # own state, so the join merely has to precede the indexer op and the
        # sparse attention, which consume the compressed KV caches.
        (q, qr_out, qr_scale_out, kv_out), _ = execute_in_parallel(
            default_chain,
            [main_compressor_chain, indexer_compressor_chain],
            self.ln_events[0],
            self.ln_events[1:3],
            aux_streams[:2],
            enable=True,
        )

        indexer_weights_out, _ = indexer.weights_proj(hidden_states)
        # The indexer compressor already ran on aux stream 1; build queries only.
        index_q, index_q_scale, weights = indexer(
            hidden_states,
            qr_out,
            None,
            indexer_weights_out,
            positions,
            self.indexer_rotary_emb,
            qr_scale_out,
            skip_compressor=True,
        )
        self._sparse_indexer_and_attn(
            hidden_states,
            index_q,
            index_q_scale,
            weights,
            q,
            kv_out,
            positions,
            o_padded,
        )

    @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,
            get_fp8_block_weight_scale,
        )
        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 = get_fp8_block_weight_scale(linear)
            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)
        if _ON_GFX950 and envs.VLLM_ROCM_USE_AITER_FP8BMM:
            self._prepare_fp8_wo_a()

    def _prepare_fp8_wo_a(self) -> None:
        try:
            from aiter.ops.batched_gemm_op_a8w8 import (
                batched_gemm_a8w8_mxscale as mxscale_op,
            )
            from aiter.ops.inverse_rope_group_quant import (
                inverse_rope_group_quant as inverse_quant_op,
            )
        except ImportError:
            logger.warning_once(
                "The DeepSeek V4 FP8 WO_A path requires AITER >= 0.1.20; "
                "falling back to BF16 WO_A."
            )
            return
        del mxscale_op, inverse_quant_op

        from vllm.model_executor.layers.quantization.utils.fp8_utils import (
            get_fp8_block_weight_scale,
            is_fp8,
        )

        weight = getattr(self.wo_a, "weight", None)
        scale = get_fp8_block_weight_scale(self.wo_a)
        if scale is None:
            # ModelOpt MXFP8 stores the multiplicative E8M0 scale without the
            # historical ``_inv`` suffix.
            scale = getattr(self.wo_a, "weight_scale", None)
        if weight is None or scale is None:
            logger.warning_once(
                "DeepSeek V4 FP8 WO_A needs a block-scaled FP8 wo_a weight; "
                "the layer exposes no weight/weight scale. Falling back to "
                "BF16 WO_A."
            )
            return
        if weight.dim() != 2 or scale.dim() != 2 or not is_fp8(weight.dtype):
            logger.warning_once(
                "DeepSeek V4 FP8 WO_A needs a 2-D FP8 wo_a weight with a 2-D "
                "block scale, got weight %s%s and scale %s. Falling back to "
                "BF16 WO_A.",
                weight.dtype,
                tuple(weight.shape),
                tuple(scale.shape),
            )
            return

        groups = self.n_local_groups
        out_per_group = self.o_lora_rank
        out_features, in_features = weight.shape
        if (
            out_features != groups * out_per_group
            or out_per_group % 128 != 0
            or in_features % 128 != 0
            or scale.shape != (out_features // 128, in_features // 128)
        ):
            logger.warning_once(
                "DeepSeek V4 FP8 WO_A needs group-128 blocks for %d groups of "
                "%d outputs, got weight %s and scale %s. Falling back to BF16 "
                "WO_A.",
                groups,
                out_per_group,
                tuple(weight.shape),
                tuple(scale.shape),
            )
            return

        e8m0_scale = _wo_a_block_scale_to_e8m0(scale)
        if e8m0_scale is None:
            logger.warning_once(
                "DeepSeek V4 FP8 WO_A could not losslessly encode the %s wo_a "
                "block scale as OCP E8M0. Falling back to BF16 WO_A.",
                scale.dtype,
            )
            return

        self._wo_a_fp8_weight = weight.view(groups, out_per_group, in_features)
        self._wo_a_e8m0_scale = e8m0_scale.view(
            groups, out_per_group // 128, in_features // 128
        )
        cache = getattr(self.rotary_emb, "cos_sin_cache_bf16", None)
        if cache is None:
            cache = self.rotary_emb.cos_sin_cache.to(dtype=torch.bfloat16)
        cos_cache, sin_cache = cache.chunk(2, dim=-1)
        self._wo_a_cos_cache = cos_cache.contiguous()
        self._wo_a_sin_cache = sin_cache.contiguous()

    def prepare_compressor_gemm_fusion(self) -> bool:
        if self._fused_compressor_weight is not None:
            return False

        from vllm.model_executor.offloader import NoopOffloader, get_offloader

        if not isinstance(get_offloader(), NoopOffloader):
            logger.warning_once(
                "DeepSeek V4 compressor GEMM fusion is incompatible with "
                "weight offloading and will remain disabled."
            )
            return False

        compressor = self.compressor
        indexer = self.indexer
        if compressor is None or indexer is None:
            return False

        main_weight = compressor.fused_wkv_wgate.weight
        indexer_weight = indexer.compressor.fused_wkv_wgate.weight
        if main_weight.ndim != 2 or indexer_weight.ndim != 2:
            raise ValueError("DeepSeek V4 compressor weights must be matrices")
        if main_weight.shape[1] != indexer_weight.shape[1]:
            raise ValueError("DeepSeek V4 compressor weights must share K")
        if main_weight.dtype != indexer_weight.dtype:
            raise ValueError("DeepSeek V4 compressor weights must share dtype")
        if main_weight.device != indexer_weight.device:
            raise ValueError("DeepSeek V4 compressor weights must share device")

        main_size = main_weight.shape[0]
        indexer_size = indexer_weight.shape[0]
        fused_weight = torch.cat((main_weight, indexer_weight), dim=0)
        with torch.no_grad():
            main_weight.set_(fused_weight[:main_size])
            indexer_weight.set_(fused_weight[main_size:])

        self._fused_compressor_weight = fused_weight
        self._fused_compressor_split_sizes = (main_size, indexer_size)
        return True

    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,
        torch.Tensor | None,
    ]:
        fused_weight = self._fused_compressor_weight
        split_sizes = self._fused_compressor_split_sizes
        if fused_weight is None or split_sizes is None:
            return super()._run_parallel_input_projections(hidden_states)

        indexer = self.indexer
        if indexer is None:
            raise RuntimeError("Fused compressor weight requires a C4 indexer")

        qr_kv = self._fused_wqa_wkv_gemm(hidden_states)
        fused_scores = torch.mm(
            hidden_states,
            fused_weight.T,
            out_dtype=torch.float32,
        )
        kv_score, indexer_kv_score = fused_scores.split(split_sizes, dim=-1)
        indexer_weights, _ = indexer.weights_proj(hidden_states)
        return qr_kv, kv_score, indexer_kv_score, indexer_weights

    @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 _o_proj(self, o: torch.Tensor, positions: torch.Tensor) -> torch.Tensor:
        if self._wo_a_fp8_weight is not None:
            from aiter.ops.batched_gemm_op_a8w8 import (
                batched_gemm_a8w8_mxscale,
            )
            from aiter.ops.inverse_rope_group_quant import (
                inverse_rope_group_quant,
            )

            assert self._wo_a_cos_cache is not None
            assert self._wo_a_sin_cache is not None
            o_fp8, o_scale = inverse_rope_group_quant(
                o.view(o.shape[0], self.n_local_heads, self.head_dim),
                positions.to(torch.int64),
                self._wo_a_cos_cache,
                self._wo_a_sin_cache,
                num_groups=self.n_local_groups,
                quant_group_size=128,
            )
            assert self._wo_a_e8m0_scale is not None
            zf = batched_gemm_a8w8_mxscale(
                o_fp8,
                self._wo_a_fp8_weight,
                o_scale,
                self._wo_a_e8m0_scale,
                dtype=o.dtype,
            ).flatten(1)
        else:
            # 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,
            )
            zf = z.flatten(1)
        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,
    ) -> None:
        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 same bf16
            # gather workspace _forward_prefill would; the dequantize / topk
            # / sparse_fwd kernels are skipped this step. The aiter prefill
            # reads the caches in place and never asks for it.
            if not self._use_aiter_sparse_mla:
                swa_only = self.compress_ratio <= 1
                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
                current_workspace_manager().get_simultaneous(
                    ((self.PREFILL_CHUNK_SIZE, M, q.shape[-1]), torch.bfloat16),
                )
            output.zero_()
            return

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

        swa_only = self.compress_ratio <= 1
        self_kv_cache = self.kv_cache if not swa_only else None
        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:],
                attn_metadata=rocm_metadata,
                swa_metadata=swa_metadata,
            )
        # The fp8 wo_a path rotates inside inverse_rope_group_quant, so folding
        # the rotation into the decode reduce would apply it twice. Only the
        # BF16 einsum path hands its rotation off to the decode.
        fuse_inv_rope = self._wo_a_fp8_weight is None
        rotated = 0
        if num_decodes > 0:
            rotated = self._forward_decode(
                q=q[:num_decode_tokens],
                positions=positions[:num_decode_tokens] if fuse_inv_rope else None,
                kv_cache=self_kv_cache,
                swa_metadata=swa_metadata,
                attn_metadata=rocm_metadata,
                swa_only=swa_only,
                output=output[:num_decode_tokens],
                adaptive_splits=(
                    _ON_GFX950
                    and not swa_only
                    and self.compress_ratio == 128
                    and rocm_metadata is not None
                    and rocm_metadata.for_cudagraph_capture
                ),
            )
        if fuse_inv_rope:
            # 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 _forward_decode(
        self,
        q: torch.Tensor,
        positions: torch.Tensor | None,
        kv_cache: torch.Tensor | None,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
        attn_metadata: DeepseekV4ROCMAiterMLASparseMetadata | None,
        swa_only: bool,
        output: torch.Tensor,
        adaptive_splits: bool,
    ) -> 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_indices = None
        topk_lens = None
        topk_ragged_indices = None
        topk_ragged_indptr = None
        if not swa_only:
            assert attn_metadata is not None
            assert swa_metadata.is_valid_token is not None
            block_size = attn_metadata.block_size // self.compress_ratio
            is_valid = swa_metadata.is_valid_token[:num_decode_tokens]
            if self.compress_ratio == 4:
                assert self.topk_indices_buffer is not None
                (
                    topk_ragged_indices,
                    topk_ragged_indptr,
                    topk_lens,
                ) = 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],
                    block_size,
                    is_valid,
                )
            else:
                topk_indices = attn_metadata.c128a_global_decode_topk_indices
                topk_lens = attn_metadata.c128a_decode_topk_lens
                topk_ragged_indices = attn_metadata.c128a_decode_topk_ragged_indices
                topk_ragged_indptr = attn_metadata.c128a_decode_topk_ragged_indptr

        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
            self._aiter_sparse_mla(
                q,
                output,
                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,
            )
            # 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=topk_indices,
            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,
            adaptive_splits=adaptive_splits,
            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,
        attn_metadata: DeepseekV4ROCMAiterMLASparseMetadata | None,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    ) -> None:
        if self._use_aiter_sparse_mla:
            self._forward_prefill_aiter(
                q=q,
                compressed_k_cache=compressed_k_cache,
                swa_k_cache=swa_k_cache,
                output=output,
                attn_metadata=attn_metadata,
                swa_metadata=swa_metadata,
            )
            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]
        left_visible = swa_metadata.prefill_left_visible
        right_visible = swa_metadata.prefill_right_visible
        if left_visible is not None:
            left_visible = left_visible[num_decode_tokens:]
            assert right_visible is not None
            right_visible = right_visible[num_decode_tokens:]

        if not swa_only:
            if self.compress_ratio == 4:
                assert self.topk_indices_buffer is not None
                topk_indices = self.topk_indices_buffer[num_decode_tokens:]
                topk_indices = topk_indices[:num_prefill_tokens]
            else:
                assert attn_metadata is not None
                topk_indices = attn_metadata.c128a_prefill_topk_indices
            assert topk_indices is not None
            top_k = topk_indices.shape[-1]
        else:
            assert self.topk_indices_buffer is not None
            topk_indices = self.topk_indices_buffer[num_decode_tokens:]
            top_k = 0

        chunk_plan = swa_metadata.get_prefill_chunk_plan(
            compress_ratio=self.compress_ratio,
            prefill_chunk_size=self.PREFILL_CHUNK_SIZE,
        )
        assert chunk_plan, "prefill chunk plan must be non-empty when num_prefills > 0"
        workspace_manager = current_workspace_manager()
        for chunk_start, chunk_end, chunk_N, chunk_M in chunk_plan:
            chunk_size = chunk_end - chunk_start
            kv = workspace_manager.get_simultaneous(
                ((chunk_size, chunk_M, q.shape[-1]), torch.bfloat16),
            )[0]
            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,
                max_image_tokens=self.max_image_tokens,
                left_visible=(
                    left_visible[query_start:query_end]
                    if left_visible is not None
                    else None
                ),
                right_visible=(
                    right_visible[query_start:query_end]
                    if right_visible is not None
                    else None
                ),
            )
            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],
            )

    def _forward_prefill_aiter(
        self,
        q: torch.Tensor,
        compressed_k_cache: torch.Tensor | None,
        swa_k_cache: torch.Tensor,
        output: torch.Tensor,
        attn_metadata: DeepseekV4ROCMAiterMLASparseMetadata | None,
        swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    ) -> None:
        """Prefill straight off the paged caches, with no bf16 gather."""
        num_decode_tokens = swa_metadata.num_decode_tokens
        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_ragged_indices = None
        topk_ragged_indptr = None
        if attn_metadata is None:
            compressed_k_cache = None
        else:
            assert compressed_k_cache is not None
            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
            if self.compress_ratio == 4:
                assert self.topk_indices_buffer is not None
                topk_indices = self.topk_indices_buffer[
                    num_decode_tokens : num_decode_tokens + num_prefill_tokens
                ]
            else:
                topk_indices = attn_metadata.c128a_prefill_topk_indices
            assert topk_indices is not None
            topk_ragged_indices, topk_ragged_indptr = build_prefill_topk_ragged_indices(
                topk_indices,
                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],
            )

        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_ragged_indices,
            topk_ragged_indptr,
        )

    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.

_forward_decode(q, positions, kv_cache, swa_metadata, attn_metadata, swa_only, output, adaptive_splits)

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

Source code in vllm/models/deepseek_v4/amd/rocm.py
def _forward_decode(
    self,
    q: torch.Tensor,
    positions: torch.Tensor | None,
    kv_cache: torch.Tensor | None,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
    attn_metadata: DeepseekV4ROCMAiterMLASparseMetadata | None,
    swa_only: bool,
    output: torch.Tensor,
    adaptive_splits: bool,
) -> 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_indices = None
    topk_lens = None
    topk_ragged_indices = None
    topk_ragged_indptr = None
    if not swa_only:
        assert attn_metadata is not None
        assert swa_metadata.is_valid_token is not None
        block_size = attn_metadata.block_size // self.compress_ratio
        is_valid = swa_metadata.is_valid_token[:num_decode_tokens]
        if self.compress_ratio == 4:
            assert self.topk_indices_buffer is not None
            (
                topk_ragged_indices,
                topk_ragged_indptr,
                topk_lens,
            ) = 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],
                block_size,
                is_valid,
            )
        else:
            topk_indices = attn_metadata.c128a_global_decode_topk_indices
            topk_lens = attn_metadata.c128a_decode_topk_lens
            topk_ragged_indices = attn_metadata.c128a_decode_topk_ragged_indices
            topk_ragged_indptr = attn_metadata.c128a_decode_topk_ragged_indptr

    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
        self._aiter_sparse_mla(
            q,
            output,
            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,
        )
        # 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=topk_indices,
        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,
        adaptive_splits=adaptive_splits,
        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_aiter(q, compressed_k_cache, swa_k_cache, output, attn_metadata, swa_metadata)

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

Source code in vllm/models/deepseek_v4/amd/rocm.py
def _forward_prefill_aiter(
    self,
    q: torch.Tensor,
    compressed_k_cache: torch.Tensor | None,
    swa_k_cache: torch.Tensor,
    output: torch.Tensor,
    attn_metadata: DeepseekV4ROCMAiterMLASparseMetadata | None,
    swa_metadata: DeepseekV4ROCMAiterSparseSWAMetadata,
) -> None:
    """Prefill straight off the paged caches, with no bf16 gather."""
    num_decode_tokens = swa_metadata.num_decode_tokens
    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_ragged_indices = None
    topk_ragged_indptr = None
    if attn_metadata is None:
        compressed_k_cache = None
    else:
        assert compressed_k_cache is not None
        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
        if self.compress_ratio == 4:
            assert self.topk_indices_buffer is not None
            topk_indices = self.topk_indices_buffer[
                num_decode_tokens : num_decode_tokens + num_prefill_tokens
            ]
        else:
            topk_indices = attn_metadata.c128a_prefill_topk_indices
        assert topk_indices is not None
        topk_ragged_indices, topk_ragged_indptr = build_prefill_topk_ragged_indices(
            topk_indices,
            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],
        )

    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_ragged_indices,
        topk_ragged_indptr,
    )

_prepare_and_attn(hidden_states, qr, kv, qr_scale, kv_score, indexer_kv_score, indexer_weights, positions, o_padded)

Run the ROCm fork/join (HCA or CSA) inside the capture boundary.

Source code in vllm/models/deepseek_v4/amd/rocm.py
def _prepare_and_attn(
    self,
    hidden_states: torch.Tensor,
    qr: torch.Tensor | None,
    kv: torch.Tensor | None,
    qr_scale: torch.Tensor | None,
    kv_score: torch.Tensor | None,
    indexer_kv_score: torch.Tensor | None,
    indexer_weights: torch.Tensor | None,
    positions: torch.Tensor,
    o_padded: torch.Tensor,
) -> None:
    """Run the ROCm fork/join (HCA or CSA) inside the capture boundary."""
    aux_streams = self.aux_stream_list
    # The sequential pipeline disables aux_stream_list before calling
    # back with real projection inputs; aux_streams is None ends that
    # recursion here.
    if aux_streams is None:
        saved_streams = self.aux_stream_list
        self.aux_stream_list = None
        try:
            super()._prepare_and_attn(
                hidden_states,
                cast(torch.Tensor, qr),
                cast(torch.Tensor, kv),
                qr_scale,
                cast(torch.Tensor, kv_score),
                cast(torch.Tensor, indexer_kv_score),
                cast(torch.Tensor, indexer_weights),
                positions,
                o_padded,
            )
        finally:
            self.aux_stream_list = saved_streams
        return

    # Re-check: forward's gate ran inside a captured segment that
    # _prepare_and_attn_eager (MRV1) then broke, making this region eager.
    if not current_platform.enable_multi_stream_overlap(
        self.aux_stream_list, get_forward_context().attn_metadata
    ):
        self._run_sequential_pipeline(hidden_states, positions, o_padded)
        return

    indexer = self.indexer
    compressor = self.compressor
    assert compressor is not None

    def default_chain():
        qr_kv = self._fused_wqa_wkv_gemm(hidden_states)
        qr_out, qr_scale_out, kv_out = self._split_qkv_and_norm(qr_kv)
        q = self._wq_b_proj(qr_out, qr_scale_out).view(
            -1, self.n_local_heads, self.head_dim
        )
        attn_metadata = get_forward_context().attn_metadata
        q = self._fused_qnorm_rope_kv_insert(q, kv_out, positions, attn_metadata)
        return q, qr_out, qr_scale_out, kv_out

    def main_compressor_chain() -> None:
        score = torch.mm(
            hidden_states,
            compressor.fused_wkv_wgate.weight.T,
            out_dtype=torch.float32,
        )
        compressor(score, positions, self.rotary_emb)

    if indexer is None:
        # HCA dual-stream: the main compressor runs on aux stream 0 while
        # the default stream produces q and inserts KV into the SWA cache.
        # Both branches only read hidden_states, so the join merely has to
        # precede the sparse attention that consumes the compressed KV.
        (q, _qr_out, _qr_scale_out, kv_out), _ = execute_in_parallel(
            default_chain,
            [main_compressor_chain],
            self.ln_events[0],
            [self.ln_events[1]],
            aux_streams[:1],
            enable=True,
        )
        self._sparse_indexer_and_attn(
            hidden_states, None, None, None, q, kv_out, positions, o_padded
        )
        return

    def indexer_compressor_chain() -> None:
        score = torch.mm(
            hidden_states,
            indexer.compressor.fused_wkv_wgate.weight.T,
            out_dtype=torch.float32,
        )
        indexer.compressor(score, positions, self.indexer_rotary_emb)

    # CSA three-stream: the main and indexer compressors run on aux
    # streams 0 and 1 while the default stream produces q and inserts KV
    # into the SWA cache. Every branch only reads hidden_states plus its
    # own state, so the join merely has to precede the indexer op and the
    # sparse attention, which consume the compressed KV caches.
    (q, qr_out, qr_scale_out, kv_out), _ = execute_in_parallel(
        default_chain,
        [main_compressor_chain, indexer_compressor_chain],
        self.ln_events[0],
        self.ln_events[1:3],
        aux_streams[:2],
        enable=True,
    )

    indexer_weights_out, _ = indexer.weights_proj(hidden_states)
    # The indexer compressor already ran on aux stream 1; build queries only.
    index_q, index_q_scale, weights = indexer(
        hidden_states,
        qr_out,
        None,
        indexer_weights_out,
        positions,
        self.indexer_rotary_emb,
        qr_scale_out,
        skip_compressor=True,
    )
    self._sparse_indexer_and_attn(
        hidden_states,
        index_q,
        index_q_scale,
        weights,
        q,
        kv_out,
        positions,
        o_padded,
    )

_run_sequential_pipeline(hidden_states, positions, o_padded)

Disable ROCm streams when the current execution region cannot overlap.

Source code in vllm/models/deepseek_v4/amd/rocm.py
def _run_sequential_pipeline(
    self,
    hidden_states: torch.Tensor,
    positions: torch.Tensor,
    o_padded: torch.Tensor,
) -> None:
    """Disable ROCm streams when the current execution region cannot overlap."""
    aux_streams = self.aux_stream_list
    self.aux_stream_list = None
    try:
        qr_kv, kv_score, indexer_kv_score, indexer_weights = (
            self._run_parallel_input_projections(hidden_states)
        )
        qr, qr_scale, kv = self._split_qkv_and_norm(qr_kv)
        self._prepare_and_attn_fn(
            hidden_states,
            qr,
            kv,
            qr_scale,
            kv_score,
            indexer_kv_score,
            indexer_weights,
            positions,
            o_padded,
        )
    finally:
        self.aux_stream_list = aux_streams

_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_v4/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,
    )

DeepseekV4ROCMAiterMLASparseMetadata dataclass

Bases: DeepseekV4FlashMLAMetadata

ROCm-specific DeepSeek V4 metadata carrying ragged decode topk.

Source code in vllm/models/deepseek_v4/amd/rocm.py
@dataclass
class DeepseekV4ROCMAiterMLASparseMetadata(DeepseekV4FlashMLAMetadata):
    """ROCm-specific DeepSeek V4 metadata carrying ragged decode topk."""

    c128a_decode_topk_ragged_indices: torch.Tensor | None = None
    c128a_decode_topk_ragged_indptr: torch.Tensor | None = None
    for_cudagraph_capture: bool = False

_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_v4/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

_wo_a_block_scale_to_e8m0(scale)

Normalize checkpoint WO_A scales to raw OCP MX E8M0 bytes.

E8M0 is an unsigned exponent-only scale format with bias 127. A finite encoded byte b in [0, 254] represents 2 ** (b - 127); 0xFF is reserved for NaN. This is the vendor-neutral OCP encoding used by AMD AITER/OPUS, not an NVIDIA-specific convention. Reference: OCP Microscaling Formats (MX) Specification, section 5.4.1: https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf

Loaders may preserve the encoded byte as float8_e8m0fnu/uint8, or decode it to a floating-point power of two. Floating-point inputs are accepted only when the original E8M0 byte can be recovered losslessly. This function does not quantize or round arbitrary scales.

Source code in vllm/models/deepseek_v4/amd/rocm.py
def _wo_a_block_scale_to_e8m0(scale: torch.Tensor) -> torch.Tensor | None:
    """Normalize checkpoint WO_A scales to raw OCP MX E8M0 bytes.

    E8M0 is an unsigned exponent-only scale format with bias 127. A finite
    encoded byte ``b`` in ``[0, 254]`` represents ``2 ** (b - 127)``;
    ``0xFF`` is reserved for NaN. This is the vendor-neutral OCP encoding used
    by AMD AITER/OPUS, not an NVIDIA-specific convention. Reference: OCP
    Microscaling Formats (MX) Specification, section 5.4.1:
    https://www.opencompute.org/documents/ocp-microscaling-formats-mx-v1-0-spec-final-pdf

    Loaders may preserve the encoded byte as ``float8_e8m0fnu``/``uint8``,
    or decode it to a floating-point power of two. Floating-point inputs are
    accepted only when the original E8M0 byte can be recovered losslessly.
    This function does not quantize or round arbitrary scales.
    """
    if scale.dtype == torch.float8_e8m0fnu:
        # Reinterpret the native E8M0 storage. ``to(uint8)`` would perform a
        # numeric conversion instead of preserving the encoded exponent byte.
        return scale.view(torch.uint8).contiguous()
    if scale.dtype == torch.uint8:
        # The checkpoint loader already exposed the E8M0 wire representation.
        return scale.contiguous()
    if not scale.dtype.is_floating_point:
        return None

    scale_f32 = scale.detach().float()
    if not bool(torch.isfinite(scale_f32).all()) or bool((scale_f32 <= 0).any()):
        return None

    # With no mantissa, E8M0 can represent only exact powers of two. Rebuild
    # the value before encoding so this adapter never silently quantizes a
    # general floating-point checkpoint scale.
    exponent = torch.round(torch.log2(scale_f32))
    if not torch.equal(torch.exp2(exponent), scale_f32):
        return None

    encoded = exponent.to(torch.int32) + 127
    # 0xFF is NaN in OCP E8M0, not a finite exponent.
    if int(encoded.min()) < 0 or int(encoded.max()) > 254:
        return None
    return encoded.to(torch.uint8).contiguous()

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_v4/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)