Skip to content

vllm.v1.attention.ops.rocm_paged_mxfp4_indexer

Sparse attention indexer on aiter's paged MXFP4 MQA-logits kernel.

gfx950 only. The kernel reads the preshuffled paged indexer K cache in place, so no layer gathers K into a contiguous buffer, prefill included. Any model whose indexer supplies the inputs rocm_mxfp4_sparse_attn_indexer documents can run the dense path. DeepSeek-V4.1's two-level indexer adds to it: the candidate source takes its block maxima from the same walk that writes its logits, and the candidate consumers can walk the candidate pool instead of the whole context: the pool is resolved once per step and shared by all of them.

Classes:

Functions:

RocmPagedMxfp4CacheLayout

Bases: NamedTuple

The byte order of an indexer K page, as the kernel reads it, and the shuffle pattern aiter's K cache op writes it with.

A page is cut into runs of n_per_tile tokens. Inside a run, each token's packed values are split into d_per_tile-byte chunks along the head dimension, and the run stores chunk 0 of every token, then chunk 1, and so on: [chunk, token, byte]. Its e8m0 scales are stored as [scale % scale_lanes, token, scale // scale_lanes].

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
class RocmPagedMxfp4CacheLayout(NamedTuple):
    """The byte order of an indexer K page, as the kernel reads it, and the
    shuffle pattern aiter's K cache op writes it with.

    A page is cut into runs of ``n_per_tile`` tokens. Inside a run, each token's
    packed values are split into ``d_per_tile``-byte chunks along the head
    dimension, and the run stores chunk 0 of every token, then chunk 1, and so
    on: ``[chunk, token, byte]``. Its e8m0 scales are stored as
    ``[scale % scale_lanes, token, scale // scale_lanes]``.
    """

    n_per_tile: int
    d_per_tile: int
    scale_lanes: int

_Layer dataclass

One indexer layer's inputs and outputs for this step.

Methods:

  • decode_rows –

    The first n token rows, one sequence each.

  • rows –

    Q, its scales and the weights of token rows [lo, hi), as seqs

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
@dataclass
class _Layer:
    """One indexer layer's inputs and outputs for this step."""

    metadata: "RocmMxfp4IndexerMetadata"
    kv: torch.Tensor
    q: torch.Tensor
    q_scale: torch.Tensor
    weights: torch.Tensor
    topk_buffer: torch.Tensor
    topk_tokens: int
    compress_ratio: int
    candidates: torch.Tensor | None
    block: int
    full_graph: bool

    @property
    def num_heads(self) -> int:
        return self.q.shape[1]

    @property
    def head_dim(self) -> int:
        return self.q.shape[2] * 2

    def rows(self, lo: int, hi: int, seqs: int = 1) -> tuple[torch.Tensor, ...]:
        """Q, its scales and the weights of token rows [lo, hi), as ``seqs``
        sequences of next_n rows each."""
        return (
            self.q[lo:hi].view(seqs, -1, *self.q.shape[1:]),
            self.q_scale[lo:hi].view(seqs, -1, *self.q_scale.shape[1:]),
            self.weights[lo:hi],
        )

    def decode_rows(self, n: int) -> tuple[torch.Tensor, ...]:
        """The first n token rows, one sequence each."""
        return (
            self.q[:n].unsqueeze(1),
            self.q_scale[:n].unsqueeze(1),
            self.weights[:n],
        )

decode_rows(n)

The first n token rows, one sequence each.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def decode_rows(self, n: int) -> tuple[torch.Tensor, ...]:
    """The first n token rows, one sequence each."""
    return (
        self.q[:n].unsqueeze(1),
        self.q_scale[:n].unsqueeze(1),
        self.weights[:n],
    )

rows(lo, hi, seqs=1)

