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_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
3667
3668
3669
3670
3671
3672
3673
3674
3675
3676
3677
3678
3679
3680
3681
3682
3683
3684
3685
3686
3687
3688
3689
3690
3691
3692
3693
3694
3695
3696
3697
3698
3699
3700
3701
3702
3703
3704
3705
3706
3707
3708
3709
3710
3711
3712
3713
3714
3715
3716
3717
3718
3719
3720
3721
3722
3723
3724
3725
3726
3727
3728
3729
3730
3731
3732
3733
3734
3735
3736
3737
3738
3739
3740
3741
3742
3743
3744
3745
3746
3747
3748
3749
3750
3751
3752
3753
3754
3755
3756
3757
3758
3759
3760
3761
3762
3763
3764
3765
3766
3767
3768
3769
3770
3771
3772
3773
3774
3775
3776
3777
3778
3779
3780
3781
3782
3783
3784
3785
3786
3787
3788
3789
3790
3791
3792
3793
3794
3795
3796
3797
3798
3799
3800
3801
3802
3803
3804
3805
3806
3807
3808
3809
3810
3811
3812
3813
3814
3815
3816
3817
3818
3819
3820
3821
3822
3823
3824
3825
3826
3827
3828
3829
3830
3831
3832
3833
3834
3835
3836
3837
3838
3839
3840
3841
3842
3843
3844
3845
3846
3847
3848
3849
3850
3851
3852
3853
3854
3855
3856
3857
3858
3859
3860
3861
3862
3863
3864
3865
3866
3867
3868
3869
3870
3871
3872
3873
3874
3875
3876
3877
3878
3879
3880
3881
3882
3883
3884
3885
3886
3887
3888
3889
3890
3891
3892
3893
3894
3895
3896
3897
3898
3899
3900
3901
3902
3903
3904
3905
3906
3907
3908
3909
3910
3911
3912
3913
3914
3915
3916
3917
3918
3919
3920
3921
3922
3923
3924
3925
3926
3927
3928
3929
3930
3931
3932
3933
3934
3935
3936
3937
3938
3939
3940
3941
3942
3943
3944
3945
3946
3947
3948
3949
3950
3951
3952
3953
3954
3955
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
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