Skip to content

vllm.v1.attention.backends.mla.indexer

Classes:

Functions:

DeepseekV32IndexerMetadataBuilder

Bases: AttentionMetadataBuilder

Source code in vllm/v1/attention/backends/mla/indexer.py
 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
1650
1651
1652
1653
1654
1655
1656
1657
1658
1659
1660
1661
1662
1663
1664
1665
1666
1667
1668
1669
class DeepseekV32IndexerMetadataBuilder(AttentionMetadataBuilder):
    # The indexer opts out of the shared reorder-threshold vote (see __init__),
    # so this is None; its own split uses self.decode_threshold.
    reorder_batch_threshold: int | None = None
    requires_block_table_width = True

    @classmethod
    def get_cudagraph_support(
        cls,
        vllm_config: VllmConfig,
        kv_cache_spec: KVCacheSpec,
    ) -> AttentionCGSupport:
        if _supports_varlen_paged_mqa_logits() or _use_flattening(vllm_config):
            return AttentionCGSupport.ALWAYS
        return AttentionCGSupport.UNIFORM_BATCH

    def __init__(self, *args, block_table_width: int, **kwargs) -> None:
        super().__init__(*args, **kwargs)
        scheduler_config = self.vllm_config.scheduler_config
        parallel_config = self.vllm_config.parallel_config
        self.dcp_world_size = parallel_config.decode_context_parallel_size
        self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0
        self.pcp_world_size = parallel_config.prefill_context_parallel_size
        self.use_pcp = self.pcp_world_size > 1
        self.pcp_rank = get_pcp_group().rank_in_group if self.use_pcp else 0
        self.cp_kv_cache_interleave_size = parallel_config.cp_kv_cache_interleave_size
        # KV compression (DeepseekV4). Default to 1 for no compression.
        self.compress_ratio = 1
        if isinstance(self.kv_cache_spec, MLAAttentionSpec):
            assert isinstance(self.kv_cache_spec.tokens_per_state, int)
            self.compress_ratio = self.kv_cache_spec.tokens_per_state
        # NOTE(Chen):an estimated max size of flattened_kv. Need to double check.
        # Counted in compressed rows, like the chunker's seq_lens and the
        # workspace.
        self.max_prefill_buffer_size = (
            get_max_prefill_buffer_size(self.vllm_config) // self.compress_ratio
        )
        self.num_speculative_tokens = (
            self.vllm_config.speculative_config.num_speculative_tokens
            if self.vllm_config.speculative_config
            else 0
        )
        self.indexer_uses_fp4 = dsa_indexer_uses_fp4(self.vllm_config)

        next_n = self.num_speculative_tokens + 1
        self.decode_threshold = next_n
        self.reorder_batch_threshold = None
        self.use_flattening = _use_flattening(self.vllm_config)
        self.supports_varlen = _supports_varlen_paged_mqa_logits()
        logger.info_once(
            "DSA indexer decode path: use_flattening=%s supports_varlen=%s "
            "(next_n=%d, use_fp4_cache=%s)",
            self.use_flattening,
            self.supports_varlen,
            next_n,
            self.indexer_uses_fp4,
        )

        sm_count = num_compute_units(self.device.index)
        self.num_sms = sm_count

        self.offsets_buffer = torch.arange(
            next_n, device=self.device, dtype=torch.int32
        )
        self.decode_lens_buffer = torch.zeros(
            (scheduler_config.max_num_batched_tokens,),
            dtype=torch.int32,
            device=self.device,
        )
        self.per_req_decode_lens_buffer = torch.zeros(
            (scheduler_config.max_num_batched_tokens,),
            dtype=torch.int32,
            device=self.device,
        )
        # Shared workspace for decode seq_lens. Native MTP views this as
        # (B, max_decode_len) at runtime, keeping context_lens contiguous even
        # when max_decode_len is smaller than next_n.
        self.decode_seq_lens_buffer = torch.zeros(
            (scheduler_config.max_num_batched_tokens,),
            dtype=torch.int32,
            device=self.device,
        )
        self.global_decode_seq_lens_buffer = torch.zeros(
            (scheduler_config.max_num_batched_tokens,),
            dtype=torch.int32,
            device=self.device,
        )
        self.decode_indices_buffer = torch.zeros(
            (scheduler_config.max_num_batched_tokens,),
            dtype=torch.int32,
            device=self.device,
        )
        self.arange_buffer = torch.arange(
            max(
                scheduler_config.max_num_seqs * next_n,
                scheduler_config.max_num_batched_tokens,
            ),
            dtype=torch.int32,
            device=self.device,
        )
        # Materialize the rank on device during builder initialization. Creating
        # this scalar in build() would introduce a GPU<->CPU sync in the decode
        # hot path.
        self.dcp_rank_tensor = torch.tensor(
            self.dcp_rank, dtype=torch.int32, device=self.device
        )
        self.expanded_block_table_buffer = torch.zeros(
            (scheduler_config.max_num_batched_tokens, block_table_width),
            dtype=torch.int32,
            device=self.device,
        )

        # See: DeepGMM/csrc/apis/attention.hpp. Sized for one slot per SM;
        # build() narrows it to whatever the kernel actually schedules.
        self.scheduler_metadata_buffer = torch.empty(
            (self.num_sms + 1, 2), dtype=torch.int32, device=self.device
        )

        if self.dcp_world_size > 1 and self.compress_ratio > 1:
            raise NotImplementedError(
                "DCP is not supported with sparse indexer KV compression "
                f"(compress_ratio={self.compress_ratio})."
            )

        # Pre-allocate buffers for CUDA graph compatibility when
        if self.compress_ratio > 1:
            # compress_ratio > 1 (DeepseekV4)
            # Compressed slot mapping output buffer
            self.compressed_slot_mapping_buffer = torch.zeros(
                (scheduler_config.max_num_batched_tokens,),
                dtype=torch.int64,
                device=self.device,
            )
            # Buffer for compressed seq_lens in decode path
            self.expanded_seq_lens_buffer = torch.zeros(
                (scheduler_config.max_num_batched_tokens,),
                dtype=torch.int32,
                device=self.device,
            )
        self.indexer_decode_block_table_buffer: torch.Tensor | None = None
        self._max_num_batched_tokens = scheduler_config.max_num_batched_tokens

    def _dcp_localize_decode_seq_lens(
        self,
        seq_lens: torch.Tensor,
        num_decodes: int,
        seq_lens_is_buffer_view: bool,
    ) -> torch.Tensor:
        local_seq_lens = get_dcp_local_seq_lens(
            seq_lens,
            self.dcp_world_size,
            self.dcp_rank_tensor,
            self.cp_kv_cache_interleave_size,
        )
        if seq_lens_is_buffer_view:
            seq_lens.copy_(local_seq_lens)
            return seq_lens

        out = self.decode_seq_lens_buffer[:num_decodes]
        out.copy_(local_seq_lens)
        return out

    def _prepare_decode_tensors(
        self,
        seq_lens: torch.Tensor,
        block_table: torch.Tensor,
        decode_lens: torch.Tensor,
        decode_lens_cpu: torch.Tensor,
        query_start_loc: torch.Tensor,
        num_decodes: int,
        num_decode_tokens: int,
        use_native: bool,
        next_n: int,
        max_decode_len: int,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, bool]:
        """Prepare native or per-token flattened decode tensors."""
        spec_config = self.vllm_config.speculative_config
        adaptive = bool(spec_config and spec_config.enable_adaptive_verification)
        min_decode_len = int(decode_lens_cpu.min().item())
        if not use_native:
            assert self.decode_seq_lens_buffer.dim() == 1
            if (
                not self.supports_varlen
                and (num_decodes == 1 or not adaptive)
                and min_decode_len == max_decode_len
                and num_decodes * max_decode_len == num_decode_tokens
            ):
                # Uniform decode lengths with no cudagraph token padding.
                _PREPARE_UNIFORM_DECODE_KERNEL(
                    seq_lens,
                    self.decode_seq_lens_buffer,
                    block_table,
                    self.expanded_block_table_buffer,
                    self.decode_lens_buffer,
                    num_decode_tokens,
                    max_decode_len,
                )
                self.decode_seq_lens_buffer[num_decode_tokens:] = 0
                seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]
                block_table = self.expanded_block_table_buffer[:num_decode_tokens]
                decode_lens = self.decode_lens_buffer[:num_decode_tokens]
                return seq_lens, block_table, decode_lens, num_decode_tokens, False
            else:
                # Variable decode lengths.
                # Assume 4 requests with seq_lens [10, 7, 12, 0] (the final req is
                # padding) and decode_lens [3, 1, 4, 0] in the below example comments.
                # The context lengths are therefore
                # [10-3, 7-1, 12-4, 0-0] = [7, 6, 8, 0].

                # 3 + 1 + 4 + 0 = 8
                actual_expanded = int(decode_lens_cpu.sum().item())

                # Fuse expanded_base and expanded_starts into a single
                # repeat_interleave:
                # seq_len_i = (context_start[b] - query_start_loc[b]) + arange[i] + 1
                # where context_start[b] = seq_lens[b] - decode_lens[b].
                # Example: offsets = [7-0, 6-3, 8-4, 0-8] = [7, 3, 4, -8]
                # expanded_offsets  = [7, 7, 7, 3, 4, 4, 4, 4]
                # result            = [8, 9, 10, 7, 9, 10, 11, 12]
                expanded_offsets = torch.repeat_interleave(
                    seq_lens - decode_lens - query_start_loc,
                    decode_lens,
                    output_size=actual_expanded,
                )

                # [8, 9, 10, 7, 9, 10, 11, 12, ...] where ... is unused buffer space
                self.decode_seq_lens_buffer[:actual_expanded] = (
                    expanded_offsets + self.arange_buffer[:actual_expanded] + 1
                )
                self.decode_seq_lens_buffer[actual_expanded:] = 0
                seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]

                # Give each of the flattened entries the same block table row as the
                # original request.
                self.expanded_block_table_buffer[:actual_expanded] = (
                    torch.repeat_interleave(
                        block_table, decode_lens, dim=0, output_size=actual_expanded
                    )
                )
                if actual_expanded < num_decode_tokens:
                    self.expanded_block_table_buffer[
                        actual_expanded:num_decode_tokens, 0
                    ] = 0
                block_table = self.expanded_block_table_buffer[:num_decode_tokens]

                # All reqs now have decode_len=1
                self.decode_lens_buffer[:num_decode_tokens] = 1
                decode_lens = self.decode_lens_buffer[:num_decode_tokens]
                return seq_lens, block_table, decode_lens, num_decode_tokens, False
        else:
            # Native path: plain decode (next_n==1) or spec decode
            # with 2D per-token context lengths (next_n > 1).
            #
            # When decode_lens are not truly uniform (e.g. some requests have
            # decode_len < next_n due to padding or short prefills), the simple
            # reshape in sparse_attn_indexer won't work. Use pack_seq_triton
            # (requires_padding) instead.
            requires_padding = min_decode_len != max_decode_len
            if use_native and next_n > 1:
                assert self.decode_seq_lens_buffer.dim() == 1
                # (B, max_decode_len): token j attends to
                # L - max_decode_len + j + 1 KV tokens.
                seq_lens_buffer = self.decode_seq_lens_buffer[
                    : num_decodes * max_decode_len
                ].view(num_decodes, max_decode_len)
                # Clamp at 0: padding requests have seq_len == 0, which would
                # otherwise make token 0 negative (next_n=2 gives 0-2+1+0 = -1).
                # Downstream kernels read these as uint32, turning -1 into ~4e9.
                seq_lens_buffer[:] = (
                    seq_lens.unsqueeze(1)
                    - max_decode_len
                    + 1
                    + self.offsets_buffer[:max_decode_len]
                ).clamp_(min=0)
                seq_lens = seq_lens_buffer
            return seq_lens, block_table, decode_lens, num_decodes, requires_padding

    def _prepare_global_decode_seq_lens(
        self,
        global_seq_lens: torch.Tensor | None,
        decode_lens: torch.Tensor,
        decode_lens_cpu: torch.Tensor,
        query_start_loc: torch.Tensor,
        num_decode_tokens: int,
        use_native: bool,
        max_decode_len: int,
    ) -> torch.Tensor | None:
        if global_seq_lens is None:
            return None
        if use_native or max_decode_len <= 1:
            return global_seq_lens

        actual_expanded = int(decode_lens_cpu.sum().item())
        if actual_expanded > 0:
            expanded_offsets = torch.repeat_interleave(
                global_seq_lens - decode_lens - query_start_loc,
                decode_lens,
                output_size=actual_expanded,
            )
            self.global_decode_seq_lens_buffer[:actual_expanded] = (
                expanded_offsets + self.arange_buffer[:actual_expanded] + 1
            )
        self.global_decode_seq_lens_buffer[actual_expanded:num_decode_tokens] = 0
        return self.global_decode_seq_lens_buffer[:num_decode_tokens]

    def _split_pcp_dcp_prefill_chunks(
        self,
        row_req_idx: np.ndarray,
        row_shard_rows: np.ndarray,
        row_query_lens_cpu: torch.Tensor,
        max_logits_bytes: int,
        request_offset: int,
    ) -> list[tuple[slice, slice]]:
        """Chunk by request rather than by row, so a split prefill's two rows
        charge their shared context once, then widen each chunk back to rows.

        ``row_shard_rows`` holds the whole request's largest DCP shard on
        every row, so the plan is identical on every PCP rank.
        """
        row_bounds = request_row_bounds(row_req_idx)
        first_rows = row_bounds[:-1]
        # Each request's context, padded to whole DCP shards.
        seq_lens = row_shard_rows[first_rows] * self.dcp_world_size
        # Every rank holds a full-size chunk of a split request (shorter ones
        # are replicated), so rows x longest row is the same on every rank,
        # whichever row holds the short tail.
        row_query_lens = row_query_lens_cpu.numpy()
        query_lens = np.diff(row_bounds) * np.maximum.reduceat(
            row_query_lens, first_rows
        )
        chunk_specs = self._split_indexer_prefill_chunks(
            torch.from_numpy(seq_lens.astype(np.int32)),
            torch.from_numpy(query_lens.astype(np.int32)),
            self.max_prefill_buffer_size,
            max_logits_bytes,
        )
        return [
            (
                slice(
                    request_offset + int(row_bounds[request_slice.start]),
                    request_offset + int(row_bounds[request_slice.stop]),
                ),
                query_slice,
            )
            for request_slice, query_slice in chunk_specs
        ]

    def _prefill_split_seq_lens(self, seq_lens_cpu: torch.Tensor) -> torch.Tensor:
        """Per-request KV lengths the prefill chunker budgets logits with;
        subclasses whose logits rows are wider than the context override."""
        return seq_lens_cpu

    @staticmethod
    def _split_indexer_prefill_chunks(
        compressed_seq_lens_cpu: torch.Tensor,
        prefill_query_lens_cpu: torch.Tensor,
        workspace_size: int,
        max_logits_bytes: int,
        request_offset: int = 0,
    ) -> list[tuple[slice, slice]]:
        """Split this step's prefill requests into chunks, respecting:
        - N constraint: total_seq_lens <= workspace_size (existing O(N)
          workspace)
        - Logits constraint: M * N * 4 <= max_logits_bytes

        When a single request-level chunk still exceeds the logits budget,
        sub-chunks on the query dimension (M) to bound peak memory.

        Returns list of (req_slice, query_slice) tuples.
        """
        chunks: list[tuple[slice, slice]] = []
        n = len(compressed_seq_lens_cpu)
        max_logits_elems = max_logits_bytes // 4
        end = 0

        while end < n:
            start, chunk_m, chunk_n = end, 0, 0

            while end < n:
                q, s = (
                    prefill_query_lens_cpu[end].item(),
                    compressed_seq_lens_cpu[end].item(),
                )
                new_m, new_n = chunk_m + q, chunk_n + s
                if new_n <= workspace_size and new_m * new_n <= max_logits_elems:
                    chunk_m, chunk_n = new_m, new_n
                    end += 1
                else:
                    break

            # A single request can exceed the budget, requiring sub-chunking
            # on the query dimension.
            if end == start:
                chunk_m, chunk_n = (
                    prefill_query_lens_cpu[end].item(),
                    compressed_seq_lens_cpu[end].item(),
                )
                end += 1

            req_slice = slice(start + request_offset, end + request_offset)
            max_q = (
                max(1, max_logits_elems // chunk_n) if chunk_n > 0 else max(1, chunk_m)
            )
            for q_off in range(0, chunk_m, max_q):
                sub_m = min(max_q, chunk_m - q_off)
                chunks.append((req_slice, slice(q_off, q_off + sub_m)))

        return chunks

    def build(
        self,
        common_prefix_len: int,
        common_attn_metadata: CommonAttentionMetadata,
        fast_build: bool = False,
    ) -> DeepseekV32IndexerMetadata:
        num_reqs = common_attn_metadata.num_reqs
        num_tokens = common_attn_metadata.num_actual_tokens
        query_start_loc = common_attn_metadata.query_start_loc
        query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu
        seq_lens = common_attn_metadata.seq_lens
        slot_mapping = common_attn_metadata.slot_mapping
        block_table = common_attn_metadata.block_table_tensor
        dcp_local_seq_lens = common_attn_metadata.dcp_local_seq_lens

        compressed_slot_mapping = slot_mapping
        indexer_block_table = block_table
        if self.compress_ratio > 1:
            kernel_block_size = self.kernel_block_size
            if (
                kernel_block_size is not None
                and self.kv_cache_spec.block_size != kernel_block_size
                and self.kv_cache_spec.block_size % kernel_block_size == 0
            ):
                factor = self.kv_cache_spec.block_size // kernel_block_size
                indexer_block_table = (block_table[:, ::factor] // factor).contiguous()
            padded_num_tokens = num_tokens
            local_slot_mapping = slot_mapping
            if self.use_pcp:
                # The gathered layout holds each rank's local tokens, padded, in
                # rank order, so this rank's segment lines up with query_start_loc.
                padded_num_tokens = slot_mapping.shape[0] // self.pcp_world_size
                local_slot_mapping = slot_mapping[
                    self.pcp_rank * padded_num_tokens : (self.pcp_rank + 1)
                    * padded_num_tokens
                ]
            compressed_slot_mapping = get_compressed_slot_mapping(
                num_tokens,
                local_slot_mapping,
                query_start_loc,
                seq_lens,
                indexer_block_table,
                self.kv_cache_spec.num_states,
                self.compress_ratio,
                out=self.compressed_slot_mapping_buffer,
            )
            if self.pcp_world_size > 1:
                compressed_slot_mapping = get_pcp_group().all_gather(
                    self.compressed_slot_mapping_buffer[:padded_num_tokens],
                    dim=0,
                )

        # PCP decode sharding keeps a zero-token placeholder on ranks that own
        # no request in a step so collectives retain a uniform rank shape. Do
        # not turn that placeholder into a zero-length indexer decode request.
        if num_tokens == 0:
            return DeepseekV32IndexerMetadata(
                seq_lens=seq_lens,
                max_seq_len=common_attn_metadata.max_seq_len,
                slot_mapping=compressed_slot_mapping,
                num_decodes=0,
                num_decode_tokens=0,
                num_prefills=0,
                num_prefill_tokens=0,
                prefill=None,
                decode=None,
            )

        num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = (
            split_decodes_and_prefills(
                common_attn_metadata,
                decode_threshold=self.decode_threshold,
                require_uniform=not (self.use_flattening or self.supports_varlen),
                treat_short_extends_as_decodes=not self.use_pcp,
            )
        )

        assert num_decodes + num_prefills == num_reqs
        assert num_decode_tokens + num_prefill_tokens == num_tokens

        prefill_metadata = None
        if num_prefills > 0:
            compressed_seq_lens = (
                seq_lens // self.compress_ratio if self.compress_ratio > 1 else seq_lens
            )
            # This CPU value is an upper bound for async-spec extend rows.  It
            # is safe for chunking/allocation because CUDA metadata below is
            # built from exact device seq_lens and gather ignores the tail.
            assert common_attn_metadata.seq_lens_cpu_upper_bound is not None
            seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
            compressed_seq_lens_cpu = (
                seq_lens_cpu // self.compress_ratio
                if self.compress_ratio > 1
                else seq_lens_cpu
            )
            prefill_query_lens_cpu = torch.diff(
                query_start_loc_cpu[num_decodes : num_decodes + num_prefills + 1]
            )
            max_logits_bytes = envs.VLLM_SPARSE_INDEXER_MAX_LOGITS_MB * 1024 * 1024
            # Upper bound is exact for prefill rows (the `[num_decodes:]`
            # slice below).
            assert common_attn_metadata.seq_lens_cpu_upper_bound is not None
            seq_lens_cpu = common_attn_metadata.seq_lens_cpu_upper_bound
            req_idx = None
            shard_rows = None
            if self.use_pcp and self.dcp_world_size > 1:
                # The gathered KV must be packed identically on every PCP rank:
                # chunk by request from its DCP shard rows, which every rank
                # holds. A dummy batch bypasses the PCP manager and has one
                # row per request, so its own extent is the request's.
                req_idx = common_attn_metadata.req_idx
                if req_idx is None:
                    req_idx = np.arange(num_reqs)
                shard_rows_cpu = common_attn_metadata.dcp_local_seq_lens_cpu_upper_bound
                if shard_rows_cpu is None:
                    shard_rows_cpu = get_dcp_local_seq_lens(
                        seq_lens_cpu,
                        self.dcp_world_size,
                        0,
                        self.cp_kv_cache_interleave_size,
                    )
                shard_rows = shard_rows_cpu.numpy()
                chunk_specs = self._split_pcp_dcp_prefill_chunks(
                    req_idx[num_decodes:],
                    shard_rows[num_decodes:],
                    prefill_query_lens_cpu,
                    max_logits_bytes,
                    request_offset=num_decodes,
                )
            else:
                chunk_specs = self._split_indexer_prefill_chunks(
                    self._prefill_split_seq_lens(compressed_seq_lens_cpu[num_decodes:]),
                    prefill_query_lens_cpu,
                    self.max_prefill_buffer_size,
                    max_logits_bytes,
                    request_offset=num_decodes,
                )

            chunks = []
            for req_slice, query_slice in chunk_specs:
                pcp_plan = None
                if req_idx is not None:
                    assert shard_rows is not None
                    pcp_plan = build_pcp_global_chunk_plan(
                        req_idx[req_slice],
                        shard_rows[req_slice],
                        self.dcp_world_size,
                        self.device,
                        self.cp_kv_cache_interleave_size,
                    )
                metadata = build_prefill_chunk_metadata(
                    req_slice.start,
                    req_slice.stop,
                    query_start_loc,
                    query_start_loc_cpu,
                    seq_lens,
                    compressed_seq_lens,
                    compressed_seq_lens_cpu,
                    indexer_block_table,
                    self.compress_ratio,
                    query_slice=query_slice,
                    skip_kv_gather=query_slice.start > 0,
                    dcp_rank=self.dcp_rank,
                    dcp_world_size=self.dcp_world_size,
                    cp_kv_cache_interleave_size=self.cp_kv_cache_interleave_size,
                    pcp_plan=pcp_plan,
                )
                # Skip when total_seq_lens is 0 (i.e., no compressed token).
                if metadata is not None:
                    chunks.append(metadata)
            prefill_metadata = DeepseekV32IndexerPrefillMetadata(
                chunks,
                max_prefill_seq_len=(
                    int(seq_lens_cpu[num_decodes:].max().item())
                    if num_prefills > 0
                    else 0
                ),
            )

        decode_metadata = None
        if num_decodes > 0:
            if not self.supports_varlen:
                torch.diff(
                    common_attn_metadata.query_start_loc[: num_decodes + 1],
                    out=self.decode_lens_buffer[:num_decodes],
                )
                self.per_req_decode_lens_buffer[:num_decodes].copy_(
                    self.decode_lens_buffer[:num_decodes]
                )
            decode_lens = self.decode_lens_buffer[:num_decodes]
            decode_lens_cpu = torch.diff(
                common_attn_metadata.query_start_loc_cpu[: num_decodes + 1]
            )

            # Under DCP the per-token decode bounds must be localized AFTER the
            # per-token expansion below, not before. Expanding from a
            # request-level localized length subtracts decode offsets in local
            # space and yields too-short bounds (e.g. world=2, rank=1, global
            # per-token bounds [8, 9, 10] -> [3, 4, 5] instead of [4, 4, 5]), so
            # the first decode token would run top-k against too short a local KV
            # range and miss valid tokens. Keep the global seq_lens here and
            # localize the expanded bounds further down.
            global_seq_lens_for_decode: torch.Tensor | None = None
            if dcp_local_seq_lens is not None:
                global_seq_lens_for_decode = common_attn_metadata.seq_lens[:num_decodes]
            seq_lens = common_attn_metadata.seq_lens[:num_decodes]
            block_table = common_attn_metadata.block_table_tensor[:num_decodes, ...]

            max_decode_len = int(decode_lens_cpu.max().item())
            min_decode_len = int(decode_lens_cpu.min().item())
            write_is_uniform = min_decode_len == max_decode_len
            next_n = 1 + self.num_speculative_tokens
            # The kernel sees max_decode_len Q rows, not the configured next_n,
            # so legality is per-step: on SM90 a uniformly 3-deep batch has no
            # native kernel. max_decode_len <= 1 always has one.
            step_next_n_ok = max_decode_len <= 1 or _supports_native_decode(
                max_decode_len
            )
            use_native = (
                not (self.use_flattening or self.supports_varlen)
                and max_decode_len <= next_n
                and step_next_n_ok
            )

            if not self.supports_varlen:
                global_seq_lens_for_decode = self._prepare_global_decode_seq_lens(
                    global_seq_lens=global_seq_lens_for_decode,
                    decode_lens=decode_lens,
                    decode_lens_cpu=decode_lens_cpu,
                    query_start_loc=common_attn_metadata.query_start_loc[:num_decodes],
                    num_decode_tokens=num_decode_tokens,
                    use_native=use_native,
                    max_decode_len=max_decode_len,
                )

            decode_indices = None
            if self.supports_varlen:
                from vllm.v1.attention.ops.metadata import (
                    _indexer_decode_metadata_kernel,
                )

                capacity = self.decode_seq_lens_buffer.numel()
                grid = max(
                    num_decodes,
                    num_decode_tokens + triton.cdiv(capacity - num_decode_tokens, 256),
                )
                _indexer_decode_metadata_kernel[(grid,)](
                    query_start_loc,
                    seq_lens,
                    block_table,
                    self.decode_seq_lens_buffer,
                    self.expanded_block_table_buffer,
                    self.decode_lens_buffer,
                    self.decode_indices_buffer,
                    self.per_req_decode_lens_buffer,
                    num_decodes,
                    num_decode_tokens,
                    capacity,
                    block_table.stride(0),
                    self.expanded_block_table_buffer.stride(0),
                    BLOCK_COLS=block_table.shape[1],
                    num_warps=4,
                )
                seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]
                block_table = self.expanded_block_table_buffer[:num_decode_tokens]
                decode_lens = self.decode_lens_buffer[:num_decode_tokens]
                decode_indices = self.decode_indices_buffer[:num_decode_tokens]
                requires_padding = False
                if global_seq_lens_for_decode is not None and max_decode_len > 1:
                    self.global_decode_seq_lens_buffer[:num_decode_tokens].copy_(
                        seq_lens
                    )
                    global_seq_lens_for_decode = self.global_decode_seq_lens_buffer[
                        :num_decode_tokens
                    ]
            else:
                seq_lens, block_table, decode_lens, batch_size, requires_padding = (
                    self._prepare_decode_tensors(
                        seq_lens=seq_lens,
                        block_table=block_table,
                        decode_lens=decode_lens,
                        decode_lens_cpu=decode_lens_cpu,
                        query_start_loc=common_attn_metadata.query_start_loc[
                            :num_decodes
                        ],
                        num_decodes=num_decodes,
                        num_decode_tokens=num_decode_tokens,
                        use_native=use_native,
                        next_n=next_n,
                        max_decode_len=max_decode_len,
                    )
                )

            if self.compress_ratio > 1:
                kernel_block_size = self.kernel_block_size
                if (
                    kernel_block_size is not None
                    and self.kv_cache_spec.block_size != kernel_block_size
                    and self.kv_cache_spec.block_size % kernel_block_size == 0
                ):
                    factor = self.kv_cache_spec.block_size // kernel_block_size
                    compressed = block_table[:, ::factor] // factor
                    rows, cols = compressed.shape
                    if self.indexer_decode_block_table_buffer is None:
                        self.indexer_decode_block_table_buffer = torch.zeros(
                            (self._max_num_batched_tokens, cols),
                            dtype=torch.int32,
                            device=self.device,
                        )
                    self.indexer_decode_block_table_buffer[:rows, :cols].copy_(
                        compressed
                    )
                    block_table = self.indexer_decode_block_table_buffer[:rows, :cols]

            # Flattening always returns a buffer view, including single-token
            # batches. Keep its address stable across varlen graph replays.
            seq_lens_is_buffer_view = not use_native or next_n > 1

            # DCP: localize the now-expanded per-token global bounds to this
            # rank's owned KV. Done here (after expansion) so each token's global
            # causal length is localized individually; see the comment above.
            if dcp_local_seq_lens is not None:
                seq_lens = self._dcp_localize_decode_seq_lens(
                    seq_lens, num_decodes, seq_lens_is_buffer_view
                )

            # For DeepseekV4 (compress_ratio > 1), the indexer KV cache stores
            # compressed tokens. Convert uncompressed seq_lens to compressed.
            if self.compress_ratio > 1:
                if seq_lens_is_buffer_view:
                    seq_lens //= self.compress_ratio
                else:
                    # Copy to avoid mutating shared state; keeps CG address stable.
                    self.expanded_seq_lens_buffer[:num_decodes] = (
                        seq_lens // self.compress_ratio
                    )
                    self.expanded_seq_lens_buffer[num_decodes:num_decode_tokens] = 0
                    seq_lens = self.expanded_seq_lens_buffer[:num_decode_tokens]

            # Non-MTP: deep_gemm paged MQA logits requires 2D context_lens
            # (csrc/apis/attention.hpp). Unsqueeze to (B, 1) so downstream
            # kernels see the same (B, next_n) layout as the MTP path.
            if seq_lens.dim() == 1:
                seq_lens = seq_lens.unsqueeze(-1)

            # DeepGEMM is required for the paged MQA logits on CUDA devices
            schedule_metadata = self.scheduler_metadata_buffer
            if current_platform.is_cuda() and has_deep_gemm():
                metadata = get_paged_mqa_logits_metadata(
                    seq_lens,
                    self.kv_cache_spec.num_states,
                    self.num_sms,
                    indices=decode_indices,
                )
                schedule_metadata = self.scheduler_metadata_buffer[: metadata.shape[0]]
                schedule_metadata[:] = metadata

            decode_metadata = DeepSeekV32IndexerDecodeMetadata(
                block_table=block_table,
                seq_lens=seq_lens,
                decode_lens=decode_lens,
                requires_padding=requires_padding,
                schedule_metadata=schedule_metadata,
                indices=decode_indices,
                global_seq_lens=global_seq_lens_for_decode,
                per_req_decode_lens=self.per_req_decode_lens_buffer[:num_decodes],
                decode_is_uniform=write_is_uniform,
                write_max_decode_len=max_decode_len,
            )

        attn_metadata = DeepseekV32IndexerMetadata(
            seq_lens=common_attn_metadata.seq_lens,
            max_seq_len=common_attn_metadata.max_seq_len,
            slot_mapping=compressed_slot_mapping,
            num_decodes=num_decodes,
            num_decode_tokens=num_decode_tokens,
            num_prefills=num_prefills,
            num_prefill_tokens=num_prefill_tokens,
            prefill=prefill_metadata,
            decode=decode_metadata,
        )

        return attn_metadata

_prefill_split_seq_lens(seq_lens_cpu)

Per-request KV lengths the prefill chunker budgets logits with; subclasses whose logits rows are wider than the context override.

Source code in vllm/v1/attention/backends/mla/indexer.py
def _prefill_split_seq_lens(self, seq_lens_cpu: torch.Tensor) -> torch.Tensor:
    """Per-request KV lengths the prefill chunker budgets logits with;
    subclasses whose logits rows are wider than the context override."""
    return seq_lens_cpu

_prepare_decode_tensors(seq_lens, block_table, decode_lens, decode_lens_cpu, query_start_loc, num_decodes, num_decode_tokens, use_native, next_n, max_decode_len)

Prepare native or per-token flattened decode tensors.

Source code in vllm/v1/attention/backends/mla/indexer.py
def _prepare_decode_tensors(
    self,
    seq_lens: torch.Tensor,
    block_table: torch.Tensor,
    decode_lens: torch.Tensor,
    decode_lens_cpu: torch.Tensor,
    query_start_loc: torch.Tensor,
    num_decodes: int,
    num_decode_tokens: int,
    use_native: bool,
    next_n: int,
    max_decode_len: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, bool]:
    """Prepare native or per-token flattened decode tensors."""
    spec_config = self.vllm_config.speculative_config
    adaptive = bool(spec_config and spec_config.enable_adaptive_verification)
    min_decode_len = int(decode_lens_cpu.min().item())
    if not use_native:
        assert self.decode_seq_lens_buffer.dim() == 1
        if (
            not self.supports_varlen
            and (num_decodes == 1 or not adaptive)
            and min_decode_len == max_decode_len
            and num_decodes * max_decode_len == num_decode_tokens
        ):
            # Uniform decode lengths with no cudagraph token padding.
            _PREPARE_UNIFORM_DECODE_KERNEL(
                seq_lens,
                self.decode_seq_lens_buffer,
                block_table,
                self.expanded_block_table_buffer,
                self.decode_lens_buffer,
                num_decode_tokens,
                max_decode_len,
            )
            self.decode_seq_lens_buffer[num_decode_tokens:] = 0
            seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]
            block_table = self.expanded_block_table_buffer[:num_decode_tokens]
            decode_lens = self.decode_lens_buffer[:num_decode_tokens]
            return seq_lens, block_table, decode_lens, num_decode_tokens, False
        else:
            # Variable decode lengths.
            # Assume 4 requests with seq_lens [10, 7, 12, 0] (the final req is
            # padding) and decode_lens [3, 1, 4, 0] in the below example comments.
            # The context lengths are therefore
            # [10-3, 7-1, 12-4, 0-0] = [7, 6, 8, 0].

            # 3 + 1 + 4 + 0 = 8
            actual_expanded = int(decode_lens_cpu.sum().item())

            # Fuse expanded_base and expanded_starts into a single
            # repeat_interleave:
            # seq_len_i = (context_start[b] - query_start_loc[b]) + arange[i] + 1
            # where context_start[b] = seq_lens[b] - decode_lens[b].
            # Example: offsets = [7-0, 6-3, 8-4, 0-8] = [7, 3, 4, -8]
            # expanded_offsets  = [7, 7, 7, 3, 4, 4, 4, 4]
            # result            = [8, 9, 10, 7, 9, 10, 11, 12]
            expanded_offsets = torch.repeat_interleave(
                seq_lens - decode_lens - query_start_loc,
                decode_lens,
                output_size=actual_expanded,
            )

            # [8, 9, 10, 7, 9, 10, 11, 12, ...] where ... is unused buffer space
            self.decode_seq_lens_buffer[:actual_expanded] = (
                expanded_offsets + self.arange_buffer[:actual_expanded] + 1
            )
            self.decode_seq_lens_buffer[actual_expanded:] = 0
            seq_lens = self.decode_seq_lens_buffer[:num_decode_tokens]

            # Give each of the flattened entries the same block table row as the
            # original request.
            self.expanded_block_table_buffer[:actual_expanded] = (
                torch.repeat_interleave(
                    block_table, decode_lens, dim=0, output_size=actual_expanded
                )
            )
            if actual_expanded < num_decode_tokens:
                self.expanded_block_table_buffer[
                    actual_expanded:num_decode_tokens, 0
                ] = 0
            block_table = self.expanded_block_table_buffer[:num_decode_tokens]

            # All reqs now have decode_len=1
            self.decode_lens_buffer[:num_decode_tokens] = 1
            decode_lens = self.decode_lens_buffer[:num_decode_tokens]
            return seq_lens, block_table, decode_lens, num_decode_tokens, False
    else:
        # Native path: plain decode (next_n==1) or spec decode
        # with 2D per-token context lengths (next_n > 1).
        #
        # When decode_lens are not truly uniform (e.g. some requests have
        # decode_len < next_n due to padding or short prefills), the simple
        # reshape in sparse_attn_indexer won't work. Use pack_seq_triton
        # (requires_padding) instead.
        requires_padding = min_decode_len != max_decode_len
        if use_native and next_n > 1:
            assert self.decode_seq_lens_buffer.dim() == 1
            # (B, max_decode_len): token j attends to
            # L - max_decode_len + j + 1 KV tokens.
            seq_lens_buffer = self.decode_seq_lens_buffer[
                : num_decodes * max_decode_len
            ].view(num_decodes, max_decode_len)
            # Clamp at 0: padding requests have seq_len == 0, which would
            # otherwise make token 0 negative (next_n=2 gives 0-2+1+0 = -1).
            # Downstream kernels read these as uint32, turning -1 into ~4e9.
            seq_lens_buffer[:] = (
                seq_lens.unsqueeze(1)
                - max_decode_len
                + 1
                + self.offsets_buffer[:max_decode_len]
            ).clamp_(min=0)
            seq_lens = seq_lens_buffer
        return seq_lens, block_table, decode_lens, num_decodes, requires_padding

