Skip to content
59 changes: 9 additions & 50 deletions vllm/v1/attention/backends/rocm_aiter_fa.py
Original file line number Diff line number Diff line change
Expand Up @@ -324,22 +324,17 @@
@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
Expand All @@ -354,9 +349,7 @@
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
Expand All @@ -365,7 +358,6 @@
@dataclass
class AiterFlashAttentionChunkPrefillMetadata:
max_query_len: int
min_query_len: int
max_seq_len: int
query_start_loc: torch.Tensor
chunk_context_metadata: AiterChunkContextMetadata
Expand All @@ -382,8 +374,6 @@
# |-- 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
Expand All @@ -395,7 +385,6 @@
num_decodes: int
num_decode_tokens: int
num_prefills: int
num_prefill_tokens: int
num_extends: int
num_extend_tokens: int

Expand All @@ -405,8 +394,6 @@

# 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.
Expand Down Expand Up @@ -441,7 +428,6 @@
# 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()
Expand Down Expand Up @@ -473,13 +459,9 @@
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,
Expand Down Expand Up @@ -515,22 +497,21 @@
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

Check failure on line 505 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E501)

vllm/v1/attention/backends/rocm_aiter_fa.py:505:89: E501 Line too long (89 > 88)
# device→host transfer.
seq_lens = common_attn_metadata.seq_lens.cpu() if num_prefills > 0 or num_extends > 0 else None

Check failure on line 507 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Ruff (E501)

vllm/v1/attention/backends/rocm_aiter_fa.py:507:89: E501 Line too long (103 > 88)

query_lens_cpu = query_start_loc_cpu[1:] - query_start_loc_cpu[:-1]

decode_metadata = None
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
Expand All @@ -541,12 +522,11 @@
]
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],
)

extend_metadata = None

Check failure on line 529 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]

Check failure on line 529 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]

Check failure on line 529 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]
if num_extends > 0:
num_extends_slice = slice(num_decodes, num_decodes + num_extends)
query_lens_for_extend = query_lens_cpu[num_extends_slice]
Expand All @@ -554,7 +534,7 @@
computed_kv_lens = seq_lens_for_extend - query_lens_for_extend
swa_metadata = None
if self.aot_sliding_window is not None:
swa_seqlen_for_extend = torch.minimum(

Check failure on line 537 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]

Check failure on line 537 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]

Check failure on line 537 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]
seq_lens_for_extend,
query_lens_for_extend + self.aot_sliding_window[0] + 1,
)
Expand Down Expand Up @@ -591,9 +571,6 @@
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),
Expand Down Expand Up @@ -638,10 +615,8 @@
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,
Expand All @@ -659,20 +634,15 @@
)
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,
)

Check failure on line 641 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]

Check failure on line 641 in vllm/v1/attention/backends/rocm_aiter_fa.py

View workflow job for this annotation

GitHub Actions / pre-commit

Value of type "Any | None" is not indexable [index]
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,
Expand All @@ -682,15 +652,12 @@
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,
)
Expand All @@ -713,15 +680,10 @@

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,
Expand All @@ -731,15 +693,12 @@
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,
)
Expand Down Expand Up @@ -932,8 +891,8 @@
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,
Expand Down Expand Up @@ -1160,8 +1119,8 @@
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
],
Expand Down
Loading