From 44aac869e52d6072b82a204a18adb49ae6af477c Mon Sep 17 00:00:00 2001 From: Patrik Torstensson Date: Fri, 4 Sep 2026 09:09:25 +0100 Subject: [PATCH] [Bugfix] FlashInfer: read the KV cache layout from attention metadata FlashInferImpl captured cache_config from the current vllm config at construction and read kv_cache_layout from it lazily, since the engine core resolves the layout after model load (#51718) and records it on the worker's CacheConfig. A draft model with a kv_cache_dtype override is built under a replace()d config whose CacheConfig is a copy, so its impls never see the resolution and the first draft forward, during memory profiling, raises "KV cache layout has not been resolved yet". Put the resolved layout on FlashInferMetadata. The builder reads it from the worker's config at build time and every impl-side read happens in forward() with the metadata in hand, so the impl needs no config of its own. This covers the DFlash, EAGLE and DSpark drafters alike, which all derive their config the same way. Co-authored-by: Claude Fable 5.1 Signed-off-by: Patrik Torstensson --- vllm/v1/attention/backends/flashinfer.py | 25 ++++++++++++------------ 1 file changed, 12 insertions(+), 13 deletions(-) diff --git a/vllm/v1/attention/backends/flashinfer.py b/vllm/v1/attention/backends/flashinfer.py index 07c19ea6b9f8..7d35ac377260 100755 --- a/vllm/v1/attention/backends/flashinfer.py +++ b/vllm/v1/attention/backends/flashinfer.py @@ -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. @@ -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, @@ -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. @@ -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. @@ -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) @@ -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; @@ -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) @@ -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) @@ -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, @@ -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,