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 @@ 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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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

Expand All @@ -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.
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -515,22 +497,21 @@ 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]

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,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],
)
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
)
Expand All @@ -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,
Expand All @@ -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,
)
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
],
Expand Down
Loading