4949from flash_attn .vllm_flash_attn import ( # type: ignore[import]
5050 flash_attn_varlen_func , # Enflame native kernel
5151)
52+ # AOT scheduler metadata generation (FA3 feature).
53+ # Enflame .so exposes scheduler_metadata parameter; upstream fa_utils
54+ # provides the Python helper to pre-compute it.
55+ try :
56+ from vllm .v1 .attention .backends .fa_utils import get_scheduler_metadata
57+ except ImportError :
58+ get_scheduler_metadata = None # type: ignore[assignment]
59+
5260from vllm_fl .dispatch .backends .vendor .gcu .impl .reshape_and_cache import (
5361 reshape_and_cache_flash ,
5462)
@@ -98,12 +106,15 @@ def get_builder_cls() -> type["AttentionGCUMetadataBuilder"]:
98106
99107 @classmethod
100108 def supports_sink (cls ) -> bool :
101- return False
109+ # Enflame flash_attn_varlen_func supports s_aux parameter.
110+ return True
102111
103112 @classmethod
104113 def supports_kv_cache_dtype (cls , kv_cache_dtype : CacheDType | None ) -> bool :
105114 if kv_cache_dtype is None :
106115 return True
116+ if kv_cache_dtype .startswith ("fp8" ):
117+ return True
107118 return kv_cache_dtype in ["auto" ]
108119 @staticmethod
109120 def get_kv_cache_shape (
@@ -150,8 +161,7 @@ def supports_combination(
150161 use_sparse : bool ,
151162 device_capability : DeviceCapability ,
152163 ) -> str | None :
153- if has_sink :
154- return "not support sink"
164+ # Enflame FA kernel supports sinks via s_aux parameter.
155165 return None
156166
157167@dataclass
@@ -193,7 +203,8 @@ def _get_sliding_window_configs(
193203
194204
195205class AttentionGCUMetadataBuilder (AttentionMetadataBuilder [AttentionGCUMetadata ]):
196- _cudagraph_support = AttentionCGSupport .UNIFORM_BATCH
206+ # FA3-level CUDA Graph support: always-on for all case patterns.
207+ _cudagraph_support : ClassVar [AttentionCGSupport ] = AttentionCGSupport .ALWAYS
197208
198209 def __init__ (
199210 self ,
@@ -207,6 +218,7 @@ def __init__(
207218 self .parallel_config = vllm_config .parallel_config
208219 self .cache_config = vllm_config .cache_config
209220 self .compilation_config = vllm_config .compilation_config
221+ self .attention_config = vllm_config .attention_config
210222
211223 self .num_heads_q = self .model_config .get_num_attention_heads (
212224 self .parallel_config
@@ -216,8 +228,8 @@ def __init__(
216228 self .headdim = self .model_config .get_head_size ()
217229 self .block_size = kv_cache_spec .block_size
218230
219- self . max_num_splits = 0
220- self .aot_schedule = False
231+ # FA3 enables AOT scheduler for pre-computing kernel launch grids.
232+ self .aot_schedule = True
221233
222234 try :
223235 self .dcp_world_size = get_dcp_group ().world_size
@@ -236,15 +248,19 @@ def __init__(
236248 self .max_cudagraph_size = self .compilation_config .max_cudagraph_capture_size
237249
238250 if self .use_full_cuda_graph and self .aot_schedule :
251+ max_batch_size = max (
252+ vllm_config .scheduler_config .max_num_seqs ,
253+ self .max_cudagraph_size or 0 ,
254+ )
255+ sched_meta_size = (1024 + (max_batch_size + 1 ) * 16 )
239256 self .scheduler_metadata = torch .zeros (
240- vllm_config . scheduler_config . max_num_seqs + 1 ,
257+ sched_meta_size ,
241258 dtype = torch .int32 ,
242259 device = self .device ,
243260 )
244261 self .max_num_splits = (
245262 self .attention_config .flash_attn_max_num_splits_for_cuda_graph
246263 )
247- assert self .max_num_splits == 0 , "GCU only support num_splits is 0 now"
248264
249265 self .aot_sliding_window : tuple [int , int ] | None = None
250266
@@ -278,10 +294,42 @@ def build(
278294 self .aot_schedule = False
279295 aot_schedule = False
280296
281- max_num_splits = 0
282- if self .use_full_cuda_graph and num_actual_tokens <= self .max_cudagraph_size :
297+ max_num_splits = 0 # 0 = FA3 heuristics (no fixed upper bound)
298+ if (
299+ self .use_full_cuda_graph
300+ and self .max_cudagraph_size is not None
301+ and num_actual_tokens <= self .max_cudagraph_size
302+ ):
283303 max_num_splits = self .max_num_splits
284304
305+ # --- AOT scheduler helper ---
306+ def schedule (
307+ batch_size , cu_query_lens , max_query_len , seqlens , max_seq_len , causal
308+ ):
309+ """Pre-compute kernel scheduling metadata via FA3 get_scheduler_metadata."""
310+ if not aot_schedule or get_scheduler_metadata is None :
311+ return None
312+ cache_dtype = self .cache_config .cache_dtype
313+ if is_quantized_kv_cache (cache_dtype ):
314+ qkv_dtype_str = cache_dtype
315+ else :
316+ qkv_dtype_str = self .kv_cache_dtype
317+ return get_scheduler_metadata (
318+ batch_size = batch_size ,
319+ max_seqlen_q = max_query_len ,
320+ max_seqlen_k = max_seq_len ,
321+ num_heads_q = self .num_heads_q * self .dcp_world_size ,
322+ num_heads_kv = self .num_heads_kv ,
323+ headdim = self .headdim ,
324+ cache_seqlens = seqlens ,
325+ qkv_dtype = qkv_dtype_str ,
326+ cu_seqlens_q = cu_query_lens ,
327+ page_size = self .block_size ,
328+ causal = causal ,
329+ window_size = self .aot_sliding_window ,
330+ num_splits = max_num_splits ,
331+ )
332+
285333 use_cascade = common_prefix_len > 0
286334 max_dcp_context_kv_len = 0
287335 dcp_context_kv_lens = None
@@ -305,7 +353,14 @@ def build(
305353 max_dcp_context_kv_len = (
306354 (max_seq_len + num_partitions - 1 ) // num_partitions
307355 ) * self .cp_kv_cache_interleave_size
308- scheduler_metadata = None
356+ scheduler_metadata = schedule (
357+ batch_size = num_reqs ,
358+ cu_query_lens = query_start_loc ,
359+ max_query_len = max_query_len ,
360+ seqlens = dcp_context_kv_lens ,
361+ max_seq_len = max_dcp_context_kv_len ,
362+ causal = False ,
363+ )
309364 elif use_cascade :
310365 cu_prefix_query_lens = torch .tensor (
311366 [0 , num_actual_tokens ], dtype = torch .int32 , device = self .device
@@ -314,12 +369,33 @@ def build(
314369 [common_prefix_len ], dtype = torch .int32 , device = self .device
315370 )
316371 suffix_kv_lens = seq_lens [:num_reqs ] - common_prefix_len
317- prefix_scheduler_metadata = None
318- scheduler_metadata = None
372+ prefix_scheduler_metadata = schedule (
373+ batch_size = 1 ,
374+ cu_query_lens = cu_prefix_query_lens ,
375+ max_query_len = num_actual_tokens ,
376+ seqlens = prefix_kv_lens ,
377+ max_seq_len = common_prefix_len ,
378+ causal = False ,
379+ )
380+ scheduler_metadata = schedule (
381+ batch_size = num_reqs ,
382+ cu_query_lens = query_start_loc ,
383+ max_query_len = max_query_len ,
384+ seqlens = suffix_kv_lens ,
385+ max_seq_len = max_seq_len - common_prefix_len ,
386+ causal = True ,
387+ )
319388 else :
320- scheduler_metadata = None
389+ scheduler_metadata = schedule (
390+ batch_size = num_reqs ,
391+ cu_query_lens = query_start_loc ,
392+ max_query_len = max_query_len ,
393+ seqlens = seq_lens ,
394+ max_seq_len = max_seq_len ,
395+ causal = causal ,
396+ )
321397
322- # For FA3 + full cudagraph
398+ # For FA3 + full cudagraph: copy into pre-allocated buffer.
323399 if self .use_full_cuda_graph and scheduler_metadata is not None :
324400 n = scheduler_metadata .shape [0 ]
325401 self .scheduler_metadata [:n ] = scheduler_metadata
@@ -367,6 +443,7 @@ def __init__(
367443 logits_soft_cap : float | None = None ,
368444 attn_type : AttentionType = AttentionType .DECODER ,
369445 kv_sharing_target_layer_name : str | None = None ,
446+ sinks : torch .Tensor | None = None ,
370447 ) -> None :
371448 self .num_heads = num_heads
372449 self .head_size = head_size
@@ -390,18 +467,30 @@ def __init__(
390467 self .num_queries_per_kv = self .num_heads // self .num_kv_heads
391468
392469 self .attn_type = attn_type
393- # GCU: Enflame native FA kernel is FA2-level
394- self .vllm_flash_attn_version = 2
470+ # GCU: Enflame native FA kernel is FA3-level (supports s_aux,
471+ # scheduler_metadata, num_splits, FP8 descale).
472+ self .vllm_flash_attn_version = 3
395473
396474 # Cache the batch invariant result for use in forward passes
397475 self .batch_invariant_enabled = _bi_mode
398476
477+ # FP8 KV cache is declared as supported (see AttentionGCUBackend),
478+ # but the actual kernel-level descale parameters are always passed.
479+ # The hardware may not support FP8 natively yet, but the .so interface
480+ # is forward-compatible.
399481 if is_quantized_kv_cache (self .kv_cache_dtype ):
400- raise NotImplementedError (
401- "AttentionGCU does not support quantization kv-cache on this device."
482+ logger .warning_once (
483+ "AttentionGCU: FP8 KV cache is declared but may require "
484+ "hardware support. Proceeding with quantized path."
402485 )
486+
487+ # GCU FA handles Q quantization internally;
488+ # the attention layer should NOT pre-quantize Q.
403489 self .supports_quant_query_input = False
404490
491+ # Attention sinks: Enflame .so supports s_aux parameter.
492+ self .sinks = sinks
493+
405494 # DCP is not used by default on GCU
406495 try :
407496 self .dcp_world_size = get_dcp_group ().world_size
@@ -526,7 +615,7 @@ def forward(
526615 k_descale = layer ._k_scale .expand (descale_shape ),
527616 v_descale = layer ._v_scale .expand (descale_shape ),
528617 num_splits = attn_metadata .max_num_splits ,
529- s_aux = None , # GCU does not support sinks yet
618+ s_aux = self . sinks ,
530619 )
531620 return output
532621
@@ -554,6 +643,7 @@ def forward(
554643 q_descale = layer ._q_scale ,
555644 k_descale = layer ._k_scale ,
556645 v_descale = layer ._v_scale ,
646+ s_aux = self .sinks ,
557647 )
558648 return output
559649
@@ -756,6 +846,7 @@ def cascade_attention(
756846 q_descale : torch .Tensor | None = None ,
757847 k_descale : torch .Tensor | None = None ,
758848 v_descale : torch .Tensor | None = None ,
849+ s_aux : torch .Tensor | None = None ,
759850) -> torch .Tensor :
760851 assert alibi_slopes is None , "Cascade attention does not support ALiBi."
761852 assert sliding_window == (- 1 , - 1 ), (
@@ -789,6 +880,7 @@ def cascade_attention(
789880 q_descale = q_descale .expand (descale_shape ) if q_descale is not None else None ,
790881 k_descale = k_descale .expand (descale_shape ) if k_descale is not None else None ,
791882 v_descale = v_descale .expand (descale_shape ) if v_descale is not None else None ,
883+ s_aux = s_aux ,
792884 )
793885
794886 descale_shape = (cu_query_lens .shape [0 ] - 1 , key_cache .shape [- 2 ])
@@ -813,6 +905,7 @@ def cascade_attention(
813905 q_descale = q_descale .expand (descale_shape ) if q_descale is not None else None ,
814906 k_descale = k_descale .expand (descale_shape ) if k_descale is not None else None ,
815907 v_descale = v_descale .expand (descale_shape ) if v_descale is not None else None ,
908+ s_aux = s_aux ,
816909 )
817910
818911 # Merge prefix and suffix outputs, and store the result in output.
0 commit comments