_split_indexer_prefill_chunks(compressed_seq_lens_cpu, prefill_query_lens_cpu, workspace_size, max_logits_bytes, request_offset=0) staticmethod

Split this step's prefill requests into chunks, respecting: - N constraint: total_seq_lens <= workspace_size (existing O(N) workspace) - Logits constraint: M * N * 4 <= max_logits_bytes

When a single request-level chunk still exceeds the logits budget, sub-chunks on the query dimension (M) to bound peak memory.

Returns list of (req_slice, query_slice) tuples.

Source code in vllm/v1/attention/backends/mla/indexer.py
@staticmethod
def _split_indexer_prefill_chunks(
    compressed_seq_lens_cpu: torch.Tensor,
    prefill_query_lens_cpu: torch.Tensor,
    workspace_size: int,
    max_logits_bytes: int,
    request_offset: int = 0,
) -> list[tuple[slice, slice]]:
    """Split this step's prefill requests into chunks, respecting:
    - N constraint: total_seq_lens <= workspace_size (existing O(N)
      workspace)
    - Logits constraint: M * N * 4 <= max_logits_bytes

    When a single request-level chunk still exceeds the logits budget,
    sub-chunks on the query dimension (M) to bound peak memory.

    Returns list of (req_slice, query_slice) tuples.
    """
    chunks: list[tuple[slice, slice]] = []
    n = len(compressed_seq_lens_cpu)
    max_logits_elems = max_logits_bytes // 4
    end = 0

    while end < n:
        start, chunk_m, chunk_n = end, 0, 0

        while end < n:
            q, s = (
                prefill_query_lens_cpu[end].item(),
                compressed_seq_lens_cpu[end].item(),
            )
            new_m, new_n = chunk_m + q, chunk_n + s
            if new_n <= workspace_size and new_m * new_n <= max_logits_elems:
                chunk_m, chunk_n = new_m, new_n
                end += 1
            else:
                break

        # A single request can exceed the budget, requiring sub-chunking
        # on the query dimension.
        if end == start:
            chunk_m, chunk_n = (
                prefill_query_lens_cpu[end].item(),
                compressed_seq_lens_cpu[end].item(),
            )
            end += 1

        req_slice = slice(start + request_offset, end + request_offset)
        max_q = (
            max(1, max_logits_elems // chunk_n) if chunk_n > 0 else max(1, chunk_m)
        )
        for q_off in range(0, chunk_m, max_q):
            sub_m = min(max_q, chunk_m - q_off)
            chunks.append((req_slice, slice(q_off, q_off + sub_m)))

    return chunks

_split_pcp_dcp_prefill_chunks(row_req_idx, row_shard_rows, row_query_lens_cpu, max_logits_bytes, request_offset)

Chunk by request rather than by row, so a split prefill's two rows charge their shared context once, then widen each chunk back to rows.

row_shard_rows holds the whole request's largest DCP shard on every row, so the plan is identical on every PCP rank.

Source code in vllm/v1/attention/backends/mla/indexer.py
def _split_pcp_dcp_prefill_chunks(
    self,
    row_req_idx: np.ndarray,
    row_shard_rows: np.ndarray,
    row_query_lens_cpu: torch.Tensor,
    max_logits_bytes: int,
    request_offset: int,
) -> list[tuple[slice, slice]]:
    """Chunk by request rather than by row, so a split prefill's two rows
    charge their shared context once, then widen each chunk back to rows.

    ``row_shard_rows`` holds the whole request's largest DCP shard on
    every row, so the plan is identical on every PCP rank.
    """
    row_bounds = request_row_bounds(row_req_idx)
    first_rows = row_bounds[:-1]
    # Each request's context, padded to whole DCP shards.
    seq_lens = row_shard_rows[first_rows] * self.dcp_world_size
    # Every rank holds a full-size chunk of a split request (shorter ones
    # are replicated), so rows x longest row is the same on every rank,
    # whichever row holds the short tail.
    row_query_lens = row_query_lens_cpu.numpy()
    query_lens = np.diff(row_bounds) * np.maximum.reduceat(
        row_query_lens, first_rows
    )
    chunk_specs = self._split_indexer_prefill_chunks(
        torch.from_numpy(seq_lens.astype(np.int32)),
        torch.from_numpy(query_lens.astype(np.int32)),
        self.max_prefill_buffer_size,
        max_logits_bytes,
    )
    return [
        (
            slice(
                request_offset + int(row_bounds[request_slice.start]),
                request_offset + int(row_bounds[request_slice.stop]),
            ),
            query_slice,
        )
        for request_slice, query_slice in chunk_specs
    ]

KpoolTailBackend

Bases: DeepseekV32IndexerBackend

Storage-only backend for the GLM-5.3-Flash kpool tail cache.

Source code in vllm/v1/attention/backends/mla/indexer.py
class KpoolTailBackend(DeepseekV32IndexerBackend):
    """Storage-only backend for the GLM-5.3-Flash kpool tail cache."""

    @classmethod
    def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...]:
        return (KVCacheLayout.LBHNC,)

    @staticmethod
    def get_name() -> str:
        return "KPOOL_TAIL"

    @classmethod
    def get_supported_head_sizes(cls) -> list[int]:
        return []

    @staticmethod
    def get_supported_kernel_block_sizes(kv_cache_spec=None) -> list[int | MultipleOf]:
        return [MultipleOf(1)]

    @staticmethod
    def get_builder_cls() -> type["KpoolTailMetadataBuilder"]:  # type: ignore[override]
        return KpoolTailMetadataBuilder

