Skip to content

vllm.v1.attention.ops.rocm_aiter_mla_sparse

Functions:

_apply_candidate_mask_strided(logits, row_ks, row_ke, candidate_blocks, block_size, row_repeat=1)

ROCm decode variant of apply_candidate_mask.

Same masking semantics over [0, end), but the grid is sized by a fixed program count rather than by the logits width. Only worth using where the width is the max_model_len workspace and the live context is far shorter, i.e. the paged decode path below; the prefill chunks pass chunk-sized logits and stay on the shared kernel.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _apply_candidate_mask_strided(
    logits: torch.Tensor,
    row_ks: torch.Tensor | None,
    row_ke: torch.Tensor,
    candidate_blocks: torch.Tensor,
    block_size: int,
    row_repeat: int = 1,
) -> None:
    """ROCm decode variant of ``apply_candidate_mask``.

    Same masking semantics over ``[0, end)``, but the grid is sized by a fixed
    program count rather than by the logits width. Only worth using where the
    width is the ``max_model_len`` workspace and the live context is far
    shorter, i.e. the paged decode path below; the prefill chunks pass
    chunk-sized logits and stay on the shared kernel.
    """
    from vllm.model_executor.kernels.attention.dsa.candidate_blocks import (
        _candidate_flags_kernel,
    )

    rows, width = logits.shape
    if not rows or not width:
        return
    nblocks = triton.cdiv(width, block_size)
    flags = torch.empty((rows, nblocks + 1), device=logits.device, dtype=torch.uint8)
    start_stride = row_ks.stride(0) if row_ks is not None else 0
    _candidate_flags_kernel[(rows,)](
        candidate_blocks,
        row_ks,
        flags,
        *candidate_blocks.stride(),
        start_stride,
        width,
        nblocks,
        block_size,
        candidate_blocks.shape[1],
        row_ks is not None,
        row_repeat,
    )
    # Derived from width, which is a tensor shape, so the grid stays static and
    # a FULL cudagraph capture remains valid across replays; only the loop trip
    # count inside the kernel is data-dependent. The min keeps narrow widths
    # from launching programs that would only fall through.
    grid_cols = min(_MASK_GRID_COLS, triton.cdiv(width, _MASK_TILE))
    _mask_candidates_strided_kernel[(rows, grid_cols)](
        logits,
        row_ks,
        row_ke,
        flags,
        *logits.stride(),
        start_stride,
        row_ke.stride(0),
        width,
        nblocks,
        block_size,
        row_ks is not None,
        row_repeat,
        _MASK_TILE,
    )

_decode_num_splits(num_queries, heads_blocks, avg_main_len=0.0, avg_extra_len=0.0, block_k=32)

Pick a flash-decode split count to keep the GPU busy across batch sizes.

Decode launches only num_queries * heads_blocks workgroups otherwise, which severely under-fills the device for the low-concurrency regime that dominates latency. Splitting the KV sequence adds parallelism.

We model the relative partial-kernel latency for a given split count s as waves * (1/s + mu) where waves = ceil(base * s / CU) and mu is a small per-wave overhead penalty:

  • waves / s captures the partial compute: each wave walks roughly total_tokens / s tokens and there are waves of them, so dividing by s makes more splits cheaper until they spill into extra waves.
  • mu * waves charges per-wave launch/tail overhead so we do not over-split into many mostly-idle waves (e.g. batch 224 on 256 CUs is best left at 1 split rather than 8 splits across 7 waves).

The minimiser naturally prefers split counts that pack the device into full waves (base * s near a multiple of CU) and falls back to 1 split once the batch already fills the device. Ties favour the smaller split count (less reduce work).

Finally we "snap down" the chosen split count to the smallest value that yields the same wave count and the same per-workgroup BLOCK_K iteration count. Because latency tracks iteration count (not raw token count), extra splits that do not lower the iteration count add only reduce/HBM overhead for no parallelism gain (e.g. batch 24: s8 and s10 both walk 4 extra iters in one wave, so s8 is strictly better). Snapping needs the average segment lengths, which the caller derives sync-free from the ragged index sizes.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _decode_num_splits(
    num_queries: int,
    heads_blocks: int,
    avg_main_len: float = 0.0,
    avg_extra_len: float = 0.0,
    block_k: int = 32,
) -> int:
    """Pick a flash-decode split count to keep the GPU busy across batch sizes.

    Decode launches only ``num_queries * heads_blocks`` workgroups otherwise,
    which severely under-fills the device for the low-concurrency regime that
    dominates latency. Splitting the KV sequence adds parallelism.

    We model the relative partial-kernel latency for a given split count ``s``
    as ``waves * (1/s + mu)`` where ``waves = ceil(base * s / CU)`` and ``mu``
    is a small per-wave overhead penalty:

      - ``waves / s`` captures the partial compute: each wave walks roughly
        ``total_tokens / s`` tokens and there are ``waves`` of them, so dividing
        by ``s`` makes more splits cheaper *until* they spill into extra waves.
      - ``mu * waves`` charges per-wave launch/tail overhead so we do not
        over-split into many mostly-idle waves (e.g. batch 224 on 256 CUs is
        best left at 1 split rather than 8 splits across 7 waves).

    The minimiser naturally prefers split counts that pack the device into full
    waves (``base * s`` near a multiple of ``CU``) and falls back to 1 split
    once the batch already fills the device. Ties favour the smaller split
    count (less reduce work).

    Finally we "snap down" the chosen split count to the smallest value that
    yields the same wave count *and* the same per-workgroup BLOCK_K iteration
    count. Because latency tracks iteration count (not raw token count), extra
    splits that do not lower the iteration count add only reduce/HBM overhead
    for no parallelism gain (e.g. batch 24: s8 and s10 both walk 4 extra iters
    in one wave, so s8 is strictly better). Snapping needs the average segment
    lengths, which the caller derives sync-free from the ragged index sizes.
    """
    base = max(1, num_queries * heads_blocks)
    # Target ~1 workgroup per CU: enough to fill the device while keeping the
    # reduce cost (which grows with split count) small. Tuned on gfx950.
    cu = max(1, _decode_cu_count())
    # Per-wave overhead penalty: higher values discourage split counts that
    # spill into extra GPU waves. Tuned on gfx950.
    mu = 0.04
    best_splits = 1
    best_cost = None
    # Search up to 16 splits; beyond that the reduce/HBM overhead dominates.
    for splits in range(1, 17):
        waves = (base * splits + cu - 1) // cu
        cost = waves * (1.0 / splits + mu)
        if best_cost is None or cost < best_cost - 1e-9:
            best_splits = splits
            best_cost = cost

    if best_splits > 1 and (avg_main_len > 0 or avg_extra_len > 0):
        target_waves = (base * best_splits + cu - 1) // cu
        target_iters = _decode_partial_iters(
            avg_main_len, avg_extra_len, best_splits, block_k
        )
        for splits in range(1, best_splits):
            waves = (base * splits + cu - 1) // cu
            iters = _decode_partial_iters(avg_main_len, avg_extra_len, splits, block_k)
            if waves == target_waves and iters == target_iters:
                best_splits = splits
                break
    return best_splits

_decode_partial_iters(avg_main_len, avg_extra_len, splits, block_k)

BLOCK_K iterations one partial workgroup walks for splits splits.

Each split processes ceil(seg_len / splits) tokens of a segment, walked BLOCK_K at a time, and the main/extra segments are handled separately.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _decode_partial_iters(
    avg_main_len: float, avg_extra_len: float, splits: int, block_k: int
) -> int:
    """BLOCK_K iterations one partial workgroup walks for ``splits`` splits.

    Each split processes ``ceil(seg_len / splits)`` tokens of a segment, walked
    ``BLOCK_K`` at a time, and the main/extra segments are handled separately.
    """
    main_iters = (
        math.ceil(math.ceil(avg_main_len / splits) / block_k) if avg_main_len > 0 else 0
    )
    extra_iters = (
        math.ceil(math.ceil(avg_extra_len / splits) / block_k)
        if avg_extra_len > 0
        else 0
    )
    return main_iters + extra_iters

_fused_inverse_rope_gptj(o, positions, cos_sin_cache, rope_head_dim, out=None)

bf16 inverse GPT-J RoPE via a single fused Triton kernel.

