From 3ebd35f41949c960676432f182c39ff56b7f69ab Mon Sep 17 00:00:00 2001 From: Patrick Schlangen Date: Tue, 28 Apr 2026 17:50:05 +0200 Subject: [PATCH 1/3] [ROCm] Clean up AITER FA backend Some of the computed meta data is not used and can be removed. Also, avoid a CPU sync when batch contains decode only. Signed-off-by: Patrick Schlangen --- vllm/v1/attention/backends/rocm_aiter_fa.py | 59 ++++----------------- 1 file changed, 9 insertions(+), 50 deletions(-) diff --git a/vllm/v1/attention/backends/rocm_aiter_fa.py b/vllm/v1/attention/backends/rocm_aiter_fa.py index 01280cd4f48b..e5f9e566d970 100644 --- a/vllm/v1/attention/backends/rocm_aiter_fa.py +++ b/vllm/v1/attention/backends/rocm_aiter_fa.py @@ -324,22 +324,17 @@ def reshape_and_cache_shuffle_triton( @dataclass class AiterFlashAttentionDecodeMetadata: max_query_len: int - min_query_len: int - max_seq_len: int - query_start_loc: torch.Tensor @dataclass class AiterFlashAttentionPrefillMetadata: max_query_len: int - min_query_len: int max_seq_len: int query_start_loc: torch.Tensor @dataclass class AiterChunkSlidingWindowMetadata: - swa_seqlens: torch.Tensor swa_cu_seqlens: torch.Tensor swa_seq_starts: torch.Tensor swa_token_to_batch: torch.Tensor @@ -354,9 +349,7 @@ class AiterChunkContextMetadata: cu_seq_lens_chunk: torch.Tensor chunk_starts: torch.Tensor token_to_batch: torch.Tensor - seq_tot: list[int] max_seq_lens: list[int] - seq_lens: torch.Tensor num_chunks: int total_token_per_batch: list[int] swa_metadata: AiterChunkSlidingWindowMetadata | None @@ -365,7 +358,6 @@ class AiterChunkContextMetadata: @dataclass class AiterFlashAttentionChunkPrefillMetadata: max_query_len: int - min_query_len: int max_seq_len: int query_start_loc: torch.Tensor chunk_context_metadata: AiterChunkContextMetadata @@ -382,8 +374,6 @@ class AiterFlashAttentionMetadata: # |-- query_len ---| num_actual_tokens: int # Number of tokens excluding padding. - num_actual_kv_tokens: int - max_query_len: int query_start_loc: torch.Tensor max_seq_len: int seq_lens: torch.Tensor @@ -395,7 +385,6 @@ class AiterFlashAttentionMetadata: num_decodes: int num_decode_tokens: int num_prefills: int - num_prefill_tokens: int num_extends: int num_extend_tokens: int @@ -405,8 +394,6 @@ class AiterFlashAttentionMetadata: # For cascade attention. use_cascade: bool - common_prefix_len: int - total_tokens: int # Only for fp8 shuffle layout kv cache, we allocate kv_scale for each layer # since we might integrate per token quant for kv cache in the future. @@ -441,7 +428,6 @@ def __init__( # Sliding window size to be used with the AOT scheduler will be # populated on first build() call. self.aot_sliding_window: tuple[int, int] | None = None - self.total_tokens: int = 0 self._init_reorder_batch_threshold(1, supports_spec_as_decode=True) sliding_window_configs: set[tuple[int, int] | None] = set() @@ -473,13 +459,9 @@ def __init__( def build_for_cudagraph_capture( self, common_attn_metadata: CommonAttentionMetadata ): - self.total_tokens = ( - self.model_config.max_model_len - * self.vllm_config.scheduler_config.max_num_partial_prefills + return self.build( + common_prefix_len=0, common_attn_metadata=common_attn_metadata ) - res = self.build(common_prefix_len=0, common_attn_metadata=common_attn_metadata) - self.total_tokens = 0 - return res def build( self, @@ -515,12 +497,14 @@ def build( num_prefills, num_decode_tokens, num_extend_tokens, - num_prefill_tokens, + _, ) = split_ret query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu - seq_lens = common_attn_metadata.seq_lens.cpu() + # Only copy seq_lens to CPU when prefill or extend is present to avoid a blocking + # device→host transfer. + seq_lens = common_attn_metadata.seq_lens.cpu() if num_prefills > 0 or num_extends > 0 else None query_lens_cpu = query_start_loc_cpu[1:] - query_start_loc_cpu[:-1] @@ -528,9 +512,6 @@ def build( if num_decodes > 0: decode_metadata = AiterFlashAttentionDecodeMetadata( max_query_len=query_lens_cpu[:num_decodes].max().item(), - min_query_len=query_lens_cpu[:num_decodes].min().item(), - max_seq_len=seq_lens[:num_decodes].max().item(), - query_start_loc=common_attn_metadata.query_start_loc[: num_decodes + 1], ) prefill_metadata = None @@ -541,7 +522,6 @@ def build( ] prefill_metadata = AiterFlashAttentionPrefillMetadata( max_query_len=query_lens_for_prefill.max().item(), - min_query_len=query_lens_for_prefill.min().item(), max_seq_len=seq_lens[num_decodes + num_extends :].max().item(), query_start_loc=query_start_loc_device - query_start_loc_device[0], ) @@ -591,9 +571,6 @@ def build( total_tokens = cu_seq_lens[-1].item() swa_metadata = AiterChunkSlidingWindowMetadata( - swa_seqlens=swa_seqlen_for_extend.to( - self.device, non_blocking=True - ), swa_cu_seqlens=cu_seq_lens.to(self.device, non_blocking=True), swa_seq_starts=seq_starts.to(self.device, non_blocking=True), swa_token_to_batch=token_to_seq.to(self.device, non_blocking=True), @@ -638,10 +615,8 @@ def build( workspace=self.extend_workspace, cu_seq_lens_chunk=cu_seq_lens_cpu.to(self.device, non_blocking=True), chunk_starts=chunk_starts.to(self.device, non_blocking=True), - seq_tot=chunk_seq_lens.sum(dim=1).tolist(), - max_seq_lens=chunk_seq_lens.max(dim=1).values.tolist(), - seq_lens=chunk_seq_lens, token_to_batch=token_to_batch_tensor.to(self.device, non_blocking=True), + max_seq_lens=chunk_seq_lens.max(dim=1).values.tolist(), num_chunks=num_chunks, total_token_per_batch=cu_seq_lens_cpu[:, -1].tolist(), swa_metadata=swa_metadata, @@ -659,20 +634,15 @@ def build( ) extend_metadata = AiterFlashAttentionChunkPrefillMetadata( max_query_len=query_lens_for_extend.max().item(), - min_query_len=query_lens_for_extend.min().item(), max_seq_len=seq_lens[num_extends_slice].max().item(), query_start_loc=query_start_loc_device - query_start_loc_device[0], chunk_context_metadata=chunk_context_metadata, ) - num_actual_kv_tokens = torch.sum(seq_lens).item() - use_cascade = common_prefix_len > 0 attn_metadata = AiterFlashAttentionMetadata( num_actual_tokens=common_attn_metadata.num_actual_tokens, - num_actual_kv_tokens=num_actual_kv_tokens, - max_query_len=common_attn_metadata.max_query_len, query_start_loc=common_attn_metadata.query_start_loc, max_seq_len=common_attn_metadata.max_seq_len, seq_lens=common_attn_metadata.seq_lens, @@ -682,15 +652,12 @@ def build( num_decodes=num_decodes, num_decode_tokens=num_decode_tokens, num_prefills=num_prefills, - num_prefill_tokens=num_prefill_tokens, num_extends=num_extends, num_extend_tokens=num_extend_tokens, decode_metadata=decode_metadata, prefill_metadata=prefill_metadata, extend_metadata=extend_metadata, use_cascade=use_cascade, - common_prefix_len=common_prefix_len, - total_tokens=self.total_tokens, k_scale=self.scale, v_scale=self.scale, ) @@ -713,15 +680,10 @@ def build_for_drafting( decode_metadata = AiterFlashAttentionDecodeMetadata( max_query_len=common_attn_metadata.max_query_len, - min_query_len=common_attn_metadata.max_query_len, # uniform batch - max_seq_len=common_attn_metadata.max_seq_len, - query_start_loc=common_attn_metadata.query_start_loc, ) return AiterFlashAttentionMetadata( num_actual_tokens=num_tokens, - num_actual_kv_tokens=0, # not used in unified_attention path - max_query_len=common_attn_metadata.max_query_len, query_start_loc=common_attn_metadata.query_start_loc, max_seq_len=common_attn_metadata.max_seq_len, seq_lens=common_attn_metadata.seq_lens, @@ -731,15 +693,12 @@ def build_for_drafting( num_decodes=num_reqs, num_decode_tokens=num_tokens, num_prefills=0, - num_prefill_tokens=0, num_extends=0, num_extend_tokens=0, decode_metadata=decode_metadata, prefill_metadata=None, extend_metadata=None, use_cascade=False, - common_prefix_len=0, - total_tokens=self.total_tokens, k_scale=self.scale, v_scale=self.scale, ) @@ -932,8 +891,8 @@ def extend_forward( output: torch.Tensor, cu_seqlens_q: torch.Tensor, max_seqlen_q: int, - max_seqlen_k: int, min_seqlen_q: int, + max_seqlen_k: int, block_table: torch.Tensor, slot_mapping: torch.Tensor, k_scale: torch.Tensor, @@ -1160,8 +1119,8 @@ def forward( output=extend_outputs, cu_seqlens_q=attn_metadata.extend_metadata.query_start_loc, max_seqlen_q=attn_metadata.extend_metadata.max_query_len, - max_seqlen_k=attn_metadata.extend_metadata.max_seq_len, min_seqlen_q=1, + max_seqlen_k=attn_metadata.extend_metadata.max_seq_len, block_table=attn_metadata.block_table[ num_decodes : num_decodes + num_extends ], From 3f3eefbf693612de014b442ff0ecc413ab570b01 Mon Sep 17 00:00:00 2001 From: Patrick Schlangen Date: Fri, 8 May 2026 11:00:41 +0200 Subject: [PATCH 2/3] Fix ruff style issue Signed-off-by: Patrick Schlangen --- vllm/v1/attention/backends/rocm_aiter_fa.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/vllm/v1/attention/backends/rocm_aiter_fa.py b/vllm/v1/attention/backends/rocm_aiter_fa.py index e5f9e566d970..4a3f894f03bf 100644 --- a/vllm/v1/attention/backends/rocm_aiter_fa.py +++ b/vllm/v1/attention/backends/rocm_aiter_fa.py @@ -504,7 +504,11 @@ def build( # Only copy seq_lens to CPU when prefill or extend is present to avoid a blocking # device→host transfer. - seq_lens = common_attn_metadata.seq_lens.cpu() if num_prefills > 0 or num_extends > 0 else None + seq_lens = ( + common_attn_metadata.seq_lens.cpu() + if num_prefills > 0 or num_extends > 0 + else None + ) query_lens_cpu = query_start_loc_cpu[1:] - query_start_loc_cpu[:-1] From 1256b9ad94367933f492e926e4f0c3b8fab01279 Mon Sep 17 00:00:00 2001 From: Patrick Schlangen Date: Fri, 8 May 2026 11:34:12 +0200 Subject: [PATCH 3/3] Fix mypy/ruff checks Signed-off-by: Patrick Schlangen --- vllm/v1/attention/backends/rocm_aiter_fa.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/vllm/v1/attention/backends/rocm_aiter_fa.py b/vllm/v1/attention/backends/rocm_aiter_fa.py index 4a3f894f03bf..5dbedc86bc02 100644 --- a/vllm/v1/attention/backends/rocm_aiter_fa.py +++ b/vllm/v1/attention/backends/rocm_aiter_fa.py @@ -502,8 +502,8 @@ def build( query_start_loc_cpu = common_attn_metadata.query_start_loc_cpu - # Only copy seq_lens to CPU when prefill or extend is present to avoid a blocking - # device→host transfer. + # Only copy seq_lens to CPU when prefill or extend is present to avoid a + # blocking device→host transfer. seq_lens = ( common_attn_metadata.seq_lens.cpu() if num_prefills > 0 or num_extends > 0 @@ -520,6 +520,7 @@ def build( prefill_metadata = None if num_prefills > 0: + assert seq_lens is not None query_lens_for_prefill = query_lens_cpu[num_decodes + num_extends :] query_start_loc_device = common_attn_metadata.query_start_loc[ num_decodes + num_extends : @@ -532,6 +533,7 @@ def build( extend_metadata = None if num_extends > 0: + assert seq_lens is not None num_extends_slice = slice(num_decodes, num_decodes + num_extends) query_lens_for_extend = query_lens_cpu[num_extends_slice] seq_lens_for_extend = seq_lens[num_extends_slice]