KpoolTailMetadataBuilder

Bases: AttentionMetadataBuilder

Build only the circular slot mapping needed by the storage-only tail.

Source code in vllm/v1/attention/backends/mla/indexer.py
class KpoolTailMetadataBuilder(AttentionMetadataBuilder):
    """Build only the circular slot mapping needed by the storage-only tail."""

    _cudagraph_support = AttentionCGSupport.ALWAYS
    supports_update_block_table = False
    reorder_batch_threshold = None

    def __init__(
        self,
        kv_cache_spec: AttentionSpec,
        layer_names: list[str],
        vllm_config: VllmConfig,
        device: torch.device,
    ):
        super().__init__(kv_cache_spec, layer_names, vllm_config, device)
        self.slot_mapping_buffer = torch.empty(
            vllm_config.scheduler_config.max_num_batched_tokens,
            dtype=torch.int64,
            device=device,
        )

    def build(
        self,
        common_prefix_len: int,
        common_attn_metadata: CommonAttentionMetadata,
        fast_build: bool = False,
    ) -> DeepseekV32IndexerMetadata:
        num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = (
            split_decodes_and_prefills(common_attn_metadata)
        )
        slot_mapping = common_attn_metadata.slot_mapping
        positions = common_attn_metadata.positions
        if positions is not None:
            slot_mapping_buffer = self.slot_mapping_buffer[
                : slot_mapping.numel()
            ].view_as(slot_mapping)
            slot_mapping = compute_kpool_tail_slot_mapping(
                slot_mapping,
                common_attn_metadata.block_table_tensor,
                common_attn_metadata.query_start_loc,
                positions,
                common_attn_metadata.num_actual_tokens,
                common_attn_metadata.num_reqs,
                self.kv_cache_spec.block_size,
                out=slot_mapping_buffer,
            )
        return DeepseekV32IndexerMetadata(
            seq_lens=common_attn_metadata.seq_lens,
            max_seq_len=common_attn_metadata.max_seq_len,
            slot_mapping=slot_mapping,
            num_decodes=num_decodes,
            num_decode_tokens=num_decode_tokens,
            num_prefills=num_prefills,
            num_prefill_tokens=num_prefill_tokens,
        )

