vllm.v1.attention.backends.mla.sparse_utils
¶
Utility functions for sparse MLA backends.
Functions:
-
flat_kv_row_view–Flat [row, head_dim] view of a paged cache and its physical rows per block.
-
prepare_sparse_mla_safe_lengths–Install dummy slots for empty queries and return nonzero kernel lengths.
-
request_row_bounds–Bounds of the runs of adjacent rows that belong to one request: run
-
triton_convert_req_index_to_global_index–out[token_id, indice_id] =
-
triton_filter_and_convert_dcp_index–Filter global per-request indices to this DCP rank's local slots.
_remap_tiling(NUM_TOPK_TOKENS, BLOCK_N, count_valid)
¶
Pick the column tiling for the index remap kernel.
Counting the valid slots per row is the only reason the column tiles have to talk to each other, so when counting give one program the whole row: the count becomes an in-register reduction plus a plain store, needing neither atomics nor a zero-initialized counter. Pad modest non-power-of-two widths (including GLM's 2176 entries) to one tile; larger widths stay tiled and atomic.
Returns:
Source code in vllm/v1/attention/backends/mla/sparse_utils.py
flat_kv_row_view(kv_cache, block_size)
¶
Flat [row, head_dim] view of a paged cache and its physical rows per block.
Token offset is block_idx * block_stride_rows + offset_in_block.
When other layers' pages sit between consecutive blocks of this cache,
block_stride_rows exceeds block_size; those in-between rows are never
indexed (triton_convert_req_index_to_global_index ensures this).
Source code in vllm/v1/attention/backends/mla/sparse_utils.py
prepare_sparse_mla_safe_lengths(physical_indices, valid_counts)
¶
Install dummy slots for empty queries and return nonzero kernel lengths.
Preserve the raw counts so empty outputs and LSE can be neutralized later.
Source code in vllm/v1/attention/backends/mla/sparse_utils.py
request_row_bounds(req_idx)
¶
Bounds of the runs of adjacent rows that belong to one request: run
r is rows [bounds[r], bounds[r + 1]).
Under PCP a rank holds two adjacent chunk rows of a split prefill; the sparse backends give such a run one KV region.
Source code in vllm/v1/attention/backends/mla/sparse_utils.py
triton_convert_req_index_to_global_index(req_id, block_table, token_indices, BLOCK_SIZE=64, BLOCK_STRIDE_ROWS=None, NUM_TOPK_TOKENS=2048, BLOCK_N=128, HAS_PREFILL_WORKSPACE=False, prefill_workspace_request_ids=None, prefill_workspace_starts=None, prefill_workspace_rank_stride=None, dcp_size=1, dcp_rank=0, cp_kv_cache_interleave_size=1, return_valid_counts=False, out=None, valid_counts_out=None)
¶
out[token_id, indice_id] = block_table[req_id[token_id], token_indices[token_id, indice_id] // BLOCK_SIZE] * BLOCK_SIZE + token_indices[token_id, indice_id] % BLOCK_SIZE
Only when token_indices[token_id, indice_id] == -1 do we output -1. For safety, we also output -1 if the derived block_id would be out-of-bounds.
When HAS_PREFILL_WORKSPACE is True, prefill tokens are mapped to workspace offsets instead of global cache slots. prefill_workspace_request_ids and prefill_workspace_starts must be provided.
int32 [num_tokens], -1 for decode else
prefill request index (maps to prefill_workspace_starts)
prefill_workspace_starts: int32 [num_prefills], 0-indexed workspace starts for each prefill request
When return_valid_counts is True, also returns the count of valid (non -1) indices per row, computed during the same kernel pass (no extra overhead).
Source code in vllm/v1/attention/backends/mla/sparse_utils.py
453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 | |
triton_filter_and_convert_dcp_index(req_id, block_table, token_indices, dcp_size, dcp_rank, cp_kv_cache_interleave_size=1, BLOCK_SIZE=64, BLOCK_STRIDE_ROWS=None, NUM_TOPK_TOKENS=2048, BLOCK_N=128, return_valid_counts=False, compact_valid_to_front=True)
¶
Filter global per-request indices to this DCP rank's local slots.
With compact_valid_to_front (default), the conversion kernel scatters
this rank's owned slots to a contiguous prefix [0, valid_count) and
leaves the rest -1. DCP filtering marks non-owned slots -1 and so
creates interior gaps; the trtllm-gen sparse kernel reads the first
valid_count entries of each row, so they must be a contiguous prefix.
Compaction is fused into the kernel (atomic slot allocator) rather than a
separate sort/gather pass. Prefix order is unspecified (only the set matters).
Source code in vllm/v1/attention/backends/mla/sparse_utils.py
591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 | |