out may alias o: the rotation is a per-row bijection whose kernel reads both lanes of a pair before storing either.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _fused_inverse_rope_gptj(
    o: torch.Tensor,
    positions: torch.Tensor,
    cos_sin_cache: torch.Tensor,
    rope_head_dim: int,
    out: torch.Tensor | None = None,
) -> torch.Tensor:
    """bf16 inverse GPT-J RoPE via a single fused Triton kernel.

    ``out`` may alias ``o``: the rotation is a per-row bijection whose kernel
    reads both lanes of a pair before storing either.
    """
    assert o.dim() == 3 and o.stride(-1) == 1, (
        "_fused_inverse_rope_gptj expects a [T, H, D] input with a contiguous last dim"
    )
    assert rope_head_dim > 0 and rope_head_dim % 2 == 0, (
        f"_fused_inverse_rope_gptj expects an even rope_head_dim, got {rope_head_dim}"
    )
    assert cos_sin_cache.shape[-1] == rope_head_dim, (
        "_fused_inverse_rope_gptj expects cos_sin_cache laid out as "
        f"[P, {rope_head_dim}] = cos | sin, got {tuple(cos_sin_cache.shape)}"
    )
    num_tokens, num_heads, head_dim = o.shape
    if out is None:
        out = torch.empty(
            (num_tokens, num_heads, head_dim), dtype=torch.bfloat16, device=o.device
        )
    else:
        assert out.dtype == torch.bfloat16, (
            f"inverse RoPE writes bf16, got an output buffer of {out.dtype}"
        )
    if num_tokens == 0:
        return out
    _inverse_rope_gptj_kernel[(num_tokens, num_heads)](
        o,
        out,
        positions,
        cos_sin_cache,
        o.stride(0),
        o.stride(1),
        out.stride(0),
        out.stride(1),
        cos_sin_cache.stride(0),
        NOPE=head_dim - rope_head_dim,
        HALF=rope_head_dim // 2,
        BLOCK_NOPE=triton.next_power_of_2(head_dim - rope_head_dim),
        BLOCK_HALF=triton.next_power_of_2(rope_head_dim // 2),
    )
    return out

_get_cached_wo_a_bf16(wo_a, n_local_groups, o_lora_rank, hidden_dim)

Dequantize wo_a to bf16 once and cache it on the module.

wo_a weights are static, so the fp8 -> fp32 -> (* block scale) -> bf16 dequant only needs to run once. Recomputing it every decode step shows up in the profile as the largest copy/mul kernels (direct_copy float ~55us and MulFunctor float ~31us per two layers). SGLang / ATOM keep wo_a in bf16 and feed a plain bf16 GEMM; this mirrors that.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _get_cached_wo_a_bf16(
    wo_a: torch.nn.Module,
    n_local_groups: int,
    o_lora_rank: int,
    hidden_dim: int,
) -> torch.Tensor:
    """Dequantize wo_a to bf16 once and cache it on the module.

    wo_a weights are static, so the fp8 -> fp32 -> (* block scale) -> bf16
    dequant only needs to run once. Recomputing it every decode step shows up
    in the profile as the largest copy/mul kernels (``direct_copy float`` ~55us
    and ``MulFunctor float`` ~31us per two layers). SGLang / ATOM keep wo_a in
    bf16 and feed a plain bf16 GEMM; this mirrors that.
    """
    cached = getattr(wo_a, "_dsv4_wo_a_bf16", None)
    if cached is not None:
        return cached
    from vllm.model_executor.layers.quantization.utils.fp8_utils import (
        get_fp8_block_weight_scale,
    )

    wo_a_scale_param = get_fp8_block_weight_scale(wo_a)
    if wo_a_scale_param is None:
        # ModelOpt MXFP8 stores the multiplicative E8M0 scale without the
        # historical ``_inv`` suffix.
        wo_a_scale_param = getattr(wo_a, "weight_scale", None)
    # Emulated MXFP8 kernels can replace the original one-byte weight with an
    # already-dequantized BF16 tensor while retaining the scale attribute for
    # metadata. Applying that retained scale again would double-dequantize the
    # weight. Block scaling is only valid while the one-byte FP8 storage remains.
    if wo_a_scale_param is not None and wo_a.weight.element_size() == 1:
        wo_a_weight = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to(
            torch.float32
        )
        wo_a_scale = _expand_2d_block_scales(
            wo_a_scale_param.view(n_local_groups, -1, wo_a_scale_param.shape[-1]),
            o_lora_rank,
            hidden_dim,
        )
        cached = (wo_a_weight * wo_a_scale).to(torch.bfloat16)
    else:
        cached = wo_a.weight.view(n_local_groups, o_lora_rank, hidden_dim).to(
            torch.bfloat16
        )
    wo_a._dsv4_wo_a_bf16 = cached
    return cached

_indexer_k_is_c4a_block_flat(compress_ratio)

V4.0 C4A is block-flat (NORMAL). Ratio 1 and 2 are 16×16 SHUFFLE.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _indexer_k_is_c4a_block_flat(compress_ratio: int) -> bool:
    """V4.0 C4A is block-flat (NORMAL). Ratio 1 and 2 are 16×16 SHUFFLE."""
    return compress_ratio == 4

_inverse_rope_gptj_kernel(o_ptr, out_ptr, pos_ptr, cos_sin_ptr, s_t, s_h, os_t, os_h, cs_stride, NOPE, HALF, BLOCK_NOPE, BLOCK_HALF)

Fused inverse GPT-J RoPE on the trailing rope_dim of each (token, head).

Mirrors DeepseekV4ScalingRotaryEmbedding.forward_native(inverse=True) for the GPT-J (non-neox) layout, writing bf16 directly. Replaces the clone + index_select + repeat_interleave + neg + stack + cat + cast chain (~10 small kernels) with a single launch.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
@triton.jit
def _inverse_rope_gptj_kernel(
    o_ptr,  # [T, H, D] input
    out_ptr,  # [T, H, D] bf16 output
    pos_ptr,  # [T] positions
    cos_sin_ptr,  # [P, rope_dim] fp32 (cos[:half] | sin[half:])
    s_t,
    s_h,  # input row strides (last dim contiguous)
    os_t,
    os_h,  # output row strides
    cs_stride,  # cos_sin_cache row stride
    NOPE: tl.constexpr,  # non-rope head dims (passed through)
    HALF: tl.constexpr,  # rope_dim // 2
    BLOCK_NOPE: tl.constexpr,
    BLOCK_HALF: tl.constexpr,
):
    """Fused inverse GPT-J RoPE on the trailing rope_dim of each (token, head).

    Mirrors ``DeepseekV4ScalingRotaryEmbedding.forward_native(inverse=True)``
    for the GPT-J (non-neox) layout, writing bf16 directly. Replaces the
    clone + index_select + repeat_interleave + neg + stack + cat + cast chain
    (~10 small kernels) with a single launch.
    """
    t = tl.program_id(0)
    h = tl.program_id(1)
    in_base = t * s_t + h * s_h
    out_base = t * os_t + h * os_h

    # NoPE lanes pass through unchanged (only cast to bf16).
    n = tl.arange(0, BLOCK_NOPE)
    nmask = n < NOPE
    vals = tl.load(o_ptr + in_base + n, mask=nmask)
    tl.store(out_ptr + out_base + n, vals.to(tl.bfloat16), mask=nmask)

    # RoPE lanes: out_even = a*cos + b*sin, out_odd = b*cos - a*sin
    # (a = even lane, b = odd lane; sin negated for the inverse rotation).
    pos = tl.load(pos_ptr + t).to(tl.int64)
    k = tl.arange(0, BLOCK_HALF)
    kmask = k < HALF
    a = tl.load(o_ptr + in_base + NOPE + 2 * k, mask=kmask).to(tl.float32)
    b = tl.load(o_ptr + in_base + NOPE + 2 * k + 1, mask=kmask).to(tl.float32)
    cos = tl.load(cos_sin_ptr + pos * cs_stride + k, mask=kmask)
    sin = tl.load(cos_sin_ptr + pos * cs_stride + HALF + k, mask=kmask)
    out_even = a * cos + b * sin
    out_odd = b * cos - a * sin
    tl.store(out_ptr + out_base + NOPE + 2 * k, out_even.to(tl.bfloat16), mask=kmask)
    tl.store(out_ptr + out_base + NOPE + 2 * k + 1, out_odd.to(tl.bfloat16), mask=kmask)

_max_decode_logits_rows(num_batched_tokens)

Upper bound on decode rows the paged-MQA logits buffer can ever hold.

rocm_fp8_paged_mqa_logits sizes its workspace as (batch_size * next_n, max_model_len). batch_size is bounded by max_num_seqs and next_n by 1 + num_speculative_tokens, which is far tighter than max_num_batched_tokens -- 192 vs 16384 for a typical 32-seq DSpark-5 deployment. The loose bound is harmless at short contexts but scales with max_model_len, so at the model's full context it asks for tens of TiB and the engine cannot start. Take whichever valid bound is smaller; the workspace is locked after profiling, so it must not be under- estimated.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _max_decode_logits_rows(num_batched_tokens: int) -> int:
    """Upper bound on decode rows the paged-MQA logits buffer can ever hold.

    ``rocm_fp8_paged_mqa_logits`` sizes its workspace as
    ``(batch_size * next_n, max_model_len)``. ``batch_size`` is bounded by
    ``max_num_seqs`` and ``next_n`` by ``1 + num_speculative_tokens``, which is
    far tighter than ``max_num_batched_tokens`` -- 192 vs 16384 for a typical
    32-seq DSpark-5 deployment. The loose bound is harmless at short contexts
    but scales with ``max_model_len``, so at the model's full context it asks
    for tens of TiB and the engine cannot start. Take whichever valid bound is
    smaller; the workspace is locked after profiling, so it must not be under-
    estimated.
    """
    try:
        vllm_config = get_current_vllm_config()
    except Exception:
        return num_batched_tokens
    scheduler_config = getattr(vllm_config, "scheduler_config", None)
    max_num_seqs = getattr(scheduler_config, "max_num_seqs", None)
    if not max_num_seqs:
        return num_batched_tokens
    speculative_config = getattr(vllm_config, "speculative_config", None)
    num_spec = getattr(speculative_config, "num_speculative_tokens", 0) or 0
    return min(num_batched_tokens, max_num_seqs * (1 + num_spec))

_mxfp8_quantize_rows(x, ROWS, COLS)

MXFP8-quantize x [ROWS, COLS] in registers, one scale per 32 lanes.

Returns the rescaled fp32 values (to be cast to e4m3 on store) and the [ROWS, COLS // 32] biased E8M0 exponents.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
@triton.jit
def _mxfp8_quantize_rows(x, ROWS: tl.constexpr, COLS: tl.constexpr):
    """MXFP8-quantize ``x`` [ROWS, COLS] in registers, one scale per 32 lanes.

    Returns the rescaled fp32 values (to be cast to e4m3 on store) and the
    [ROWS, COLS // 32] biased E8M0 exponents.
    """
    blocks = tl.reshape(x, (ROWS, COLS // 32, 32))
    bits = _mxfp8_scale_bits(tl.max(tl.abs(blocks), axis=2))
    # Multiply by the reciprocal: a divisor of 2**-127 would be subnormal and
    # flush to zero, turning an all-zero block into NaN.
    q = blocks * tl.exp2(127.0 - bits)[:, :, None]
    return tl.reshape(q, (ROWS, COLS)), bits

_mxfp8_scale_bits(amax)

Biased E8M0 exponent that puts amax at the top of the e4m3 range.

Same rounding as mxfp8_e4m3_quantize, so the output is bit-identical to quantizing the tensor there.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
@triton.jit
def _mxfp8_scale_bits(amax):
    """Biased E8M0 exponent that puts ``amax`` at the top of the e4m3 range.

    Same rounding as ``mxfp8_e4m3_quantize``, so the output is bit-identical to
    quantizing the tensor there.
    """
    amax = tl.maximum(amax, 1.1754943508222875e-38)
    bits = tl.ceil(tl.log2(amax / 448.0)) + 127.0
    return tl.minimum(tl.maximum(bits, 0.0), 254.0)

_mxfp8_wo_a_bmm_config(num_tokens, n_groups)

(BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages) for gfx950.

Tuned under HIP graphs with a cold weight at G = 4 and 2, over every decode shape of conc 1-128 x 0-5 spec tokens plus prefill chunks up to 8K tokens. The best tile tracks the total work T * G, so the tiers are keyed on it.

This will be replaced after new GEMM kernel from AITER with proper 32x32 scale shape GEMM fp8 enabled.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _mxfp8_wo_a_bmm_config(num_tokens: int, n_groups: int) -> tuple[int, ...]:
    """(BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages) for gfx950.

    Tuned under HIP graphs with a cold weight at G = 4 and 2, over every
    decode shape of conc 1-128 x 0-5 spec tokens plus prefill chunks up to
    8K tokens. The best tile tracks the total work T * G, so the tiers are
    keyed on it.

    This will be replaced after new GEMM kernel from AITER with proper 32x32 scale
    shape GEMM fp8 enabled.
    """
    work = num_tokens * n_groups
    if work <= 64:
        return 16, 16, 1024, 2, 3
    if work <= 128:
        return 32, 16, 1024, 2, 3
    if work <= 256:
        return 32, 32, 512, 2, 3
    if work <= 512:
        return 64, 32, 512, 2, 3
    if work <= 1024:
        return 64, 64, 512, 4, 2
    if work <= 2048:
        return 64, 64, 256, 4, 2
    if work <= 3072:
        return 64, 64, 256, 2, 1
    if work <= 4096:
        return 128, 128, 256, 8, 2
    return 128, 128, 128, 4, 2

_rocm_sparse_attn_decode_ragged_bf16_triton(q, kv, indices, indptr, scale, attn_sink, nope_head_dim, rope_head_dim, num_splits, out=None)

Split-K decode over an bf16 ragged KV cache.

Partitions each query's selected tokens across different workgroups and combines the partials through reduction.

Parameters:

  • q

    (Tensor) –

    Queries laid out as [sq, h, d].

  • kv

    (Tensor) –

    Unquantized KV rows laid out as [skv, d].

  • indices

    (Tensor) –

    Flattened per-query KV slots.

  • indptr

    (Tensor) –

    Segment offsets into indices, [sq + 1].

  • scale

    (float) –

    Softmax scale.

  • attn_sink

    (Tensor | None) –

    Optional per-head sink logits.

  • nope_head_dim

    (int) –

    NoPE width of d.

  • rope_head_dim

    (int) –

    RoPE width of d.

  • num_splits

    (int) –

    Number of KV splits per query.

  • out

    (Tensor | None, default: None ) –

    Optional destination with d trailing elements.

Returns:

  • Tensor –

    The attention output, out when provided.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def _rocm_sparse_attn_decode_ragged_bf16_triton(
    q: torch.Tensor,
    kv: torch.Tensor,
    indices: torch.Tensor,
    indptr: torch.Tensor,
    scale: float,
    attn_sink: torch.Tensor | None,
    nope_head_dim: int,
    rope_head_dim: int,
    num_splits: int,
    out: torch.Tensor | None = None,
) -> torch.Tensor:
    """Split-K decode over an bf16 ragged KV cache.

    Partitions each query's selected tokens across different workgroups and combines
    the partials through reduction.

    Args:
        q: Queries laid out as ``[sq, h, d]``.
        kv: Unquantized KV rows laid out as ``[skv, d]``.
        indices: Flattened per-query KV slots.
        indptr: Segment offsets into ``indices``, ``[sq + 1]``.
        scale: Softmax scale.
        attn_sink: Optional per-head sink logits.
        nope_head_dim: NoPE width of ``d``.
        rope_head_dim: RoPE width of ``d``.
        num_splits: Number of KV splits per query.
        out: Optional destination with ``d`` trailing elements.

    Returns:
        The attention output, ``out`` when provided.

    """
    assert q.ndim == 3, f"expected q=[sq,h,d], got {q.shape}"
    assert kv.ndim == 2, f"expected kv=[skv,d], got {kv.shape}"
    assert indices.ndim == 1, f"expected indices=[nnz], got {indices.shape}"
    assert indptr.ndim == 1, f"expected indptr=[sq+1], got {indptr.shape}"
    assert not q.is_cpu and not kv.is_cpu and not indices.is_cpu and not indptr.is_cpu

    indices = _as_int32_contiguous_1d(indices)
    indptr = _as_int32_contiguous_1d(indptr)
    has_attn_sink = attn_sink is not None
    if attn_sink is None:
        attn_sink = torch.empty(1, device=q.device, dtype=torch.float32)
    else:
        attn_sink = attn_sink.contiguous()

    num_queries, num_heads, head_dim = q.shape
    assert indptr.numel() == num_queries + 1, (
        f"expected indptr shape [{num_queries + 1}], got {indptr.shape}"
    )
    _validate_sparse_dims(
        head_dim,
        nope_head_dim,
        rope_head_dim,
        "_rocm_sparse_attn_decode_ragged_bf16_triton",
    )
    if out is None:
        out = torch.empty_like(q)
    assert out.shape[-1] == head_dim, (
        f"expected out trailing dim {head_dim}, got {out.shape[-1]}"
    )

    block_h = _SPARSE_DECODE_BF16_BLOCK_H
    block_d = triton.next_power_of_2(head_dim)
    block_k = _SPARSE_DECODE_BF16_BLOCK_K
    heads_blocks = triton.cdiv(num_heads, block_h)
    comb_dim = nope_head_dim + rope_head_dim

    part_m = torch.empty(
        (num_queries, num_splits, num_heads), dtype=torch.float32, device=q.device
    )
    part_l = torch.empty_like(part_m)
    part_acc = torch.empty(
        (num_queries, num_splits, num_heads, comb_dim),
        dtype=torch.float32,
        device=q.device,
    )

    _sparse_attn_decode_ragged_bf16_partial_kernel[
        (num_queries, num_splits, heads_blocks)
    ](
        q,
        kv,
        indices,
        indptr,
        part_m,
        part_l,
        part_acc,
        q.stride(0),
        q.stride(1),
        q.stride(2),
        kv.stride(0),
        kv.stride(1),
        part_m.stride(0),
        part_m.stride(1),
        part_acc.stride(0),
        part_acc.stride(1),
        part_acc.stride(2),
        num_heads,
        head_dim,
        kv.shape[0],
        float(scale),
        BLOCK_H=block_h,
        BLOCK_D=block_d,
        BLOCK_K=block_k,
        NUM_SPLITS=num_splits,
        num_warps=4,
    )

    _sparse_attn_decode_reduce_kernel[(num_queries, num_heads)](
        part_m,
        part_l,
        part_acc,
        attn_sink,
        out,
        None,
        None,
        None,
        out.stride(0),
        out.stride(1),
        0,
        head_dim // 32,
        part_m.stride(0),
        part_m.stride(1),
        part_acc.stride(0),
        part_acc.stride(1),
        part_acc.stride(2),
        0,
        num_heads,
        HAS_ATTN_SINK=has_attn_sink,
        ADAPTIVE_SPLITS=False,
        COMB_DIM=comb_dim,
        BLOCK_H=1,
        NUM_SPLITS=num_splits,
        SPLITS_PAD=triton.next_power_of_2(num_splits),
        FUSE_INV_ROPE=False,
        NOPE=nope_head_dim,
        HALF=rope_head_dim // 2,
        QUANT_OUT=False,
        num_warps=4,
    )
    return out

_rocm_sparse_attn_decode_ragged_triton(q, main_cache, main_indices, main_indptr, scale, attn_sink, nope_head_dim, rope_head_dim, extra_cache=None, extra_indices=None, extra_indptr=None, out=None, extra_cache_nan_free=False, adaptive_splits=False, inv_rope_positions=None, inv_rope_cos_sin_cache=None, out_mxfp8=None)

Split-K sparse decode; returns the attention output.

With out_mxfp8 = (data, scale) the reduce writes MXFP8 instead of bf16: data is [b, h * d] e4m3 and scale [b, h * d // 32] E8M0, and data viewed as [b, h, d] is returned.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
3956
3957
3958
3959
3960
3961
3962
3963
3964
3965
3966
3967
3968
3969
3970
3971
3972
3973
3974
3975
3976
3977
3978
3979
3980
3981
3982
3983
3984
3985
3986
3987
3988
3989
3990
3991
3992
3993
3994
3995
3996
3997
3998
3999
4000
4001
4002
4003
4004
4005
4006
4007
4008
4009
4010
4011
4012
4013
4014
4015
4016
4017
4018
4019
4020
4021
4022
4023
4024
4025
4026
4027
4028
4029
4030
4031
4032
4033
4034
4035
4036
4037
4038
4039
4040
4041
4042
4043
4044
4045
4046
4047
4048
4049
4050
4051
4052
4053
4054
4055
4056
4057
4058
4059
4060
4061
4062
4063
4064
4065
4066
4067
4068
4069
4070
4071
4072
4073
4074
4075
4076
4077
4078
4079
4080
4081
4082
4083
4084
4085
4086
4087
4088
4089
4090
4091
4092
4093
4094
4095
4096
4097
4098
4099
4100
4101
4102
4103
4104
4105
4106
4107
4108
4109
4110
4111
4112
4113
4114
4115
4116
4117
4118
4119
4120
4121
4122
4123
4124
4125
4126
4127
4128
4129
4130
4131
4132
4133
4134
4135
4136
4137
4138
4139
4140
4141
4142
4143
4144
4145
4146
4147
4148
4149
4150
4151
4152
4153
4154
4155
4156
4157
4158
4159
4160
4161
4162
4163
4164
4165
4166
4167
4168
4169
4170
4171
4172
4173
4174
4175
4176
4177
4178
4179
4180
4181
4182
4183
4184
4185
4186
4187
4188
4189
4190
4191
4192
4193
4194
4195
4196
4197
4198
4199
4200
4201
4202
4203
4204
4205
4206
4207
4208
4209
4210
4211
4212
4213
4214
4215
4216
4217
4218
4219
4220
4221
4222
4223
4224
4225
4226
4227
4228
4229
4230
4231
4232
4233
4234
4235
4236
4237
4238
4239
4240
4241
4242
4243
4244
4245
4246
4247
4248
4249
4250
4251
4252
4253
4254
4255
4256
4257
4258
4259
4260
4261
4262
4263
4264
4265
4266
4267
4268
4269
4270
4271
4272
4273
4274
4275
4276
4277
4278
4279
4280
4281
4282
4283
4284
4285
4286
4287
4288
def _rocm_sparse_attn_decode_ragged_triton(
    q: torch.Tensor,
    main_cache: torch.Tensor,
    main_indices: torch.Tensor,
    main_indptr: torch.Tensor,
    scale: float,
    attn_sink: torch.Tensor | None,
    nope_head_dim: int,
    rope_head_dim: int,
    extra_cache: torch.Tensor | None = None,
    extra_indices: torch.Tensor | None = None,
    extra_indptr: torch.Tensor | None = None,
    out: torch.Tensor | None = None,
    extra_cache_nan_free: bool = False,
    adaptive_splits: bool = False,
    inv_rope_positions: torch.Tensor | None = None,
    inv_rope_cos_sin_cache: torch.Tensor | None = None,
    out_mxfp8: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
    """Split-K sparse decode; returns the attention output.

    With ``out_mxfp8 = (data, scale)`` the reduce writes MXFP8 instead of
    bf16: ``data`` is [b, h * d] e4m3 and ``scale`` [b, h * d // 32] E8M0, and
    ``data`` viewed as [b, h, d] is returned.
    """
    assert q.ndim == 3, f"expected q=[b,h,d], got {q.shape}"
    assert main_cache.ndim == 3, (
        f"expected main_cache=[blocks,block,bytes], got {main_cache.shape}"
    )
    assert main_indices.ndim == 1, (
        f"expected main_indices=[nnz], got {main_indices.shape}"
    )
    assert main_indptr.ndim == 1, f"expected main_indptr=[b+1], got {main_indptr.shape}"
    assert (
        not q.is_cpu
        and not main_cache.is_cpu
        and not main_indices.is_cpu
        and not main_indptr.is_cpu
    )

    main_indices = _as_int32_contiguous_1d(main_indices)
    main_indptr = _as_int32_contiguous_1d(main_indptr)
    has_attn_sink = attn_sink is not None
    if attn_sink is None:
        attn_sink = torch.empty(1, device=q.device, dtype=torch.float32)
    else:
        attn_sink = attn_sink.contiguous()

    num_queries, num_heads, head_dim = q.shape
    assert main_indptr.numel() == num_queries + 1, (
        f"expected main_indptr shape [{num_queries + 1}], got {main_indptr.shape}"
    )
    _validate_dsv4_sparse_dims(
        head_dim,
        nope_head_dim,
        rope_head_dim,
        "_rocm_sparse_attn_decode_ragged_triton",
    )

    has_extra = (
        extra_cache is not None
        and extra_indices is not None
        and extra_indptr is not None
    )
    assert not extra_cache_nan_free or (_ON_GFX950 and has_extra), (
        "extra_cache_nan_free requires a gfx950 compressed cache with trusted "
        "canonical-writer provenance"
    )
    if has_extra:
        assert extra_cache is not None
        assert extra_indices is not None
        assert extra_indptr is not None
        assert extra_indices.ndim == 1, (
            f"expected extra_indices=[nnz], got {extra_indices.shape}"
        )
        assert extra_indptr.ndim == 1, (
            f"expected extra_indptr=[b+1], got {extra_indptr.shape}"
        )
        extra_indices = _as_int32_contiguous_1d(extra_indices)
        extra_indptr = _as_int32_contiguous_1d(extra_indptr)
        assert extra_indptr.numel() == num_queries + 1, (
            f"expected extra_indptr shape [{num_queries + 1}], got {extra_indptr.shape}"
        )
    else:
        extra_cache = main_cache
        extra_indices = torch.empty(0, device=q.device, dtype=torch.int32)
        extra_indptr = torch.zeros(num_queries + 1, device=q.device, dtype=torch.int32)

    block_h = 16
    out_scale = None
    if out_mxfp8 is not None:
        assert out is None, "out and out_mxfp8 are mutually exclusive"
        assert _ON_GFX950, "the MXFP8 reduce epilogue is gfx950-only"
        assert inv_rope_positions is not None, (
            "the MXFP8 output feeds wo_a, so it must be inverse-RoPE'd first"
        )
        out_data, out_scale = out_mxfp8
        assert out_data.dtype == torch.float8_e4m3fn and out_scale.dtype == (
            torch.uint8
        ), f"expected e4m3/uint8 MXFP8 buffers, got {out_data.dtype}/{out_scale.dtype}"
        assert out_data.shape == (num_queries, num_heads * head_dim), (
            f"expected MXFP8 data [{num_queries}, {num_heads * head_dim}], "
            f"got {tuple(out_data.shape)}"
        )
        assert out_scale.shape == (num_queries, num_heads * head_dim // 32), (
            f"expected MXFP8 scale [{num_queries}, {num_heads * head_dim // 32}], "
            f"got {tuple(out_scale.shape)}"
        )
        assert out_data.stride(-1) == 1 and out_scale.stride(-1) == 1
        out = out_data.view(num_queries, num_heads, head_dim)
    elif out is None:
        out = torch.empty_like(q, dtype=torch.bfloat16)
    else:
        assert out.shape == q.shape, f"expected out shape {q.shape}, got {out.shape}"
        assert out.device == q.device, (
            f"expected out on device {q.device}, got {out.device}"
        )
        assert out.dtype == torch.bfloat16, (
            f"expected out dtype {torch.bfloat16}, got {out.dtype}"
        )
    heads_blocks = triton.cdiv(num_heads, block_h)
    nope_block = triton.next_power_of_2(nope_head_dim)
    comb_dim = nope_head_dim + rope_head_dim
    is_fnuz = current_platform.is_fp8_fnuz()

    if not (_ON_GFX942 or _ON_GFX950):  # Fallback path for un-tuned architectures.
        block_k = 16 if head_dim >= 256 else 32
        _sparse_attn_decode_ragged_kernel[(num_queries, heads_blocks)](
            q,
            main_cache,
            main_indices,
            main_indptr,
            extra_cache,
            extra_indices,
            extra_indptr,
            attn_sink,
            out,
            q.stride(0),
            q.stride(1),
            out.stride(0),
            out.stride(1),
            main_cache.stride(0),
            extra_cache.stride(0),
            main_cache.shape[0] * main_cache.shape[1],
            extra_cache.shape[0] * extra_cache.shape[1],
            main_cache.shape[1],
            extra_cache.shape[1],
            scale,
            num_heads,
            HAS_ATTN_SINK=has_attn_sink,
            HAS_EXTRA=has_extra,
            NOPE_DIM=nope_head_dim,
            NOPE_BLOCK=nope_block,
            ROPE_DIM=rope_head_dim,
            IS_FNUZ_MAIN=is_fnuz,
            IS_FNUZ_EXTRA=False,
            BLOCK_H=block_h,
            BLOCK_K=block_k,
            num_warps=8,
        )
        return out

    block_k = 32  # KV tokens walked per split-K iteration. Tuned on gfx950.
    if _ON_GFX950:
        inv_q = 1.0 / max(1, num_queries)
        avg_main_len = main_indices.numel() * inv_q
        avg_extra_len = (extra_indices.numel() * inv_q) if has_extra else 0.0
        num_splits = _decode_gfx950_num_splits(
            num_queries,
            heads_blocks,
            avg_main_len,
            avg_extra_len,
            block_k,
        )
    else:
        # Average per-query segment lengths, read sync-free from the ragged
        # index sizes, let the split heuristic avoid over-splitting.
        inv_q = 1.0 / max(1, num_queries)
        avg_main_len = main_indices.numel() * inv_q
        avg_extra_len = (extra_indices.numel() * inv_q) if has_extra else 0.0
        num_splits = _decode_num_splits(
            num_queries, heads_blocks, avg_main_len, avg_extra_len, block_k
        )

    base_workgroups = num_queries * heads_blocks
    adaptive_splits = (
        _ON_GFX950 and adaptive_splits and base_workgroups >= 16 and num_splits > 4
    )
    one_wave_splits = (
        max(1, _decode_cu_count() // base_workgroups)
        if adaptive_splits and 16 <= base_workgroups < 64
        else num_splits
    )

    part_m = torch.empty(
        (num_queries, num_splits, num_heads), dtype=torch.float32, device=q.device
    )
    part_l = torch.empty_like(part_m)
    part_acc = torch.empty(
        (num_queries, num_splits, num_heads, comb_dim),
        dtype=torch.float32,
        device=q.device,
    )

    if _ON_GFX950:
        _sparse_attn_decode_gfx950_partial_kernel[
            (num_queries, num_splits, heads_blocks)
        ](
            q,
            main_cache,
            main_indices,
            main_indptr,
            extra_cache,
            extra_indices,
            extra_indptr,
            part_m,
            part_l,
            part_acc,
            q.stride(0),
            q.stride(1),
            main_cache.stride(0),
            extra_cache.stride(0),
            main_cache.shape[0] * main_cache.shape[1],
            extra_cache.shape[0] * extra_cache.shape[1],
            main_cache.shape[1],
            extra_cache.shape[1],
            scale,
            num_heads,
            HAS_EXTRA=has_extra,
            NOPE_DIM=nope_head_dim,
            ROPE_DIM=rope_head_dim,
            IS_FNUZ_MAIN=is_fnuz,
            IS_FNUZ_EXTRA=False,
            TRUST_EXTRA_CACHE_NAN_FREE=extra_cache_nan_free,
            ADAPTIVE_SPLITS=adaptive_splits,
            ONE_WAVE_SPLITS=one_wave_splits,
            BLOCK_H=block_h,
            BLOCK_K=block_k,
            NUM_SPLITS=num_splits,
            NUM_STAGES=1,
            num_warps=4,
            waves_per_eu=0,
        )
    else:
        _sparse_attn_decode_partial_kernel[(num_queries, num_splits, heads_blocks)](
            q,
            main_cache,
            main_indices,
            main_indptr,
            extra_cache,
            extra_indices,
            extra_indptr,
            part_m,
            part_l,
            part_acc,
            q.stride(0),
            q.stride(1),
            main_cache.stride(0),
            extra_cache.stride(0),
            part_m.stride(0),
            part_m.stride(1),
            part_acc.stride(0),
            part_acc.stride(1),
            part_acc.stride(2),
            main_cache.shape[0] * main_cache.shape[1],
            extra_cache.shape[0] * extra_cache.shape[1],
            main_cache.shape[1],
            extra_cache.shape[1],
            scale,
            num_heads,
            HAS_EXTRA=has_extra,
            NOPE_DIM=nope_head_dim,
            NOPE_BLOCK=nope_block,
            ROPE_DIM=rope_head_dim,
            # main_cache = swa_k_cache (C++ encoder, FNUZ on gfx942 / OCP on gfx950).
            # extra_cache = compressed kv_cache (Triton encoder, OCP everywhere).
            # Reading both with a single IS_FNUZ would decode one of them with the
            # wrong FNUZ/OCP scale ratio (~1.87×).
            IS_FNUZ_MAIN=is_fnuz,
            IS_FNUZ_EXTRA=False,
            BLOCK_H=block_h,
            BLOCK_K=block_k,
            NUM_SPLITS=num_splits,
            NUM_STAGES=1,
            num_warps=4,
        )

    fuse_inv_rope = inv_rope_positions is not None
    if fuse_inv_rope:
        assert inv_rope_cos_sin_cache is not None
        assert inv_rope_cos_sin_cache.shape[-1] == rope_head_dim, (
            "fused inverse RoPE expects cos_sin_cache laid out as "
            f"[P, {rope_head_dim}] = cos | sin, got "
            f"{tuple(inv_rope_cos_sin_cache.shape)}"
        )
        assert nope_head_dim % 2 == 0 and rope_head_dim % 2 == 0, (
            "fused inverse RoPE pairs adjacent lanes, so both head dims must "
            f"be even, got nope={nope_head_dim} rope={rope_head_dim}"
        )

    _sparse_attn_decode_reduce_kernel[(num_queries, num_heads)](
        part_m,
        part_l,
        part_acc,
        attn_sink,
        out,
        inv_rope_positions,
        inv_rope_cos_sin_cache,
        out_scale,
        out.stride(0),
        out.stride(1),
        out_scale.stride(0) if out_scale is not None else 0,
        head_dim // 32,
        part_m.stride(0),
        part_m.stride(1),
        part_acc.stride(0),
        part_acc.stride(1),
        part_acc.stride(2),
        inv_rope_cos_sin_cache.stride(0) if inv_rope_cos_sin_cache is not None else 0,
        num_heads,
        HAS_ATTN_SINK=has_attn_sink,
        ADAPTIVE_SPLITS=adaptive_splits,
        COMB_DIM=comb_dim,
        BLOCK_H=1,
        NUM_SPLITS=num_splits,
        SPLITS_PAD=triton.next_power_of_2(num_splits),
        FUSE_INV_ROPE=fuse_inv_rope,
        NOPE=nope_head_dim,
        HALF=rope_head_dim // 2,
        QUANT_OUT=out_scale is not None,
        num_warps=4,
    )
    return out

build_prefill_topk_ragged_indices(topk_indices, token_to_req_indices, query_start_loc, seq_lens, is_valid_token, block_table, block_size, compress_ratio, num_compressed, token_offset, num_rows=-1)

Map prefill top-k rows to a ragged stream of compressed-cache slots.

topk_indices holds local compressed positions for the prefill tokens, which sit at token_offset in the batch; token_to_req_indices, query_start_loc, seq_lens and block_table are batch-wide. block_size is the compressed cache's, i.e. already divided by the ratio.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def build_prefill_topk_ragged_indices(
    topk_indices: torch.Tensor,
    token_to_req_indices: torch.Tensor,
    query_start_loc: torch.Tensor,
    seq_lens: torch.Tensor,
    is_valid_token: torch.Tensor,
    block_table: torch.Tensor,
    block_size: int,
    compress_ratio: int,
    num_compressed: int,
    token_offset: int,
    num_rows: int = -1,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Map prefill top-k rows to a ragged stream of compressed-cache slots.

    ``topk_indices`` holds local compressed positions for the prefill tokens,
    which sit at ``token_offset`` in the batch; ``token_to_req_indices``,
    ``query_start_loc``, ``seq_lens`` and ``block_table`` are batch-wide.
    ``block_size`` is the compressed cache's, i.e. already divided by the ratio.
    """
    topk_indices = topk_indices.reshape(topk_indices.shape[0], -1)
    num_tokens, width = topk_indices.shape
    dense = torch.empty(
        (num_tokens, width), dtype=torch.int32, device=topk_indices.device
    )
    lens = torch.empty(num_tokens, dtype=torch.int32, device=topk_indices.device)
    if num_tokens > 0 and width > 0:
        _prefill_topk_global_slots_kernel[(num_tokens,)](
            dense,
            lens,
            topk_indices,
            topk_indices.stride(0),
            token_to_req_indices,
            query_start_loc,
            seq_lens,
            is_valid_token,
            block_table,
            block_table.stride(0),
            token_offset,
            num_compressed,
            TOPK=width,
            COMPRESS_RATIO=compress_ratio,
            BLOCK_SIZE=block_size,
            BLOCK_W=min(triton.next_power_of_2(width), 1024),
        )
    else:
        lens.zero_()
    return build_ragged_indices_from_dense(dense, lens, num_rows=num_rows)

fp8_mqa_logits_torch(q, kv, weights, cu_seqlen_ks, cu_seqlen_ke)

Compute FP8 MQA logits for a single sequence without KV paging.

Parameters:

  • q

    (Tensor) –

    Query tensor of shape [M, H, D]. Casted to torch.float8_e4m3fn by caller.

  • kv

    (tuple[Tensor, Tensor]) –

    Tuple (k_fp8, k_scales) where k_fp8 has shape [N, D] with dtype torch.float8_e4m3fn and k_scales has shape [N] (or [N, 1]) with dtype torch.float32.

  • weights

    (Tensor) –

    weights of shape [M, H], dtype torch.float32.

  • cu_seqlen_ks

    (Tensor) –

    Start indices (inclusive) for valid K per query position, shape [M], dtype int32.

  • cu_seqlen_ke

    (Tensor) –

    End indices (exclusive) for valid K per query position, shape [M], dtype int32.

Returns:

  • Tensor –

    Logits tensor of shape [M, N], dtype torch.float32.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def fp8_mqa_logits_torch(
    q: torch.Tensor,
    kv: tuple[torch.Tensor, torch.Tensor],
    weights: torch.Tensor,
    cu_seqlen_ks: torch.Tensor,
    cu_seqlen_ke: torch.Tensor,
) -> torch.Tensor:
    """Compute FP8 MQA logits for a single sequence without KV paging.

    Args:
        q: Query tensor of shape [M, H, D]. Casted to
            `torch.float8_e4m3fn` by caller.
        kv: Tuple `(k_fp8, k_scales)` where `k_fp8` has shape [N, D] with
            dtype `torch.float8_e4m3fn` and `k_scales` has shape [N] (or
            [N, 1]) with dtype `torch.float32`.
        weights: weights of shape [M, H], dtype `torch.float32`.
        cu_seqlen_ks: Start indices (inclusive) for valid K per query position,
            shape [M], dtype int32.
        cu_seqlen_ke: End indices (exclusive) for valid K per query position,
            shape [M], dtype int32.

    Returns:
        Logits tensor of shape [M, N], dtype `torch.float32`.

    """
    k_fp8, scale = kv
    seq_len_kv = k_fp8.shape[0]
    k = k_fp8.to(torch.bfloat16)
    q = q.to(torch.bfloat16)
    device = q.device

    mask_lo = (
        torch.arange(0, seq_len_kv, device=device)[None, :] >= cu_seqlen_ks[:, None]
    )
    mask_hi = (
        torch.arange(0, seq_len_kv, device=device)[None, :] < cu_seqlen_ke[:, None]
    )
    mask = mask_lo & mask_hi

    # ``score`` is [H, M, N]; ``scale`` is the per-KV-token scale, which
    # vLLM callers hand us as ``[N, 1]`` (a ``[N, 4]`` uint8 buffer cast
    # to fp32). PyTorch right-aligns dimensions for broadcasting, so a
    # naked ``score * scale`` would align ``scale``'s leading dim with
    # ``score``'s M dim and raise a shape mismatch. Flatten to ``[N]`` so
    # broadcasting lines up with the last dim of ``score``.
    score = torch.einsum("mhd,nd->hmn", q, k).float() * scale.reshape(-1)
    logits = (score.relu() * weights.unsqueeze(-1).transpose(0, 1)).sum(dim=0)
    logits = logits.masked_fill(~mask, float("-inf"))

    return logits

rocm_fp8_mqa_logits(q, kv, weights, cu_seqlen_ks, cu_seqlen_ke)

Compute FP8 MQA logits for a single sequence without KV paging.

Parameters:

  • q

    (Tensor) –

    Query tensor of shape [M, H, D]. Casted to torch.float8_e4m3fn by caller.

  • kv

    (tuple[Tensor, Tensor]) –

    Tuple (k_fp8, k_scales) where k_fp8 has shape [N, D] with dtype torch.float8_e4m3fn and k_scales has shape [N] (or [N, 1]) with dtype torch.float32.

  • weights

    (Tensor) –

    weights of shape [M, H], dtype torch.float32.

  • cu_seqlen_ks

    (Tensor) –

    Start indices (inclusive) for valid K per query position, shape [M], dtype int32.

  • cu_seqlen_ke

    (Tensor) –

    End indices (exclusive) for valid K per query position, shape [M], dtype int32.

Returns:

  • Tensor –

    Logits tensor of shape [M, N], dtype torch.float32.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_fp8_mqa_logits(
    q: torch.Tensor,
    kv: tuple[torch.Tensor, torch.Tensor],
    weights: torch.Tensor,
    cu_seqlen_ks: torch.Tensor,
    cu_seqlen_ke: torch.Tensor,
) -> torch.Tensor:
    """Compute FP8 MQA logits for a single sequence without KV paging.

    Args:
        q: Query tensor of shape [M, H, D]. Casted to
            `torch.float8_e4m3fn` by caller.
        kv: Tuple `(k_fp8, k_scales)` where `k_fp8` has shape [N, D] with
            dtype `torch.float8_e4m3fn` and `k_scales` has shape [N] (or
            [N, 1]) with dtype `torch.float32`.
        weights: weights of shape [M, H], dtype `torch.float32`.
        cu_seqlen_ks: Start indices (inclusive) for valid K per query position,
            shape [M], dtype int32.
        cu_seqlen_ke: End indices (exclusive) for valid K per query position,
            shape [M], dtype int32.

    Returns:
        Logits tensor of shape [M, N], dtype `torch.float32`.

    """
    from vllm._aiter_ops import rocm_aiter_ops

    k_fp8, scale = kv

    if _ON_GFX942 and rocm_aiter_ops.is_enabled():
        from aiter.ops.flydsl import flydsl_fp8_mqa_logits

        return flydsl_fp8_mqa_logits(
            q, k_fp8, scale, weights, cu_seqlen_ks, cu_seqlen_ke
        )

    aiter_mqa_logits_module = None
    if rocm_aiter_ops.is_enabled() or rocm_aiter_ops.is_rdna_aiter_enabled():
        aiter_mqa_logits_module = mqa_logits_module()

    if aiter_mqa_logits_module is not None:
        fp8_mqa_logits = aiter_mqa_logits_module.fp8_mqa_logits
        return fp8_mqa_logits(q, k_fp8, scale, weights, cu_seqlen_ks, cu_seqlen_ke)
    else:
        return fp8_mqa_logits_torch(q, kv, weights, cu_seqlen_ks, cu_seqlen_ke)

rocm_fp8_paged_mqa_logits(q_fp8, kv_cache_fp8, weights, context_lens, block_tables, schedule_metadata, max_model_len, *, compress_ratio=1)

Compute FP8 MQA logits using paged KV-cache.

Parameters:

  • q_fp8

    (Tensor) –

    Query tensor of shape [B, next_n, H, D]. Casted to torch.float8_e4m3fn by caller.

  • kv_cache_fp8

    (Tensor) –

    Paged KV-cache in packed FP8+scale layout with shape [num_blocks, block_size, 1, D+4], dtype torch.uint8.

  • weights

    (Tensor) –

    Tensor of shape [B * next_n, H], dtype torch.float32.

  • context_lens

    (Tensor) –

    Tensor of shape [B], dtype int32; effective context length for each batch element.

  • block_tables

    (Tensor) –

    Tensor of shape [B, max_blocks], dtype int32; maps logical block indices to physical blocks in the paged cache.

  • schedule_metadata

    (Tensor) –

    Returned by get_paged_mqa_logits_metadata; used to distribute work across SMs.

  • max_model_len

    (int) –

    Maximum sequence length used to size the logits output.

  • compress_ratio

    (int, default: 1 ) –

    C4A (4) takes block-flat Triton; 1 and 2 stay on AITER.

Returns:

  • Tensor –

    Logits tensor of shape [B * next_n, max_model_len], dtype

  • Tensor –

    torch.float32.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_fp8_paged_mqa_logits(
    q_fp8: torch.Tensor,
    kv_cache_fp8: torch.Tensor,
    weights: torch.Tensor,
    context_lens: torch.Tensor,
    block_tables: torch.Tensor,
    schedule_metadata: torch.Tensor,
    max_model_len: int,
    *,
    compress_ratio: int = 1,
) -> torch.Tensor:
    """Compute FP8 MQA logits using paged KV-cache.

    Args:
        q_fp8: Query tensor of shape [B, next_n, H, D]. Casted to
            `torch.float8_e4m3fn` by caller.
        kv_cache_fp8: Paged KV-cache in packed FP8+scale layout with shape
            [num_blocks, block_size, 1, D+4], dtype `torch.uint8`.
        weights: Tensor of shape [B * next_n, H], dtype `torch.float32`.
        context_lens: Tensor of shape [B], dtype int32; effective context length
            for each batch element.
        block_tables: Tensor of shape [B, max_blocks], dtype int32; maps logical
            block indices to physical blocks in the paged cache.
        schedule_metadata: Returned by `get_paged_mqa_logits_metadata`;
            used to distribute work across SMs.
        max_model_len: Maximum sequence length used to size the logits output.
        compress_ratio: C4A (4) takes block-flat Triton; 1 and 2 stay on AITER.

    Returns:
        Logits tensor of shape [B * next_n, max_model_len], dtype
        `torch.float32`.

    """
    from vllm._aiter_ops import rocm_aiter_ops

    batch_size, next_n = q_fp8.shape[:2]
    block_size = kv_cache_fp8.shape[1]

    # C4A only: Flash/DSv3.2 also skip insert but still write SHUFFLE.
    if (
        (_ON_GFX950 or _ON_GFX942)
        and _indexer_k_is_c4a_block_flat(compress_ratio)
        and block_size > 1
    ):
        if block_size % 64 == 0:
            return rocm_fp8_paged_mqa_logits_triton(
                q_fp8, kv_cache_fp8, weights, context_lens, block_tables, max_model_len
            )
        # Non 64-aligned page size (not used in prod): eager torch ref.
        return fp8_paged_mqa_logits_torch(
            q_fp8, kv_cache_fp8, weights, context_lens, block_tables, max_model_len
        )

    aiter_paged_mqa_logits_module = None

    if rocm_aiter_ops.is_enabled() or rocm_aiter_ops.is_rdna_aiter_enabled():
        aiter_paged_mqa_logits_module = paged_mqa_logits_module()

    if aiter_paged_mqa_logits_module is not None:
        if _ON_GFX942 or _ON_GFX950:
            deepgemm_fp8_paged_mqa_logits = (
                aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits
            )
            batch_size, next_n, heads, _ = q_fp8.shape
            (out_logits,) = current_workspace_manager().get_simultaneous(
                ((batch_size * next_n, max_model_len), torch.float32),
            )
            deepgemm_fp8_paged_mqa_logits(
                q_fp8,
                kv_cache_fp8,
                weights,
                out_logits,
                context_lens,
                block_tables,
                max_model_len,
                ChunkK=256,
                Preshuffle=block_size > 1,
                KVBlockSize=block_size,
                WavePerEU=2,
            )
            return out_logits
        deepgemm_fp8_paged_mqa_logits_stage1 = (
            aiter_paged_mqa_logits_module.deepgemm_fp8_paged_mqa_logits_stage1
        )
        batch_size, next_n, heads, _ = q_fp8.shape
        (out_qk,) = current_workspace_manager().get_simultaneous(
            ((heads, batch_size * next_n, max_model_len), torch.float32),
        )
        out_qk.fill_(float("-inf"))
        deepgemm_fp8_paged_mqa_logits_stage1(
            q_fp8,
            kv_cache_fp8,
            weights,
            out_qk,
            context_lens,
            block_tables,
            max_model_len,
            ChunkQ=heads,
        )
        return out_qk.sum(dim=0)
    else:
        return fp8_paged_mqa_logits_torch(
            q_fp8, kv_cache_fp8, weights, context_lens, block_tables, max_model_len
        )

rocm_fp8_paged_mqa_logits_triton(q_fp8, kv_cache_fp8, weights, context_lens, block_tables, max_model_len)

Triton paged MQA-logits for decode and MTP; matches the torch ref but has no host sync, so it is safe to capture under a full CUDA graph.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_fp8_paged_mqa_logits_triton(
    q_fp8: torch.Tensor,
    kv_cache_fp8: torch.Tensor,
    weights: torch.Tensor,
    context_lens: torch.Tensor,
    block_tables: torch.Tensor,
    max_model_len: int,
) -> torch.Tensor:
    """Triton paged MQA-logits for decode and MTP; matches the torch ref but
    has no host sync, so it is safe to capture under a full CUDA graph."""
    batch_size, next_n, num_heads, head_size = q_fp8.shape
    block_size = kv_cache_fp8.shape[1]
    BLOCK_KV = 64
    assert block_size % BLOCK_KV == 0

    fp8_dtype = current_platform.fp8_dtype()
    num_blocks = kv_cache_fp8.shape[0]
    kv_flat = kv_cache_fp8.reshape(
        num_blocks, -1
    )  # uint8 [num_blocks, block_size*(D+4)]
    kv_val = kv_flat.view(fp8_dtype)  # [num_blocks, block_size*(D+4)] fp8
    kv_scale = kv_flat.view(torch.float32)  # [num_blocks, block_size*(D+4)//4] fp32

    cl = context_lens.reshape(-1)
    ctx_per_row = not (next_n > 1 and cl.numel() == batch_size)

    max_blocks = block_tables.shape[1]
    (out_logits,) = current_workspace_manager().get_simultaneous(
        ((batch_size * next_n, max_model_len), torch.float32),
    )

    # Memory-bound over the KV range: split each row's keys across programs so
    # few-row / long-context launches still fill the GPU. Cap splits at the
    # device CU count (304 on gfx942, 256 on gfx950) rather than a gfx950-sized
    # constant. All terms are static at launch, so the grid stays CUDA-graph-safe.
    rows = batch_size * next_n
    tiles_cap = (max_model_len + BLOCK_KV - 1) // BLOCK_KV
    N_SPLITS = max(1, min(max(1, _decode_cu_count()), tiles_cap, 1024 // rows))
    _fp8_paged_mqa_logits_decode_kernel[(rows, N_SPLITS)](
        q_fp8,
        kv_val,
        kv_scale,
        weights,
        cl,
        block_tables,
        out_logits,
        q_fp8.stride(0),
        q_fp8.stride(1),
        q_fp8.stride(2),
        weights.stride(0),
        kv_val.stride(0),
        kv_scale.stride(0),
        (block_size * head_size) // 4,
        block_tables.stride(0),
        out_logits.stride(0),
        max_blocks,
        max_model_len,
        NUM_HEADS=num_heads,
        HEAD_SIZE=head_size,
        BLOCK_SIZE=block_size,
        BLOCK_KV=BLOCK_KV,
        N_SPLITS=N_SPLITS,
        NEXT_N=next_n,
        CTX_PER_ROW=ctx_per_row,
        num_warps=4,
        num_stages=2,
    )
    return out_logits

rocm_inv_rope_einsum(rotary_emb, o, positions, rope_head_dim, n_local_groups, o_lora_rank, wo_a, inverse_rope=True)

Inverse-RoPE + WO_A bmm path used on ROCm.

Fuses the inverse GPT-J RoPE into one Triton kernel and caches the bf16 wo_a weight so the per-step dequant disappears. Callers whose attention already rotated every row pass inverse_rope=False; that is a property of the attention backend, not of the batch, so it stays constant across steps and is safe to read from compiled code.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_inv_rope_einsum(
    rotary_emb: torch.nn.Module,
    o: torch.Tensor,
    positions: torch.Tensor,
    rope_head_dim: int,
    n_local_groups: int,
    o_lora_rank: int,
    wo_a: torch.nn.Module,
    inverse_rope: bool = True,
) -> torch.Tensor:
    """Inverse-RoPE + WO_A bmm path used on ROCm.

    Fuses the inverse GPT-J RoPE into one Triton kernel and caches the bf16
    wo_a weight so the per-step dequant disappears. Callers whose attention
    already rotated every row pass ``inverse_rope=False``; that is a property
    of the attention backend, not of the batch, so it stays constant across
    steps and is safe to read from compiled code.
    """
    if inverse_rope:
        o_ref = _fused_inverse_rope_gptj(
            o, positions, rotary_emb.cos_sin_cache, rope_head_dim
        )
    else:
        assert o.dtype == torch.bfloat16, (
            "a pre-rotated attention output feeds the wo_a bmm directly, so it "
            f"must already be bf16, got {o.dtype}"
        )
        o_ref = o
    o_ref = o_ref.reshape(o.shape[0], n_local_groups, -1)

    wo_a_weight = _get_cached_wo_a_bf16(
        wo_a, n_local_groups, o_lora_rank, o_ref.shape[-1]
    )

    return torch.einsum("tgd,grd->tgr", o_ref, wo_a_weight)

rocm_inverse_rope_mxfp8_rows(o, positions, cos_sin_cache, rope_head_dim, out_data, out_scale)

Inverse-RoPE bf16 attention rows and MXFP8-quantize them for wo_a.

The counterpart of rocm_inverse_rope_rows_ for layers whose attention output is MXFP8: rows the decode reduce did not emit (prefill) go through here. o is [T, H, D]; out_data [T, H * D] e4m3 and out_scale [T, H * D // 32] E8M0, the layout the reduce epilogue writes.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_inverse_rope_mxfp8_rows(
    o: torch.Tensor,
    positions: torch.Tensor,
    cos_sin_cache: torch.Tensor,
    rope_head_dim: int,
    out_data: torch.Tensor,
    out_scale: torch.Tensor,
) -> None:
    """Inverse-RoPE bf16 attention rows and MXFP8-quantize them for wo_a.

    The counterpart of ``rocm_inverse_rope_rows_`` for layers whose attention
    output is MXFP8: rows the decode reduce did not emit (prefill) go through
    here. ``o`` is [T, H, D]; ``out_data`` [T, H * D] e4m3 and ``out_scale``
    [T, H * D // 32] E8M0, the layout the reduce epilogue writes.
    """
    num_tokens, num_heads, head_dim = o.shape
    if num_tokens == 0:
        return
    assert o.stride(-1) == 1 and out_data.stride(-1) == 1 and out_scale.stride(-1) == 1
    assert out_data.shape == (num_tokens, num_heads * head_dim)
    assert out_scale.shape == (num_tokens, num_heads * head_dim // 32)
    assert cos_sin_cache.shape[-1] == rope_head_dim
    _inverse_rope_mxfp8_quant_kernel[(num_tokens, num_heads)](
        o,
        out_data,
        out_scale,
        positions,
        cos_sin_cache,
        o.stride(0),
        o.stride(1),
        out_data.stride(0),
        out_scale.stride(0),
        cos_sin_cache.stride(0),
        HEAD_DIM=head_dim,
        NOPE=head_dim - rope_head_dim,
        HALF=rope_head_dim // 2,
        num_warps=4,
    )

rocm_inverse_rope_rows_(o, positions, cos_sin_cache, rope_head_dim)

Inverse-RoPE attention output rows in place.

For rows no attention kernel rotated in its epilogue. Call it from the eager attention segment: which rows still owe a rotation depends on the prefill/decode split, and the o_proj that used to do this runs inside the compiled region, where a batch-dependent Python value would be frozen at trace time.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_inverse_rope_rows_(
    o: torch.Tensor,
    positions: torch.Tensor,
    cos_sin_cache: torch.Tensor,
    rope_head_dim: int,
) -> None:
    """Inverse-RoPE attention output rows in place.

    For rows no attention kernel rotated in its epilogue. Call it from the
    eager attention segment: which rows still owe a rotation depends on the
    prefill/decode split, and the o_proj that used to do this runs inside the
    compiled region, where a batch-dependent Python value would be frozen at
    trace time.
    """
    if o.shape[0] == 0:
        return
    _fused_inverse_rope_gptj(o, positions, cos_sin_cache, rope_head_dim, out=o)

rocm_mxfp8_wo_a_bmm(a, a_scale, wo_a, n_groups, o_lora_rank)

Grouped MXFP8 wo_a: out[t, g, :] = a[t, g, :] @ W[g].T, bf16 out.

a is the [T, G * K] e4m3 attention output and a_scale its [T, G * K // 32] E8M0 scales, as the sparse decode reduce writes them. The weight is the checkpoint's MXFP8 wo_a as loaded, [G * R, K] with either [G * R // 32, K // 32] block scales or [G * R, K // 32] per-row scales, so there is no dequantized copy to keep. Returns [T, G * R].

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_mxfp8_wo_a_bmm(
    a: torch.Tensor,
    a_scale: torch.Tensor,
    wo_a: torch.nn.Module,
    n_groups: int,
    o_lora_rank: int,
) -> torch.Tensor:
    """Grouped MXFP8 wo_a: ``out[t, g, :] = a[t, g, :] @ W[g].T``, bf16 out.

    ``a`` is the [T, G * K] e4m3 attention output and ``a_scale`` its
    [T, G * K // 32] E8M0 scales, as the sparse decode reduce writes them.
    The weight is the checkpoint's MXFP8 ``wo_a`` as loaded, [G * R, K] with
    either [G * R // 32, K // 32] block scales or [G * R, K // 32] per-row
    scales, so there is no dequantized copy to keep.
    Returns [T, G * R].
    """
    return torch.ops.vllm.rocm_dsv41_mxfp8_wo_a_bmm(
        a, a_scale, wo_a.weight, wo_a.weight_scale, n_groups, o_lora_rank
    )

rocm_sparse_attn_decode(q, kv_cache, swa_k_cache, swa_only, topk_indices, topk_lens, swa_indices, swa_lens, swa_ragged_indices, swa_ragged_indptr, topk_ragged_indices, topk_ragged_indptr, attn_sink, scale, head_dim, nope_head_dim, rope_head_dim, output, extra_cache_nan_free=False, adaptive_splits=False, inv_rope_positions=None, inv_rope_cos_sin_cache=None, output_mxfp8=None)

Run sparse MLA decode into output.

Passing inv_rope_positions folds the inverse RoPE into the reduce epilogue. Returns how many leading rows of output came back rotated, so a caller mixing in a decode path that does not fuse still knows what it owes the standalone pass. Read it from the eager attention segment only.

output_mxfp8 = (data, scale) replaces output: the reduce also MXFP8-quantizes the rotated rows for the FP8 wo_a (see _rocm_sparse_attn_decode_ragged_triton). It needs gfx950 and the fused inverse RoPE, and always covers every row.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_sparse_attn_decode(
    q: torch.Tensor,
    kv_cache: torch.Tensor | None,
    swa_k_cache: torch.Tensor,
    swa_only: bool,
    topk_indices: torch.Tensor | None,
    topk_lens: torch.Tensor | None,
    swa_indices: torch.Tensor,
    swa_lens: torch.Tensor,
    swa_ragged_indices: torch.Tensor | None,
    swa_ragged_indptr: torch.Tensor | None,
    topk_ragged_indices: torch.Tensor | None,
    topk_ragged_indptr: torch.Tensor | None,
    attn_sink: torch.Tensor | None,
    scale: float,
    head_dim: int,
    nope_head_dim: int,
    rope_head_dim: int,
    output: torch.Tensor | None,
    extra_cache_nan_free: bool = False,
    adaptive_splits: bool = False,
    inv_rope_positions: torch.Tensor | None = None,
    inv_rope_cos_sin_cache: torch.Tensor | None = None,
    output_mxfp8: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> int:
    """Run sparse MLA decode into ``output``.

    Passing ``inv_rope_positions`` folds the inverse RoPE into the reduce
    epilogue. Returns how many leading rows of ``output`` came back rotated,
    so a caller mixing in a decode path that does not fuse still knows what it
    owes the standalone pass. Read it from the eager attention segment only.

    ``output_mxfp8 = (data, scale)`` replaces ``output``: the reduce also
    MXFP8-quantizes the rotated rows for the FP8 wo_a (see
    ``_rocm_sparse_attn_decode_ragged_triton``). It needs gfx950 and the fused
    inverse RoPE, and always covers every row.
    """
    assert swa_k_cache.dtype == torch.uint8, (
        "ROCm Triton sparse decode expects uint8 fp8_ds_mla SWA cache, "
        f"got {swa_k_cache.dtype}"
    )
    _validate_dsv4_sparse_dims(
        head_dim,
        nope_head_dim,
        rope_head_dim,
        "rocm_sparse_attn_decode",
    )

    main_indices = swa_indices.reshape(swa_indices.shape[0], -1)

    extra_cache = None
    extra_indices = None
    if not swa_only:
        assert kv_cache is not None
        assert topk_indices is not None or (
            topk_ragged_indices is not None and topk_ragged_indptr is not None
        )
        assert kv_cache.dtype == torch.uint8, (
            "ROCm Triton sparse decode expects uint8 fp8_ds_mla extra cache, "
            f"got {kv_cache.dtype}"
        )
        extra_cache = kv_cache
        if topk_indices is not None:
            extra_indices = topk_indices.reshape(topk_indices.shape[0], -1)

    if output_mxfp8 is not None:
        assert output is None, "output and output_mxfp8 are mutually exclusive"
        direct_out = None
    else:
        assert output is not None
        direct_out = output if _ON_GFX950 and output.dtype == torch.bfloat16 else None
    attn_out = _rocm_sparse_attn_decode_triton(
        q=q,
        main_cache=swa_k_cache,
        main_indices=main_indices,
        scale=scale,
        attn_sink=None if attn_sink is None else attn_sink[: q.shape[1]],
        nope_head_dim=nope_head_dim,
        rope_head_dim=rope_head_dim,
        extra_cache=extra_cache,
        extra_indices=extra_indices,
        main_lengths=swa_lens,
        extra_lengths=topk_lens,
        main_ragged_indices=swa_ragged_indices,
        main_ragged_indptr=swa_ragged_indptr,
        extra_ragged_indices=topk_ragged_indices,
        extra_ragged_indptr=topk_ragged_indptr,
        out=direct_out,
        extra_cache_nan_free=extra_cache_nan_free,
        adaptive_splits=adaptive_splits,
        inv_rope_positions=inv_rope_positions,
        inv_rope_cos_sin_cache=inv_rope_cos_sin_cache,
        out_mxfp8=output_mxfp8,
    )
    if output_mxfp8 is not None:
        return q.shape[0]
    assert output is not None
    if direct_out is None:
        output.copy_(attn_out.to(output.dtype))
    return output.shape[0] if inv_rope_positions is not None else 0

rocm_sparse_attn_decode_bf16(q, kv, scale, head_dim, nope_head_dim, rope_head_dim, attn_sink, output, ragged_indices, ragged_indptr, num_splits)

Run split-K sparse attention over decode rows using an unquantized KV cache.

Parameters:

  • q

    (Tensor) –

    Decode queries laid out as [sq, h, d].

  • kv

    (Tensor) –

    KV cache laid out as [skv, 1, d].

  • scale

    (float) –

    Softmax scale.

  • head_dim

    (int) –

    Post-absorption head width.

  • nope_head_dim

    (int) –

    NoPE width of head_dim.

  • rope_head_dim

    (int) –

    RoPE width of head_dim.

  • attn_sink

    (Tensor | None) –

    Optional per-head sink logits.

  • output

    (Tensor) –

    Destination, written in place.

  • ragged_indices

    (Tensor) –

    Flattened per-query KV slots.

  • ragged_indptr

    (Tensor) –

    Segment offsets into ragged_indices, [sq + 1].

  • num_splits

    (int) –

    KV splits per query, from :func:rocm_sparse_decode_bf16_num_splits.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_sparse_attn_decode_bf16(
    q: torch.Tensor,
    kv: torch.Tensor,
    scale: float,
    head_dim: int,
    nope_head_dim: int,
    rope_head_dim: int,
    attn_sink: torch.Tensor | None,
    output: torch.Tensor,
    ragged_indices: torch.Tensor,
    ragged_indptr: torch.Tensor,
    num_splits: int,
) -> None:
    """Run split-K sparse attention over decode rows using an unquantized KV cache.

    Args:
        q: Decode queries laid out as ``[sq, h, d]``.
        kv: KV cache laid out as ``[skv, 1, d]``.
        scale: Softmax scale.
        head_dim: Post-absorption head width.
        nope_head_dim: NoPE width of ``head_dim``.
        rope_head_dim: RoPE width of ``head_dim``.
        attn_sink: Optional per-head sink logits.
        output: Destination, written in place.
        ragged_indices: Flattened per-query KV slots.
        ragged_indptr: Segment offsets into ``ragged_indices``, ``[sq + 1]``.
        num_splits: KV splits per query, from
            :func:`rocm_sparse_decode_bf16_num_splits`.

    """
    assert kv.ndim == 3 and kv.shape[1] == 1, (
        f"ROCm Triton sparse decode expects kv=[skv,1,d], got {kv.shape}"
    )
    _validate_sparse_dims(
        head_dim,
        nope_head_dim,
        rope_head_dim,
        "rocm_sparse_attn_decode_bf16",
    )
    num_queries, num_heads = q.shape[0], q.shape[1]
    direct = output.shape[-1] == head_dim
    out = (
        output
        if direct
        else torch.empty(
            (num_queries, num_heads, head_dim),
            dtype=output.dtype,
            device=output.device,
        )
    )
    _rocm_sparse_attn_decode_ragged_bf16_triton(
        q=q,
        kv=kv.squeeze(1),
        indices=ragged_indices,
        indptr=ragged_indptr,
        scale=scale,
        attn_sink=None if attn_sink is None else attn_sink[: q.shape[1]],
        nope_head_dim=nope_head_dim,
        rope_head_dim=rope_head_dim,
        num_splits=num_splits,
        out=out,
    )
    if not direct:
        output.copy_(out[..., : output.shape[-1]])

rocm_sparse_decode_bf16_num_splits(num_queries, num_heads, sparse_len)

Number or kv splits in splitK for the sparse bf16 decode, or 1 for single-pass.

Parameters:

  • num_queries

    (int) –

    Decode rows in the batch.

  • num_heads

    (int) –

    Query heads per row.

  • sparse_len

    (int) –

    Longest selected KV run any decode row can walk.

Returns:

  • int –

    The split count, or 1 when the caller should use the single-pass kernel.

Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
def rocm_sparse_decode_bf16_num_splits(
    num_queries: int, num_heads: int, sparse_len: int
) -> int:
    """Number or kv splits in splitK for the sparse bf16 decode, or 1 for single-pass.

    Args:
        num_queries: Decode rows in the batch.
        num_heads: Query heads per row.
        sparse_len: Longest selected KV run any decode row can walk.

    Returns:
        The split count, or 1 when the caller should use the single-pass kernel.

    """
    if sparse_len < _SPARSE_DECODE_BF16_MIN_SPLIT_LEN:
        return 1
    block_k = _SPARSE_DECODE_BF16_BLOCK_K
    heads_blocks = triton.cdiv(num_heads, _SPARSE_DECODE_BF16_BLOCK_H)
    select = _decode_gfx950_num_splits if _ON_GFX950 else _decode_num_splits
    num_splits = select(num_queries, heads_blocks, sparse_len, 0.0, block_k)
    # Number of splits cannot exceed the available k tiles
    return max(1, min(num_splits, math.ceil(sparse_len / block_k)))