Q, its scales and the weights of token rows [lo, hi), as seqs sequences of next_n rows each.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def rows(self, lo: int, hi: int, seqs: int = 1) -> tuple[torch.Tensor, ...]:
    """Q, its scales and the weights of token rows [lo, hi), as ``seqs``
    sequences of next_n rows each."""
    return (
        self.q[lo:hi].view(seqs, -1, *self.q.shape[1:]),
        self.q_scale[lo:hi].view(seqs, -1, *self.q_scale.shape[1:]),
        self.weights[lo:hi],
    )

_aiter() cached

The aiter module with the paged MXFP4 MQA-logits launcher, its schedule and cache_format.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
@functools.cache
def _aiter():
    """The aiter module with the paged MXFP4 MQA-logits launcher, its schedule
    and cache_format."""
    from aiter.ops.triton.attention import pa_mqa_logits_mxfp4

    return pa_mqa_logits_mxfp4

_aiter_topk() cached

The aiter per-row top-k.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
@functools.cache
def _aiter_topk() -> Callable[..., None]:
    """The aiter per-row top-k."""
    from aiter.ops.topk import top_k_per_row_decode

    return top_k_per_row_decode

_kv_view(kv_cache, head_dim)

The indexer cache as [pages, entries, 1, bytes]. The page stride is the block-major pool's, not the page's own size.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def _kv_view(kv_cache: torch.Tensor, head_dim: int) -> torch.Tensor:
    """The indexer cache as [pages, entries, 1, bytes]. The page stride is the
    block-major pool's, not the page's own size."""
    num_pages, entries, width = kv_cache.shape
    assert kv_cache.dtype == torch.uint8
    assert width == head_dim // 2 + head_dim // MXFP4_BLOCK_SIZE, width
    return torch.as_strided(
        kv_cache, (num_pages, entries, 1, width), (kv_cache.stride(0), width, width, 1)
    )

_remap_compact_topk(indices, positions, block)

Candidate slots to context positions, in place: slot j sits in pool block j // block, which starts at positions[j // block].

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def _remap_compact_topk(
    indices: torch.Tensor, positions: torch.Tensor, block: int
) -> None:
    """Candidate slots to context positions, in place: slot j sits in pool
    block j // block, which starts at positions[j // block]."""
    rows, k = indices.shape
    if rows == 0:
        return
    _remap_compact_topk_kernel[(rows,)](
        indices,
        indices.stride(0),
        positions,
        positions.stride(0),
        K=k,
        BLOCK=block,
        PADDED_K=triton.next_power_of_2(k),
        num_warps=4,
    )

_topk(logits, lengths, out, k)

Top-k over each row's [0, lengths[row]); -1 pads rows shorter than k.

aiter's kernel is faster on the candidate top-k and from 128 rows on; below that, a wide decode step is faster on vLLM's. Shape alone decides, so a FULL graph replays the same choice at every context length.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def _topk(
    logits: torch.Tensor, lengths: torch.Tensor, out: torch.Tensor, k: int
) -> None:
    """Top-k over each row's [0, lengths[row]); -1 pads rows shorter than k.

    aiter's kernel is faster on the candidate top-k and from 128 rows on; below
    that, a wide decode step is faster on vLLM's. Shape alone decides, so a FULL
    graph replays the same choice at every context length.
    """
    rows, width = logits.shape
    if rows == 0:
        return
    if width >= k and (k >= 2048 or rows >= 128):
        _aiter_topk()(
            logits,
            1,
            lengths.reshape(-1),
            out,
            rows,
            logits.stride(0),
            logits.stride(1),
            k=k,
        )
    else:
        torch.ops._C.top_k_per_row_decode(
            logits, 1, lengths, out, rows, logits.stride(0), logits.stride(1), k
        )

build_rocm_mxfp4_decode_schedule(row_lens, num_heads, head_dim, page_entries, out, logits_width, native=None)

