vllm.v1.attention.ops.rocm_aiter_mla_sparse
¶
Functions:
-
build_prefill_topk_ragged_indices–Map prefill top-k rows to a ragged stream of compressed-cache slots.
-
fp8_mqa_logits_torch–Compute FP8 MQA logits for a single sequence without KV paging.
-
rocm_fp8_mqa_logits–Compute FP8 MQA logits for a single sequence without KV paging.
-
rocm_fp8_paged_mqa_logits–Compute FP8 MQA logits using paged KV-cache.
-
rocm_fp8_paged_mqa_logits_triton–Triton paged MQA-logits for decode and MTP; matches the torch ref but
-
rocm_inv_rope_einsum–Inverse-RoPE + WO_A bmm path used on ROCm.
-
rocm_inverse_rope_mxfp8_rows–Inverse-RoPE bf16 attention rows and MXFP8-quantize them for wo_a.
-
rocm_inverse_rope_rows_–Inverse-RoPE attention output rows in place.
-
rocm_mxfp8_wo_a_bmm–Grouped MXFP8 wo_a:
out[t, g, :] = a[t, g, :] @ W[g].T, bf16 out. -
rocm_sparse_attn_decode–Run sparse MLA decode into
output.
_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
_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 / scaptures the partial compute: each wave walks roughlytotal_tokens / stokens and there arewavesof them, so dividing bysmakes more splits cheaper until they spill into extra waves.mu * wavescharges 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
3556 3557 3558 3559 3560 3561 3562 3563 3564 3565 3566 3567 3568 3569 3570 3571 3572 3573 3574 3575 3576 3577 3578 3579 3580 3581 3582 3583 3584 3585 3586 3587 3588 3589 3590 3591 3592 3593 3594 3595 3596 3597 3598 3599 3600 3601 3602 3603 3604 3605 3606 3607 3608 3609 3610 3611 3612 3613 3614 3615 3616 3617 3618 3619 3620 3621 | |
_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
_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
_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
_indexer_k_is_c4a_block_flat(compress_ratio)
¶
_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
_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
_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
_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
_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
_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 | |
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
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_e4m3fnby caller. -
(kv¶tuple[Tensor, Tensor]) –Tuple
(k_fp8, k_scales)wherek_fp8has shape [N, D] with dtypetorch.float8_e4m3fnandk_scaleshas shape [N] (or [N, 1]) with dtypetorch.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
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_e4m3fnby caller. -
(kv¶tuple[Tensor, Tensor]) –Tuple
(k_fp8, k_scales)wherek_fp8has shape [N, D] with dtypetorch.float8_e4m3fnandk_scaleshas shape [N] (or [N, 1]) with dtypetorch.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
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_e4m3fnby 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:
Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 | |
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
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
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
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
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
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
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 | |