PCPGlobalChunkPlan dataclass

PCP packing for one indexer prefill chunk under PCP + DCP.

Source code in vllm/v1/attention/backends/mla/indexer.py
@dataclass(frozen=True)
class PCPGlobalChunkPlan:
    """PCP packing for one indexer prefill chunk under PCP + DCP."""

    # [num_reqs+1] cumsum of the padded per-request context (shard rows x W):
    # the global layout.
    row_start_cu: torch.Tensor
    # [num_reqs+1] like row_start_cu, but only a region's FIRST row carries its
    # extent.
    global_cu: torch.Tensor
    # [num_reqs+1] cumsum of shard rows: this rank's padded layout.
    padded_local_cu: torch.Tensor
    padded_local_total: int
    total: int
    # [total] for each global row, its position in the rank-major gathered buffer.
    deinterleave_idx: torch.Tensor

_supports_native_decode(next_n)

Whether decode can pass next_n Q rows per request to the kernel instead of flattening to one single-token row per query, which re-reads the KV tile once per row.

Source code in vllm/v1/attention/backends/mla/indexer.py
def _supports_native_decode(next_n: int) -> bool:
    """Whether decode can pass `next_n` Q rows per request to the kernel
    instead of flattening to one single-token row per query, which re-reads
    the KV tile once per row.
    """
    if not (current_platform.is_cuda() and has_deep_gemm()):
        return next_n in (1, 2)
    if current_platform.is_device_capability_family(100):
        return True
    if current_platform.is_device_capability_family(90):
        return native_next_n_supported(next_n)
    return next_n in (1, 2)