Work descriptors that even out a decode step, or None where the static grid already fills the machine. Depends only on the rows' lengths, the cache geometry and the logits width, so one serves every layer of a group. The width sizes the slices, capped the way the static grid caps them.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def build_rocm_mxfp4_decode_schedule(
    row_lens: torch.Tensor,
    num_heads: int,
    head_dim: int,
    page_entries: int,
    out: torch.Tensor,
    logits_width: int,
    native: "RocmMxfp4NativeDecode | None" = None,
) -> torch.Tensor | None:
    """Work descriptors that even out a decode step, or None where the static
    grid already fills the machine. Depends only on the rows' lengths, the
    cache geometry and the logits width, so one serves every layer of a group.
    The width sizes the slices, capped the way the static grid caps them."""
    if native is None:
        return _aiter().build_schedule(
            row_lens,
            1,
            num_heads,
            head_dim,
            page_entries,
            out=out,
            max_model_len=logits_width,
        )
    return _aiter().build_schedule(
        native.context_lens,
        native.next_n,
        num_heads,
        head_dim,
        page_entries,
        out=out,
        row_ends=row_lens,
        max_model_len=logits_width,
    )

reserve_rocm_mxfp4_indexer_workspace(hidden_states, logits_width, candidate_block_size=0, gather_block_size=0, num_candidate_cols=0)

