Skip to content

Commit 74031c7

Browse files
update attention backend
1 parent a1ab318 commit 74031c7

1 file changed

Lines changed: 113 additions & 20 deletions

File tree

vllm_fl/dispatch/backends/vendor/gcu/impl/attention.py

Lines changed: 113 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -49,6 +49,14 @@
4949
from 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+
5260
from 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

195205
class 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

Comments
 (0)