build_pcp_global_chunk_plan(row_req_idx, row_shard_rows, dcp_world_size, device, interleave=1)

Plan the PCP packing for one chunk from its requests' KV shard rows.

The global layout is padded to shard rows x W per request. Positions past a request's true extent lie beyond every token's causal bound, so the kernel never reads them.

Source code in vllm/v1/attention/backends/mla/indexer.py
def build_pcp_global_chunk_plan(
    row_req_idx: np.ndarray,
    row_shard_rows: np.ndarray,
    dcp_world_size: int,
    device: torch.device,
    interleave: int = 1,
) -> PCPGlobalChunkPlan:
    """Plan the PCP packing for one chunk from its requests' KV shard rows.

    The global layout is padded to ``shard rows x W`` per request. Positions
    past a request's true extent lie beyond every token's causal bound, so
    the kernel never reads them.
    """
    shard_rows = np.ascontiguousarray(row_shard_rows, dtype=np.int64)
    assert row_req_idx.shape == shard_rows.shape
    num_rows = len(shard_rows)

    row_bounds = request_row_bounds(row_req_idx)
    region_first_row = row_bounds[:-1]
    region_of_row = np.repeat(np.arange(len(region_first_row)), np.diff(row_bounds))

    region_padded = shard_rows[region_first_row]
    assert np.all(region_padded > 0), (
        f"PCP+DCP prefill got an empty context: {region_padded.tolist()}"
    )

    # Pinned host buffers, filled in place so the copies below are async.
    cu = torch.zeros((3, num_rows + 1), dtype=torch.int32, pin_memory=PIN_MEMORY)
    row_start_cu, global_cu, padded_cu = cu
    # Only a region's first row carries its extent.
    first_row = torch.from_numpy(region_first_row)
    row_padded = torch.zeros(num_rows, dtype=torch.int32)
    row_padded[first_row] = torch.from_numpy(region_padded).int()
    torch.cumsum(row_padded * dcp_world_size, 0, out=global_cu[1:])
    torch.cumsum(row_padded, 0, out=padded_cu[1:])
    row_start_cu[:num_rows] = global_cu[first_row][torch.from_numpy(region_of_row)]
    row_start_cu[num_rows] = global_cu[num_rows]
    total = int(global_cu[num_rows])
    padded_total = int(padded_cu[num_rows])

    idx = torch.empty(total, dtype=torch.int64, pin_memory=PIN_MEMORY)
    idx_np = idx.numpy()
    region_start = global_cu[first_row].numpy()
    region_padded_start = padded_cu[first_row].numpy()
    for i in range(len(region_first_row)):
        g = int(region_padded[i]) * dcp_world_size
        t = np.arange(g, dtype=np.int64)
        local = round_down(t // dcp_world_size, interleave) + t % interleave
        # Padded positions are outside causal bounds but must remain in-bounds.
        local = np.minimum(local, region_padded[i] - 1)
        idx_np[region_start[i] : region_start[i] + g] = (
            ((t // interleave) % dcp_world_size) * padded_total
            + region_padded_start[i]
            + local
        )

    cu = async_tensor_h2d(cu, device=device)
    return PCPGlobalChunkPlan(
        row_start_cu=cu[0],
        global_cu=cu[1],
        padded_local_cu=cu[2],
        padded_local_total=padded_total,
        total=total,
        deinterleave_idx=async_tensor_h2d(idx, device=device),
    )

compute_kpool_tail_slot_mapping(slot_mapping, block_table, query_start_loc, positions, num_actual_tokens, num_reqs, kpool, out=None)

Map every token to its request's one circular tail block.

Source code in vllm/v1/attention/backends/mla/indexer.py
def compute_kpool_tail_slot_mapping(
    slot_mapping: torch.Tensor,
    block_table: torch.Tensor,
    query_start_loc: torch.Tensor,
    positions: torch.Tensor,
    num_actual_tokens: int,
    num_reqs: int,
    kpool: int,
    out: torch.Tensor | None = None,
) -> torch.Tensor:
    """Map every token to its request's one circular tail block."""
    if out is None:
        out = torch.empty_like(slot_mapping)
    else:
        assert out.shape == slot_mapping.shape
    if slot_mapping.is_cuda and slot_mapping.dim() == 1 and num_reqs > 0:
        block = 256
        num_tokens = slot_mapping.shape[0]
        num_actual_tokens = min(num_actual_tokens, num_tokens)
        grid = (num_reqs + triton.cdiv(num_tokens - num_actual_tokens, block),)
        _kpool_tail_slot_mapping_kernel[grid](
            slot_mapping,
            block_table,
            block_table.stride(0),
            query_start_loc,
            positions,
            out,
            num_reqs,
            num_actual_tokens,
            num_tokens,
            kpool,
            BLOCK=block,
            num_warps=4,
        )
        return out
    # Torch fallback: CPU tensors (the CPU unit tests), non-1D or empty inputs.
    # Production always takes the Triton path above — spec-decode tokens arrive
    # flattened token-major, so slot_mapping is 1D there too.
    out.copy_(slot_mapping)
    if num_actual_tokens == 0:
        return out
    tokens = torch.arange(num_actual_tokens, device=slot_mapping.device)
    req = torch.searchsorted(query_start_loc, tokens, right=True) - 1
    req = req.clamp_(min=0, max=num_reqs - 1)
    own_block = block_table[:num_reqs, 0].index_select(0, req).to(torch.int64)
    pos = positions[:num_actual_tokens].to(torch.int64)
    out[:num_actual_tokens] = own_block * kpool + torch.remainder(pos, kpool)
    return out

dsa_indexer_uses_fp4(vllm_config)

Whether the DeepSeek sparse indexer should use the MXFP4 K cache.

Source code in vllm/v1/attention/backends/mla/indexer.py
def dsa_indexer_uses_fp4(vllm_config: VllmConfig) -> bool:
    """Whether the DeepSeek sparse indexer should use the MXFP4 K cache."""
    kv_dtype = vllm_config.attention_config.resolve_indexer_kv_dtype("fp8")
    if kv_dtype not in DSA_INDEXER_KV_DTYPES:
        raise ValueError(
            f"indexer_kv_dtype={kv_dtype!r} is not supported by the DeepSeek "
            f"sparse indexer (expected one of {DSA_INDEXER_KV_DTYPES})."
        )
    use_fp4 = kv_dtype == "mxfp4"
    if use_fp4 and current_platform.is_rocm():
        from vllm._aiter_ops import rocm_aiter_ops
        from vllm.platforms.rocm import get_cdna_version

        # Only DeepSeek-V4.1 is wired to the ROCm MXFP4 cache; other DSA models
        # would silently keep their FP8 one.
        model_config = vllm_config.model_config
        if model_config is None or model_config.hf_config.model_type != "deepseek_v41":
            raise ValueError(
                "indexer_kv_dtype='mxfp4' on ROCm is only supported for "
                "DeepSeek-V4.1-Flash."
            )
        if get_cdna_version() != 4:
            raise ValueError(
                "indexer_kv_dtype='mxfp4' on ROCm requires CDNA4 (MI350X/MI355X)."
            )
        if not rocm_aiter_ops.is_enabled():
            raise ValueError(
                "indexer_kv_dtype='mxfp4' on ROCm runs on aiter's kernels; enable "
                "aiter with VLLM_ROCM_USE_AITER=1."
            )
        return True
    if use_fp4 and not current_platform.is_device_capability_family(100):
        raise ValueError(
            "indexer_kv_dtype='mxfp4' requires Blackwell datacenter GPUs "
            "(sm_10x, e.g. B200/GB200); sm_120 (consumer Blackwell) and "
            "earlier architectures are not supported."
        )
    return use_fp4