Profiling run: claim the decode logits workspace and the peak prefill logits, block scores included when the layer writes them, candidate lists when it gathers.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def reserve_rocm_mxfp4_indexer_workspace(
    hidden_states: torch.Tensor,
    logits_width: int,
    candidate_block_size: int = 0,
    gather_block_size: int = 0,
    num_candidate_cols: int = 0,
) -> None:
    """Profiling run: claim the decode logits workspace and the peak prefill
    logits, block scores included when the layer writes them, candidate lists
    when it gathers."""
    rows = _max_decode_logits_rows(hidden_states.shape[0])
    specs = [((rows, logits_width), torch.float32)]
    budget = envs.VLLM_SPARSE_INDEXER_MAX_LOGITS_MB * 1024 * 1024
    if candidate_block_size:
        nblocks = triton.cdiv(logits_width, candidate_block_size)
        specs.append(((rows, nblocks), torch.float32))
        budget += budget // candidate_block_size
    if gather_block_size:
        # aiter allocates the candidate lists: each gathered row, 8 B per
        # candidate block (int32 slot and position). The first consumer builds
        # them and they live until the step ends, across the later layers, so
        # one reservation is held for the rest of the forward.
        held = get_forward_context().additional_kwargs
        if "rocm_mxfp4_candidate_lists" not in held:
            held["rocm_mxfp4_candidate_lists"] = torch.empty(
                hidden_states.shape[0] * (num_candidate_cols // gather_block_size) * 8,
                dtype=torch.uint8,
                device=hidden_states.device,
            )
    current_workspace_manager().get_simultaneous(*specs)
    torch.empty(budget, dtype=torch.uint8, device=hidden_states.device)

rocm_mxfp4_consumer_rows(num_candidate_cols)

Query rows per candidate-consumer launch. Its logits are [rows, pool] fp32 whatever the context, so the logits budget alone sizes it.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def rocm_mxfp4_consumer_rows(num_candidate_cols: int) -> int:
    """Query rows per candidate-consumer launch. Its logits are [rows, pool]
    fp32 whatever the context, so the logits budget alone sizes it."""
    budget = envs.VLLM_SPARSE_INDEXER_MAX_LOGITS_MB * 1024 * 1024
    return max(1, budget // (4 * num_candidate_cols))

rocm_mxfp4_decode_schedule_words(num_heads, head_dim, page_entries, next_n=1)

int32 words of the largest schedule a decode step can build, flattened or not. The slice cap can take it past target_wgs, up to SCHED_SLOT_CAP.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
def rocm_mxfp4_decode_schedule_words(
    num_heads: int, head_dim: int, page_entries: int, next_n: int = 1
) -> int:
    """int32 words of the largest schedule a decode step can build, flattened or
    not. The slice cap can take it past target_wgs, up to SCHED_SLOT_CAP."""
    words = 4 * _aiter().SCHED_SLOT_CAP
    for rows in {1, next_n}:
        config = _aiter().select_config(num_heads, head_dim, rows, page_entries)
        words = max(words, 4 * config["target_wgs"])
    return words

rocm_mxfp4_sparse_attn_indexer(hidden_states, k_cache_prefix, kv_cache, q_values, q_scale, weights, topk_tokens, head_dim, max_model_len, topk_indices_buffer, compress_ratio=1, candidate_blocks=None, candidate_block_size=0, candidate_write=False)

Dense indexer: every layer that scores the whole context, the candidate source included. With candidate_blocks and not candidate_write the scores are masked to the pool first, as the shared path does.

The model supplies q_values [T, H, D // 2] packed e2m1, q_scale holding each head's D // 32 ue8m0 bytes (one int32 per head at D = 128), fp32 weights with the softmax and head scales folded in, and kv_cache pages in rocm_paged_mxfp4_cache_layout order for H heads.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
@eager_break_during_capture
def rocm_mxfp4_sparse_attn_indexer(
    hidden_states: torch.Tensor,
    k_cache_prefix: str,
    kv_cache: torch.Tensor,
    q_values: torch.Tensor,
    q_scale: torch.Tensor,
    weights: torch.Tensor,
    topk_tokens: int,
    head_dim: int,
    max_model_len: int,
    topk_indices_buffer: torch.Tensor,
    compress_ratio: int = 1,
    candidate_blocks: torch.Tensor | None = None,
    candidate_block_size: int = 0,
    candidate_write: bool = False,
) -> torch.Tensor:
    """Dense indexer: every layer that scores the whole context, the
    candidate source included. With ``candidate_blocks`` and not
    ``candidate_write`` the scores are masked to the pool first, as the
    shared path does.

    The model supplies ``q_values`` [T, H, D // 2] packed e2m1, ``q_scale``
    holding each head's D // 32 ue8m0 bytes (one int32 per head at D = 128),
    fp32 ``weights`` with the softmax and head scales folded in, and
    ``kv_cache`` pages in `rocm_paged_mxfp4_cache_layout` order for H heads.
    """
    if not isinstance(get_forward_context().attn_metadata, dict):
        reserve_rocm_mxfp4_indexer_workspace(
            hidden_states,
            max_model_len,
            candidate_block_size if candidate_write else 0,
        )
        return topk_indices_buffer
    layer = _layer(
        k_cache_prefix,
        kv_cache,
        q_values,
        q_scale,
        weights,
        topk_indices_buffer,
        topk_tokens,
        head_dim,
        compress_ratio,
        candidate_blocks,
        candidate_block_size,
    )
    metadata = layer.metadata
    topk_indices_buffer[: hidden_states.shape[0]] = -1
    if metadata.prefill is not None:
        for chunk, plan in zip(metadata.prefill.chunks, metadata.prefill_plans):
            _dense_prefill(layer, chunk, plan, candidate_write)
    if metadata.decode is not None:
        _dense_decode(layer, max_model_len, candidate_write)
    return topk_indices_buffer

rocm_mxfp4_sparse_mqa_indexer(hidden_states, k_cache_prefix, kv_cache, q_values, q_scale, weights, topk_tokens, head_dim, max_model_len, topk_indices_buffer, compress_ratio, candidate_blocks, candidate_block_size, num_candidate_cols)

Candidate consumer: score only the source's pool where the builder's length gate says it pays, else the dense walk masked to the pool.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
@eager_break_during_capture
def rocm_mxfp4_sparse_mqa_indexer(
    hidden_states: torch.Tensor,
    k_cache_prefix: str,
    kv_cache: torch.Tensor,
    q_values: torch.Tensor,
    q_scale: torch.Tensor,
    weights: torch.Tensor,
    topk_tokens: int,
    head_dim: int,
    max_model_len: int,
    topk_indices_buffer: torch.Tensor,
    compress_ratio: int,
    candidate_blocks: torch.Tensor,
    candidate_block_size: int,
    num_candidate_cols: int,
) -> torch.Tensor:
    """Candidate consumer: score only the source's pool where the builder's
    length gate says it pays, else the dense walk masked to the pool."""
    if not isinstance(get_forward_context().attn_metadata, dict):
        reserve_rocm_mxfp4_indexer_workspace(
            hidden_states,
            max(max_model_len, num_candidate_cols),
            gather_block_size=candidate_block_size,
            num_candidate_cols=num_candidate_cols,
        )
        return topk_indices_buffer
    layer = _layer(
        k_cache_prefix,
        kv_cache,
        q_values,
        q_scale,
        weights,
        topk_indices_buffer,
        topk_tokens,
        head_dim,
        compress_ratio,
        candidate_blocks,
        candidate_block_size,
    )
    metadata = layer.metadata
    topk_indices_buffer[: hidden_states.shape[0]] = -1
    if metadata.prefill is not None:
        for chunk, plan in zip(metadata.prefill.chunks, metadata.prefill_plans):
            if not plan.use_gather:
                _dense_prefill(layer, chunk, plan, candidate_write=False)
        for launch in metadata.gather_launches:
            _gather_prefill(layer, launch, num_candidate_cols)
    if metadata.decode is not None:
        # A FULL graph cannot follow the per-step gate; the gather is the side
        # that stays flat in the context length.
        if metadata.decode_use_gather or layer.full_graph:
            _gather_decode(layer, num_candidate_cols)
        else:
            _dense_decode(layer, max_model_len, candidate_write=False)
    return topk_indices_buffer

rocm_paged_mxfp4_cache_layout(num_heads, head_dim, page_entries) cached

The K page layout from aiter's cache_format, the only place vLLM reads it, so a change there cannot reach the writer unchecked.

Raises:

  • ValueError –

    aiter rejects the page size, or describes an order this layout cannot express.

Source code in vllm/v1/attention/ops/rocm_paged_mxfp4_indexer.py
@functools.cache
def rocm_paged_mxfp4_cache_layout(
    num_heads: int, head_dim: int, page_entries: int
) -> RocmPagedMxfp4CacheLayout:
    """The K page layout from aiter's ``cache_format``, the only place vLLM
    reads it, so a change there cannot reach the writer unchecked.

    Raises:
        ValueError: aiter rejects the page size, or describes an order this
            layout cannot express.

    """
    fmt = _aiter().cache_format(num_heads, head_dim, page_entries)
    try:
        n_per_tile = int(fmt["n_per_tile"])
        d_per_tile = int(fmt["d_per_tile"])
        scale_order = fmt.get("scale_mode", _SCALE_ORDER)
        scale_lanes = int(fmt.get("scale_lanes", _WAVE_SIZE // n_per_tile))
    except (KeyError, TypeError, ValueError, ZeroDivisionError) as e:
        raise ValueError(f"aiter's cache_format() returned {fmt!r}") from e
    if (
        scale_order != _SCALE_ORDER
        or min(n_per_tile, d_per_tile, scale_lanes) <= 0
        or page_entries % n_per_tile
        or (head_dim // 2) % d_per_tile
        or (head_dim // MXFP4_BLOCK_SIZE) % scale_lanes
    ):
        raise ValueError(
            f"aiter's cache_format() returned {fmt!r} for {num_heads} heads of "
            f"{head_dim} in {page_entries}-entry pages, a K page order "
            "RocmPagedMxfp4CacheLayout cannot express"
        )
    return RocmPagedMxfp4CacheLayout(n_per_tile, d_per_tile, scale_lanes)