Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 12 additions & 13 deletions vllm/v1/attention/backends/flashinfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -644,6 +644,10 @@ class FlashInferMetadata:
num_prefill_tokens: int
causal: bool

kv_cache_layout: KVCacheLayout
"""From the builder's config: a draft impl's construction-time config is a
derived copy that never sees the layout resolved after loading."""

prefill: FIPrefill | TRTLLMPrefill | None
"""
Holds the metadata for the prefill portion of the batch.
Expand Down Expand Up @@ -1388,6 +1392,7 @@ def build(
num_prefills=num_prefills,
num_prefill_tokens=num_prefill_tokens,
causal=causal,
kv_cache_layout=self.kv_cache_layout,
use_cascade=use_cascade,
prefill=None,
decode=None,
Expand Down Expand Up @@ -1813,8 +1818,6 @@ def __init__(
num_heads, num_kv_heads, is_prefill=False
)
vllm_config = get_current_vllm_config_or_none()
# The layout is resolved after model construction, so read it lazily.
self.cache_config = vllm_config.cache_config if vllm_config else None
# Query pre-quantization needs a single dtype for the whole query tensor.
# SM90 XQA needs BF16/FP16-Q for decode and FP8 for prefill,
# so only enable this for SM100 trtllm-gen where both use FP8-Q.
Expand Down Expand Up @@ -1850,11 +1853,6 @@ def __init__(
else:
self.dcp_combine = partial(cp_lse_ag_out_rs, is_lse_base_on_e=False)

@property
def kv_cache_layout(self) -> KVCacheLayout:
assert self.cache_config is not None
return self.cache_config.get_resolved_kv_cache_layout()

def fused_output_quant_supported(self, quant_key: QuantKey):
# XQA does not support FP8/NVFP4 output, so require trtllm-gen
# (SM100+) here. Without that we cannot fuse the output quant.
Expand Down Expand Up @@ -2020,11 +2018,12 @@ def forward(
value = value[:num_actual_tokens]
output_padded = output
output = output[:num_actual_tokens]
kv_cache_layout = attn_metadata.kv_cache_layout

if attn_metadata.use_cascade:
# Cascade attention (rare case).
assert attn_metadata.cascade_wrapper is not None
stride_order = self.kv_cache_layout.layer_view_order
stride_order = kv_cache_layout.layer_view_order
if self.is_kvcache_nvfp4:
kv_cache_views = tuple(
cache.permute(*stride_order)
Expand All @@ -2044,7 +2043,7 @@ def forward(
num_decode_tokens = attn_metadata.num_decode_tokens
num_prefill_tokens = attn_metadata.num_prefill_tokens

stride_order = self.kv_cache_layout.layer_view_order
stride_order = kv_cache_layout.layer_view_order
kv_cache_permute = kv_cache.permute(*stride_order) # HND and contiguous
# Fix degenerate strides on any size-1 dimension (e.g. num_kv_heads=1
# with TP=8). PyTorch permits non-canonical strides on size-1 dims;
Expand Down Expand Up @@ -2200,7 +2199,7 @@ def forward(
seq_lens_prefill = attn_metadata.prefill.seq_lens

# This path needs to be enabled with VLLM_KV_CACHE_LAYOUT = HND
assert get_flashinfer_layout_string(self.kv_cache_layout) == "HND"
assert get_flashinfer_layout_string(kv_cache_layout) == "HND"
assert is_strictly_contiguous(prefill_query)
assert is_strictly_contiguous(workspace_buffer)
assert is_strictly_contiguous(block_tables_prefill)
Expand Down Expand Up @@ -2385,7 +2384,7 @@ def forward(
# trtllm-gen needs HND layout on SM100. XQA is selected
# separately on SM90 and does not use this SM100 layout gate.
if decode_with_trtllm_gen:
assert get_flashinfer_layout_string(self.kv_cache_layout) == "HND"
assert get_flashinfer_layout_string(kv_cache_layout) == "HND"
else:
assert decode_with_xqa
assert is_strictly_contiguous(decode_query)
Expand Down Expand Up @@ -2431,7 +2430,7 @@ def forward(
window_left=self.window_left,
out=output[:num_decode_tokens],
sinks=self.sinks,
kv_layout=get_flashinfer_layout_string(self.kv_cache_layout),
kv_layout=get_flashinfer_layout_string(kv_cache_layout),
q_len_per_req=q_len_per_req,
mask=attn_metadata.decode.mask,
q_cu_seq_lens=attn_metadata.decode.q_cu_seq_lens,
Expand Down Expand Up @@ -2510,7 +2509,7 @@ def forward(
sinks=self.sinks,
o_sf_scale=self.o_sf_scale,
out=out,
kv_layout=get_flashinfer_layout_string(self.kv_cache_layout),
kv_layout=get_flashinfer_layout_string(kv_cache_layout),
backend=attn_metadata.decode.kernel.value,
q_len_per_req=q_len_per_req,
max_q_len=max_q_len,
Expand Down
Loading