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. -
rocm_sparse_attn_decode_bf16–Run split-K sparse attention over decode rows using an unquantized KV cache.
-
rocm_sparse_decode_bf16_num_splits–Number or kv splits in splitK for the sparse bf16 decode, or 1 for single-pass.
_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
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 | |
_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_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
dtrailing elements.
Returns:
-
Tensor–The attention output,
outwhen provided.
Source code in vllm/v1/attention/ops/rocm_aiter_mla_sparse.py
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 | |
_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 | |
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
4501 4502 4503 4504 4505 4506 4507 4508 4509 4510 4511 4512 4513 4514 4515 4516 4517 4518 4519 4520 4521 4522 4523 4524 4525 4526 4527 4528 4529 4530 4531 4532 4533 4534 4535 4536 4537 4538 4539 4540 4541 4542 4543 4544 4545 4546 4547 4548 4549 4550 4551 4552 4553 4554 4555 4556 4557 4558 4559 4560 4561 4562 4563 4564 4565 4566 4567 4568 4569 4570 4571 4572 4573 4574 4575 4576 4577 4578 4579 4580 4581 4582 4583 4584 4585 4586 4587 4588 4589 4590 4591 4592 4593 4594 4595 4596 4597 4598 4599 4600 | |
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
4435 4436 4437 4438 4439 4440 4441 4442 4443 4444 4445 4446 4447 4448 4449 4450 4451 4452 4453 4454 4455 4456 4457 4458 4459 4460 4461 4462 4463 4464 4465 4466 4467 4468 4469 4470 4471 4472 4473 4474 4475 4476 4477 4478 4479 4480 4481 4482 4483 4484 4485 4486 4487 4488 4489 4490 4491 4492 4493 4494 4495 4496 4497 4498 | |
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.