diff --git a/vllm_fl/dispatch/backends/vendor/gcu/gcu.py b/vllm_fl/dispatch/backends/vendor/gcu/gcu.py index de7d51efd..0acceff98 100644 --- a/vllm_fl/dispatch/backends/vendor/gcu/gcu.py +++ b/vllm_fl/dispatch/backends/vendor/gcu/gcu.py @@ -73,15 +73,11 @@ def rotary_embedding( ) def attention_backend(self, use_mla: bool = False, use_sparse: bool = False) -> str: - from vllm.v1.attention.backends.registry import AttentionBackendEnum - if use_mla: if use_sparse: raise NotImplementedError("GCU does not support sparse attention yet") raise NotImplementedError("GCU does not support MLA yet") - - import flash_attn.vllm_flash_attn - - sys.modules["vllm.vllm_flash_attn"] = flash_attn.vllm_flash_attn - - return AttentionBackendEnum.FLASH_ATTN.get_path() + # GCU uses a standalone flash_attn backend (AttentionGCUBackend) that calls + # Enflame's native flash_attn_varlen_func directly, without depending on + # vllm upstream FlashAttentionBackend / FlashAttentionImpl. + return "vllm_fl.dispatch.backends.vendor.gcu.impl.attention.AttentionGCUBackend" diff --git a/vllm_fl/dispatch/backends/vendor/gcu/impl/attention.py b/vllm_fl/dispatch/backends/vendor/gcu/impl/attention.py new file mode 100644 index 000000000..953dd7154 --- /dev/null +++ b/vllm_fl/dispatch/backends/vendor/gcu/impl/attention.py @@ -0,0 +1,912 @@ +# Copyright (c) 2025 BAAI. All rights reserved. +# Adapted from https://github.com/vllm-project/vllm/blob/v0.11.0/vllm/v1/attention/backends/flash_attn.py +# Below is the original copyright: +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# GCU variant: uses flash_attn.vllm_flash_attn (Enflame native) instead of vllm upstream +# flash_attn backend. Does NOT depend on vllm.v1.attention.backends.flash_attn. + +from dataclasses import dataclass +from typing import ClassVar + +import numpy as np +import torch + +from vllm.v1.attention.backend import ( + AttentionBackend, + AttentionImpl, + AttentionType, + MultipleOf, +) +from vllm.utils.torch_utils import is_quantized_kv_cache +from vllm.model_executor.layers.attention.attention import Attention +from vllm.v1.attention.ops.common import cp_lse_ag_out_rs +from vllm.v1.attention.ops.merge_attn_states import merge_attn_states + + +from vllm.config import VllmConfig, get_current_vllm_config, get_layers_from_vllm_config +from vllm.config.cache import CacheDType +from vllm.distributed.parallel_state import get_dcp_group +from vllm.logger import init_logger +from vllm.model_executor.layers.batch_invariant import _batch_invariant_MODE as _bi_mode +from vllm.utils.math_utils import cdiv +from vllm.v1.attention.backend import ( + AttentionCGSupport, + AttentionMetadataBuilder, +) +from vllm.v1.attention.backends.utils import ( + CommonAttentionMetadata, + get_dcp_local_seq_lens, + get_kv_cache_layout, +) +from vllm.v1.kv_cache_interface import AttentionSpec +from vllm.platforms.interface import DeviceCapability + +# GCU: use Enflame-native flash_attn_varlen_func for best performance, +# plus PyTorch native reshape_and_cache_flash (upstream triton kernel has +# GCU compatibility issues — see reshape_and_cache.py for details). +from flash_attn.vllm_flash_attn import ( # type: ignore[import] + flash_attn_varlen_func, # Enflame native kernel +) +# AOT scheduler metadata generation (FA3 feature). +# Enflame .so exposes scheduler_metadata parameter; upstream fa_utils +# provides the Python helper to pre-compute it. +try: + from vllm.v1.attention.backends.fa_utils import get_scheduler_metadata +except ImportError: + get_scheduler_metadata = None # type: ignore[assignment] + +from vllm_fl.dispatch.backends.vendor.gcu.impl.reshape_and_cache import ( + reshape_and_cache_flash, +) + +logger = init_logger(__name__) + + +class AttentionGCUBackend(AttentionBackend): + accept_output_buffer: bool = True + supported_dtypes: ClassVar[list[torch.dtype]] = [torch.float16, torch.bfloat16] + + @staticmethod + def get_supported_kernel_block_sizes() -> list[int | MultipleOf]: + vllm_config = get_current_vllm_config() + model_config = vllm_config.model_config + cache_config = vllm_config.cache_config + if ( + model_config + and model_config.is_hybrid + and ( + cache_config.mamba_ssm_cache_dtype == "float32" + or cache_config.mamba_cache_dtype == "float32" + ) + ): + return [16, 32, 64] + return [MultipleOf(16)] + + @staticmethod + def get_name() -> str: + return "FLASH_ATTN" + + @classmethod + def supports_attn_type(cls, attn_type: str) -> bool: + return attn_type in ( + AttentionType.DECODER, + AttentionType.ENCODER, + AttentionType.ENCODER_ONLY, + AttentionType.ENCODER_DECODER, + ) + @staticmethod + def get_impl_cls() -> type["AttentionGCUImpl"]: + return AttentionGCUImpl + + @staticmethod + def get_builder_cls() -> type["AttentionGCUMetadataBuilder"]: + return AttentionGCUMetadataBuilder + + @classmethod + def supports_sink(cls) -> bool: + # Enflame flash_attn_varlen_func supports s_aux parameter. + return True + + @classmethod + def supports_kv_cache_dtype(cls, kv_cache_dtype: CacheDType | None) -> bool: + if kv_cache_dtype is None: + return True + if kv_cache_dtype.startswith("fp8"): + return True + return kv_cache_dtype in ["auto"] + @staticmethod + def get_kv_cache_shape( + num_blocks: int, + block_size: int, + num_kv_heads: int, + head_size: int, + cache_dtype_str: str = "auto", + ) -> tuple[int, ...]: + if block_size % 16 != 0: + raise ValueError("Block size must be a multiple of 16.") + return (2, num_blocks, block_size, num_kv_heads, head_size) + + @staticmethod + def get_kv_cache_stride_order( + include_num_layers_dimension: bool = False, + ) -> tuple[int, ...]: + cache_layout = get_kv_cache_layout() + if cache_layout == "NHD" and include_num_layers_dimension: + return (2, 0, 1, 3, 4, 5) + elif cache_layout == "NHD": + stride_order = (0, 1, 2, 3, 4) + elif cache_layout == "HND" and include_num_layers_dimension: + return (2, 4, 0, 1, 3, 5) + elif cache_layout == "HND": + stride_order = (0, 1, 3, 2, 4) + else: + raise ValueError(f"Unknown cache layout format {cache_layout}.") + return stride_order + + @classmethod + def supports_head_size(cls, head_size: int) -> bool: + return head_size % 8 == 0 and head_size <= 256 + + @classmethod + def supports_combination( + cls, + head_size: int, + dtype: torch.dtype, + kv_cache_dtype: CacheDType | None, + block_size: int, + use_mla: bool, + has_sink: bool, + use_sparse: bool, + device_capability: DeviceCapability, + ) -> str | None: + # Enflame FA kernel supports sinks via s_aux parameter. + return None + +@dataclass +class AttentionGCUMetadata: + num_actual_tokens: int + max_query_len: int + query_start_loc: torch.Tensor + max_seq_len: int + seq_lens: torch.Tensor + block_table: torch.Tensor + slot_mapping: torch.Tensor + + use_cascade: bool + common_prefix_len: int + cu_prefix_query_lens: torch.Tensor | None + prefix_kv_lens: torch.Tensor | None + suffix_kv_lens: torch.Tensor | None + + max_dcp_context_kv_len: int | None = None + dcp_context_kv_lens: torch.Tensor | None = None + + scheduler_metadata: torch.Tensor | None = None + prefix_scheduler_metadata: torch.Tensor | None = None + max_num_splits: int = 0 + + causal: bool = True + + +def _get_sliding_window_configs( + vllm_config: VllmConfig, +) -> set[tuple[int, int] | None]: + """Get the set of all sliding window configs used in the model.""" + sliding_window_configs: set[tuple[int, int] | None] = set() + layers = get_layers_from_vllm_config(vllm_config, Attention) + for layer in layers.values(): + assert isinstance(layer.impl, AttentionGCUImpl) + sliding_window_configs.add(layer.impl.sliding_window) + return sliding_window_configs + + +class AttentionGCUMetadataBuilder(AttentionMetadataBuilder[AttentionGCUMetadata]): + # FA3-level CUDA Graph support: always-on for all case patterns. + _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.ALWAYS + + def __init__( + self, + kv_cache_spec: AttentionSpec, + layer_names: list[str], + vllm_config: VllmConfig, + device: torch.device, + ): + super().__init__(kv_cache_spec, layer_names, vllm_config, device) + self.model_config = vllm_config.model_config + self.parallel_config = vllm_config.parallel_config + self.cache_config = vllm_config.cache_config + self.compilation_config = vllm_config.compilation_config + self.attention_config = vllm_config.attention_config + + self.num_heads_q = self.model_config.get_num_attention_heads( + self.parallel_config + ) + self.num_heads_kv = self.model_config.get_num_kv_heads(self.parallel_config) + self.kv_cache_dtype = kv_cache_spec.dtype + self.headdim = self.model_config.get_head_size() + self.block_size = kv_cache_spec.block_size + + # FA3 enables AOT scheduler for pre-computing kernel launch grids. + self.aot_schedule = True + + try: + self.dcp_world_size = get_dcp_group().world_size + self.dcp_rank = get_dcp_group().rank_in_group + except AssertionError: + self.dcp_world_size = 1 + self.dcp_rank = 0 + + self.cp_kv_cache_interleave_size = ( + self.parallel_config.cp_kv_cache_interleave_size + ) + + self.use_full_cuda_graph = ( + self.compilation_config.cudagraph_mode.has_full_cudagraphs() + ) + self.max_cudagraph_size = self.compilation_config.max_cudagraph_capture_size + + if self.use_full_cuda_graph and self.aot_schedule: + max_batch_size = max( + vllm_config.scheduler_config.max_num_seqs, + self.max_cudagraph_size or 0, + ) + sched_meta_size = (1024 + (max_batch_size + 1) * 16) + self.scheduler_metadata = torch.zeros( + sched_meta_size, + dtype=torch.int32, + device=self.device, + ) + self.max_num_splits = ( + self.attention_config.flash_attn_max_num_splits_for_cuda_graph + ) + + self.aot_sliding_window: tuple[int, int] | None = None + + def build( + self, + common_prefix_len: int, + common_attn_metadata: CommonAttentionMetadata, + fast_build: bool = False, + ) -> AttentionGCUMetadata: + num_reqs = common_attn_metadata.num_reqs + num_actual_tokens = common_attn_metadata.num_actual_tokens + max_query_len = common_attn_metadata.max_query_len + max_seq_len = common_attn_metadata.max_seq_len + query_start_loc = common_attn_metadata.query_start_loc + seq_lens = common_attn_metadata.seq_lens + block_table_tensor = common_attn_metadata.block_table_tensor + slot_mapping = common_attn_metadata.slot_mapping + causal = common_attn_metadata.causal + + aot_schedule = self.aot_schedule and not fast_build + + if self.aot_sliding_window is None: + self.aot_sliding_window = (-1, -1) + if aot_schedule: + sliding_window_configs = _get_sliding_window_configs(self.vllm_config) + if len(sliding_window_configs) == 1: + sliding_window_config = sliding_window_configs.pop() + if sliding_window_config is not None: + self.aot_sliding_window = sliding_window_config + elif len(sliding_window_configs) > 1: + self.aot_schedule = False + aot_schedule = False + + max_num_splits = 0 # 0 = FA3 heuristics (no fixed upper bound) + if ( + self.use_full_cuda_graph + and self.max_cudagraph_size is not None + and num_actual_tokens <= self.max_cudagraph_size + ): + max_num_splits = self.max_num_splits + + # --- AOT scheduler helper --- + def schedule( + batch_size, cu_query_lens, max_query_len, seqlens, max_seq_len, causal + ): + """Pre-compute kernel scheduling metadata via FA3 get_scheduler_metadata.""" + if not aot_schedule or get_scheduler_metadata is None: + return None + cache_dtype = self.cache_config.cache_dtype + if is_quantized_kv_cache(cache_dtype): + qkv_dtype_str = cache_dtype + else: + qkv_dtype_str = self.kv_cache_dtype + return get_scheduler_metadata( + batch_size=batch_size, + max_seqlen_q=max_query_len, + max_seqlen_k=max_seq_len, + num_heads_q=self.num_heads_q * self.dcp_world_size, + num_heads_kv=self.num_heads_kv, + headdim=self.headdim, + cache_seqlens=seqlens, + qkv_dtype=qkv_dtype_str, + cu_seqlens_q=cu_query_lens, + page_size=self.block_size, + causal=causal, + window_size=self.aot_sliding_window, + num_splits=max_num_splits, + ) + + use_cascade = common_prefix_len > 0 + max_dcp_context_kv_len = 0 + dcp_context_kv_lens = None + + cu_prefix_query_lens = None + prefix_kv_lens = None + suffix_kv_lens = None + prefix_scheduler_metadata = None + + if self.dcp_world_size > 1: + query_kv_lens = query_start_loc[1:] - query_start_loc[:-1] + dcp_context_kv_lens = seq_lens - query_kv_lens + + dcp_context_kv_lens = get_dcp_local_seq_lens( + dcp_context_kv_lens, + self.dcp_world_size, + self.dcp_rank, + self.cp_kv_cache_interleave_size, + ) + num_partitions = self.dcp_world_size * self.cp_kv_cache_interleave_size + max_dcp_context_kv_len = ( + (max_seq_len + num_partitions - 1) // num_partitions + ) * self.cp_kv_cache_interleave_size + scheduler_metadata = schedule( + batch_size=num_reqs, + cu_query_lens=query_start_loc, + max_query_len=max_query_len, + seqlens=dcp_context_kv_lens, + max_seq_len=max_dcp_context_kv_len, + causal=False, + ) + elif use_cascade: + cu_prefix_query_lens = torch.tensor( + [0, num_actual_tokens], dtype=torch.int32, device=self.device + ) + prefix_kv_lens = torch.tensor( + [common_prefix_len], dtype=torch.int32, device=self.device + ) + suffix_kv_lens = seq_lens[:num_reqs] - common_prefix_len + prefix_scheduler_metadata = schedule( + batch_size=1, + cu_query_lens=cu_prefix_query_lens, + max_query_len=num_actual_tokens, + seqlens=prefix_kv_lens, + max_seq_len=common_prefix_len, + causal=False, + ) + scheduler_metadata = schedule( + batch_size=num_reqs, + cu_query_lens=query_start_loc, + max_query_len=max_query_len, + seqlens=suffix_kv_lens, + max_seq_len=max_seq_len - common_prefix_len, + causal=True, + ) + else: + scheduler_metadata = schedule( + batch_size=num_reqs, + cu_query_lens=query_start_loc, + max_query_len=max_query_len, + seqlens=seq_lens, + max_seq_len=max_seq_len, + causal=causal, + ) + + # For FA3 + full cudagraph: copy into pre-allocated buffer. + if self.use_full_cuda_graph and scheduler_metadata is not None: + n = scheduler_metadata.shape[0] + self.scheduler_metadata[:n] = scheduler_metadata + self.scheduler_metadata[n:] = 0 + scheduler_metadata = self.scheduler_metadata[:n] + + attn_metadata = AttentionGCUMetadata( + num_actual_tokens=num_actual_tokens, + max_query_len=max_query_len, + query_start_loc=query_start_loc, + max_seq_len=max_seq_len, + seq_lens=seq_lens, + block_table=block_table_tensor, + slot_mapping=slot_mapping, + max_dcp_context_kv_len=max_dcp_context_kv_len, + dcp_context_kv_lens=dcp_context_kv_lens, + use_cascade=use_cascade, + common_prefix_len=common_prefix_len, + scheduler_metadata=scheduler_metadata, + cu_prefix_query_lens=cu_prefix_query_lens, + prefix_kv_lens=prefix_kv_lens, + suffix_kv_lens=suffix_kv_lens, + prefix_scheduler_metadata=prefix_scheduler_metadata, + max_num_splits=max_num_splits, + causal=causal, + ) + return attn_metadata + + def use_cascade_attention(self, *args, **kwargs) -> bool: + return use_cascade_attention(*args, **kwargs) + + +class AttentionGCUImpl(AttentionImpl): + can_return_lse_for_decode: bool = True + + def __init__( + self, + num_heads: int, + head_size: int, + scale: float, + num_kv_heads: int, + alibi_slopes: list[float] | None, + sliding_window: int | None, + kv_cache_dtype: str, + logits_soft_cap: float | None = None, + attn_type: AttentionType = AttentionType.DECODER, + kv_sharing_target_layer_name: str | None = None, + sinks: torch.Tensor | None = None, + ) -> None: + self.num_heads = num_heads + self.head_size = head_size + self.scale = float(scale) + self.num_kv_heads = num_kv_heads + if alibi_slopes is not None: + alibi_slopes = torch.tensor(alibi_slopes, dtype=torch.float32) + self.alibi_slopes = alibi_slopes + if sliding_window is None: + self.sliding_window = (-1, -1) + elif attn_type == AttentionType.ENCODER_ONLY: + self.sliding_window = (sliding_window - 1, sliding_window - 1) + else: + self.sliding_window = (sliding_window - 1, 0) + self.kv_cache_dtype = kv_cache_dtype + if logits_soft_cap is None: + logits_soft_cap = 0 + self.logits_soft_cap = logits_soft_cap + self.kv_sharing_target_layer_name = kv_sharing_target_layer_name + + self.num_queries_per_kv = self.num_heads // self.num_kv_heads + + self.attn_type = attn_type + # GCU: Enflame native FA kernel is FA3-level (supports s_aux, + # scheduler_metadata, num_splits, FP8 descale). + self.vllm_flash_attn_version = 3 + + # Cache the batch invariant result for use in forward passes + self.batch_invariant_enabled = _bi_mode + + # FP8 KV cache is declared as supported (see AttentionGCUBackend), + # but the actual kernel-level descale parameters are always passed. + # The hardware may not support FP8 natively yet, but the .so interface + # is forward-compatible. + if is_quantized_kv_cache(self.kv_cache_dtype): + logger.warning_once( + "AttentionGCU: FP8 KV cache is declared but may require " + "hardware support. Proceeding with quantized path." + ) + + # GCU FA handles Q quantization internally; + # the attention layer should NOT pre-quantize Q. + self.supports_quant_query_input = False + + # Attention sinks: Enflame .so supports s_aux parameter. + self.sinks = sinks + + # DCP is not used by default on GCU + try: + self.dcp_world_size = get_dcp_group().world_size + except AssertionError: + self.dcp_world_size = 1 + + def forward( + self, + layer: torch.nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + kv_cache: torch.Tensor, + attn_metadata: AttentionGCUMetadata, + output: torch.Tensor | None = None, + output_scale: torch.Tensor | None = None, + output_block_scale: torch.Tensor | None = None, + ) -> torch.Tensor: + """Forward pass with GCU flash attention. + + Args: + query: shape = [num_tokens, num_heads, head_size] + key: shape = [num_tokens, num_kv_heads, head_size] + value: shape = [num_tokens, num_kv_heads, head_size] + kv_cache: shape = + [2, num_blocks, block_size, num_kv_heads, head_size] + attn_metadata: Metadata for attention. + Returns: + shape = [num_tokens, num_heads * head_size] + NOTE: FP8 quantization, flash-attn expect the size of + {q,k,v}_descale to be (num_sequences, num_kv_heads). + We use torch's .expand() to avoid duplicating values + """ + assert output is not None, "Output tensor must be provided." + + if output_scale is not None or output_block_scale is not None: + raise NotImplementedError( + "fused output quantization is not yet supported for AttentionGCU" + ) + + if attn_metadata is None: + # Profiling run. + return output.fill_(0) + + attn_type = self.attn_type + + num_actual_tokens = attn_metadata.num_actual_tokens + + # Handle encoder attention differently - no KV cache needed + if attn_type in (AttentionType.ENCODER_ONLY, AttentionType.ENCODER): + return self._forward_encoder_attention( + query[:num_actual_tokens], + key[:num_actual_tokens], + value[:num_actual_tokens], + output[:num_actual_tokens], + attn_metadata, + layer, + ) + + # For decoder and cross-attention, use KV cache as before + key_cache, value_cache = kv_cache.unbind(0) + + if ( + self.kv_sharing_target_layer_name is None + and key is not None + and value is not None + ): + reshape_and_cache_flash( + key, + value, + key_cache, + value_cache, + attn_metadata.slot_mapping, + self.kv_cache_dtype, + layer._k_scale, + layer._v_scale, + ) + + if not attn_metadata.use_cascade: + cu_seqlens_q = attn_metadata.query_start_loc + seqused_k = attn_metadata.seq_lens + max_seqlen_q = attn_metadata.max_query_len + max_seqlen_k = attn_metadata.max_seq_len + block_table = attn_metadata.block_table + scheduler_metadata = attn_metadata.scheduler_metadata + + descale_shape = (cu_seqlens_q.shape[0] - 1, self.num_kv_heads) + + if self.dcp_world_size > 1: + self._forward_with_dcp( + query[:num_actual_tokens], + key[:num_actual_tokens], + value[:num_actual_tokens], + key_cache, + value_cache, + output[:num_actual_tokens], + attn_metadata, + q_descale=layer._q_scale.expand(descale_shape), + k_descale=layer._k_scale.expand(descale_shape), + v_descale=layer._v_scale.expand(descale_shape), + ) + return output + else: + flash_attn_varlen_func( + q=query[:num_actual_tokens], + k=key_cache, + v=value_cache, + out=output[:num_actual_tokens], + cu_seqlens_q=cu_seqlens_q, + max_seqlen_q=max_seqlen_q, + seqused_k=seqused_k, + max_seqlen_k=max_seqlen_k, + softmax_scale=self.scale, + causal=attn_metadata.causal, + alibi_slopes=self.alibi_slopes, + window_size=self.sliding_window, + block_table=block_table, + softcap=self.logits_soft_cap, + scheduler_metadata=scheduler_metadata, + fa_version=self.vllm_flash_attn_version, + q_descale=layer._q_scale.expand(descale_shape), + k_descale=layer._k_scale.expand(descale_shape), + v_descale=layer._v_scale.expand(descale_shape), + num_splits=attn_metadata.max_num_splits, + s_aux=self.sinks, + ) + return output + + # Cascade attention (rare case). + cascade_attention( + output[:num_actual_tokens], + query[:num_actual_tokens], + key_cache, + value_cache, + cu_query_lens=attn_metadata.query_start_loc, + max_query_len=attn_metadata.max_query_len, + cu_prefix_query_lens=attn_metadata.cu_prefix_query_lens, + prefix_kv_lens=attn_metadata.prefix_kv_lens, + suffix_kv_lens=attn_metadata.suffix_kv_lens, + max_kv_len=attn_metadata.max_seq_len, + softmax_scale=self.scale, + alibi_slopes=self.alibi_slopes, + sliding_window=self.sliding_window, + logits_soft_cap=self.logits_soft_cap, + block_table=attn_metadata.block_table, + common_prefix_len=attn_metadata.common_prefix_len, + fa_version=self.vllm_flash_attn_version, + prefix_scheduler_metadata=attn_metadata.prefix_scheduler_metadata, + suffix_scheduler_metadata=attn_metadata.scheduler_metadata, + q_descale=layer._q_scale, + k_descale=layer._k_scale, + v_descale=layer._v_scale, + s_aux=self.sinks, + ) + return output + + def _forward_with_dcp( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + key_cache: torch.Tensor, + value_cache: torch.Tensor, + output: torch.Tensor, + attn_metadata: AttentionGCUMetadata, + q_descale: torch.Tensor | None = None, + k_descale: torch.Tensor | None = None, + v_descale: torch.Tensor | None = None, + ) -> torch.Tensor: + cu_seqlens_q = attn_metadata.query_start_loc + max_seqlen_q = attn_metadata.max_query_len + block_table = attn_metadata.block_table + + query = query.contiguous() + query_across_dcp = get_dcp_group().all_gather(query, dim=1) + context_attn_out, context_lse = flash_attn_varlen_func( + q=query_across_dcp, + k=key_cache, + v=value_cache, + out=None, + cu_seqlens_q=cu_seqlens_q, + max_seqlen_q=max_seqlen_q, + seqused_k=attn_metadata.dcp_context_kv_lens, + max_seqlen_k=attn_metadata.max_dcp_context_kv_len, + softmax_scale=self.scale, + causal=False, + alibi_slopes=self.alibi_slopes, + window_size=self.sliding_window, + block_table=block_table, + softcap=self.logits_soft_cap, + return_softmax_lse=True, + scheduler_metadata=attn_metadata.scheduler_metadata, + fa_version=self.vllm_flash_attn_version, + q_descale=q_descale, + k_descale=k_descale, + v_descale=v_descale, + ) + context_attn_out_cor, context_lse_cor = cp_lse_ag_out_rs( + context_attn_out, + context_lse.transpose(0, 1), + get_dcp_group(), + return_lse=True, + ) + context_lse_cor = context_lse_cor.transpose(0, 1).contiguous() + + query_attn_out, query_lse = flash_attn_varlen_func( + q=query, + k=key, + v=value, + out=None, + cu_seqlens_q=cu_seqlens_q, + max_seqlen_q=max_seqlen_q, + cu_seqlens_k=cu_seqlens_q, + max_seqlen_k=max_seqlen_q, + softmax_scale=self.scale, + causal=attn_metadata.causal, + alibi_slopes=self.alibi_slopes, + window_size=self.sliding_window, + softcap=self.logits_soft_cap, + return_softmax_lse=True, + fa_version=self.vllm_flash_attn_version, + q_descale=q_descale, + k_descale=k_descale, + v_descale=v_descale, + ) + assert context_attn_out_cor.shape == query_attn_out.shape + assert context_lse_cor.shape == query_lse.shape + merge_attn_states( + output, + context_attn_out_cor, + context_lse_cor, + query_attn_out, + query_lse, + ) + + def _forward_encoder_attention( + self, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + output: torch.Tensor, + attn_metadata: AttentionGCUMetadata, + layer: torch.nn.Module, + ) -> torch.Tensor: + """Forward pass for encoder attention without KV cache.""" + if self.kv_cache_dtype.startswith("fp8"): + raise NotImplementedError( + "quantization is not supported for encoder attention" + ) + + cu_seqlens_q = attn_metadata.query_start_loc + cu_seqlens_k = attn_metadata.query_start_loc + max_seqlen_q = attn_metadata.max_query_len + max_seqlen_k = attn_metadata.max_query_len + + descale_shape = ( + cu_seqlens_q.shape[0] - 1, + self.num_kv_heads, + ) + + flash_attn_varlen_func( + q=query, + k=key, + v=value, + out=output, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlen_q=max_seqlen_q, + max_seqlen_k=max_seqlen_k, + softmax_scale=self.scale, + causal=False, + alibi_slopes=self.alibi_slopes, + window_size=self.sliding_window, + softcap=self.logits_soft_cap, + fa_version=self.vllm_flash_attn_version, + q_descale=layer._q_scale.expand(descale_shape), + k_descale=layer._k_scale.expand(descale_shape), + v_descale=layer._v_scale.expand(descale_shape), + ) + + return output + + +def use_cascade_attention( + common_prefix_len: int, + query_lens: np.ndarray, + num_query_heads: int, + num_kv_heads: int, + use_alibi: bool, + use_sliding_window: bool, + use_local_attention: bool, + num_sms: int, + dcp_world_size: int, +) -> bool: + if common_prefix_len < 256: + return False + if use_alibi or use_sliding_window or use_local_attention: + return False + num_reqs = len(query_lens) + if num_reqs < 8: + return False + if dcp_world_size > 1: + return False + + num_queries_per_kv = num_query_heads // num_kv_heads + use_flash_decoding = ( + num_queries_per_kv > 1 + and not use_sliding_window + and not use_alibi + and np.all(query_lens == 1) + ) + if not use_flash_decoding: + return True + + num_tokens = num_reqs + q_tile_size = 128 + kv_tile_size = 128 + num_prefix_tiles = cdiv(common_prefix_len, kv_tile_size) + + cascade_ctas = num_query_heads * cdiv(num_tokens, q_tile_size) + cascade_waves = cdiv(cascade_ctas, num_sms) + cascade_time = cascade_waves * num_prefix_tiles + + flash_decoding_ctas = ( + num_reqs * num_kv_heads * cdiv(num_queries_per_kv, q_tile_size) + ) + flash_decoding_ctas *= num_prefix_tiles + flash_decoding_time = cdiv(flash_decoding_ctas, num_sms) + + return cascade_time < flash_decoding_time + + +def cascade_attention( + output: torch.Tensor, + query: torch.Tensor, + key_cache: torch.Tensor, + value_cache: torch.Tensor, + cu_query_lens: torch.Tensor, + max_query_len: int, + cu_prefix_query_lens: torch.Tensor, + prefix_kv_lens: torch.Tensor, + suffix_kv_lens: torch.Tensor, + max_kv_len: int, + softmax_scale: float, + alibi_slopes: torch.Tensor | None, + sliding_window: tuple[int, int], + logits_soft_cap: float, + block_table: torch.Tensor, + common_prefix_len: int, + fa_version: int, + prefix_scheduler_metadata: torch.Tensor | None = None, + suffix_scheduler_metadata: torch.Tensor | None = None, + q_descale: torch.Tensor | None = None, + k_descale: torch.Tensor | None = None, + v_descale: torch.Tensor | None = None, + s_aux: torch.Tensor | None = None, +) -> torch.Tensor: + assert alibi_slopes is None, "Cascade attention does not support ALiBi." + assert sliding_window == (-1, -1), ( + "Cascade attention does not support sliding window." + ) + + num_tokens = query.shape[0] + block_size = key_cache.shape[-3] + assert common_prefix_len % block_size == 0 + num_common_kv_blocks = common_prefix_len // block_size + assert num_common_kv_blocks > 0 + descale_shape = (cu_prefix_query_lens.shape[0] - 1, key_cache.shape[-2]) + + # Process shared prefix. + prefix_output, prefix_lse = flash_attn_varlen_func( + q=query, + k=key_cache, + v=value_cache, + cu_seqlens_q=cu_prefix_query_lens, + seqused_k=prefix_kv_lens, + max_seqlen_q=num_tokens, + max_seqlen_k=common_prefix_len, + softmax_scale=softmax_scale, + causal=False, + window_size=sliding_window, + block_table=block_table[:1], + softcap=logits_soft_cap, + return_softmax_lse=True, + scheduler_metadata=prefix_scheduler_metadata, + fa_version=fa_version, + q_descale=q_descale.expand(descale_shape) if q_descale is not None else None, + k_descale=k_descale.expand(descale_shape) if k_descale is not None else None, + v_descale=v_descale.expand(descale_shape) if v_descale is not None else None, + s_aux=s_aux, + ) + + descale_shape = (cu_query_lens.shape[0] - 1, key_cache.shape[-2]) + + # Process suffix per query. + suffix_output, suffix_lse = flash_attn_varlen_func( + q=query, + k=key_cache, + v=value_cache, + cu_seqlens_q=cu_query_lens, + seqused_k=suffix_kv_lens, + max_seqlen_q=max_query_len, + max_seqlen_k=max_kv_len - common_prefix_len, + softmax_scale=softmax_scale, + causal=True, + window_size=sliding_window, + block_table=block_table[:, num_common_kv_blocks:], + softcap=logits_soft_cap, + return_softmax_lse=True, + scheduler_metadata=suffix_scheduler_metadata, + fa_version=fa_version, + q_descale=q_descale.expand(descale_shape) if q_descale is not None else None, + k_descale=k_descale.expand(descale_shape) if k_descale is not None else None, + v_descale=v_descale.expand(descale_shape) if v_descale is not None else None, + s_aux=s_aux, + ) + + # Merge prefix and suffix outputs, and store the result in output. + merge_attn_states(output, prefix_output, prefix_lse, suffix_output, suffix_lse) diff --git a/vllm_fl/dispatch/backends/vendor/gcu/impl/reshape_and_cache.py b/vllm_fl/dispatch/backends/vendor/gcu/impl/reshape_and_cache.py new file mode 100644 index 000000000..cee71e0e5 --- /dev/null +++ b/vllm_fl/dispatch/backends/vendor/gcu/impl/reshape_and_cache.py @@ -0,0 +1,64 @@ +# GCU300 reshape_and_cache_flash: PyTorch native implementation. +# +# The upstream vLLM triton_reshape_and_cache_flash kernel has GCU +# compatibility issues: +# 1. torch.cuda.get_device_capability() crashes (not CUDA) +# 2. tl.load(slot_mapping).to(tl.int64) — int64 in triton kernel +# requires TORCH_GCU_ENABLE_INT64_AND_UINT64=1 env var, and even +# then the int64 division/modulo inside the kernel may not work +# correctly on GCU300. +# +# The PyTorch native implementation below is correct, portable, and +# CUDA Graph compatible (no data-dependent control flow, no dynamic +# shapes, no CPU-GPU synchronization). +# +# When flag_gems is enabled (USE_FLAGGEMS=1), ATen operators (clamp, +# div, remainder, index, etc.) are replaced by Triton kernels that +# support int64 with ENABLE_I64_CHECK=0 env var. + +from __future__ import annotations + +import logging + +import torch + +logger = logging.getLogger(__name__) + + +def reshape_and_cache_flash( + key: torch.Tensor, # [num_tokens, num_kv_heads, head_size] + value: torch.Tensor, # [num_tokens, num_kv_heads, head_size] + key_cache: torch.Tensor, # [num_blocks, block_size, num_kv_heads, head_size] + value_cache: torch.Tensor, # [num_blocks, block_size, num_kv_heads, head_size] + slot_mapping: torch.Tensor, # [num_tokens] int64 + kv_cache_dtype: str, # e.g. "auto", "fp8" (unused for non-quantized) + k_scale: torch.Tensor, # scalar or per-token scale (unused for non-quantized) + v_scale: torch.Tensor, # scalar or per-token scale (unused for non-quantized) +) -> None: + """Write per-token K/V into the paged KV cache (PyTorch native). + + Writes each token's key/value tensor into the paged KV cache at the + flat slot position given by *slot_mapping*. Tokens whose slot is + negative (padding / speculative-draft rejects) are redirected to + block 0 (NULL_BLOCK_ID, reserved for padding in vLLM), which is + a harmless no-op write. + + CUDA Graph compatible: + - No data-dependent control flow (no if/return based on tensor values) + - No dynamic shapes (no boolean masking) + - No CPU-GPU synchronization (no .item()/.cpu()/.tolist()) + """ + block_size = key_cache.size(1) + + # Clamp negative slots to 0 (NULL_BLOCK_ID). Block 0 is reserved + # for padding in vLLM, so writing padding K/V there is harmless. + # This avoids boolean indexing (dynamic shape) and CPU-GPU sync. + slots = torch.clamp(slot_mapping, min=0) + + # Decompose flat slot → (block_idx, offset_within_block) + block_idx = slots // block_size + block_offset = slots % block_size + + # In-place write into paged cache via advanced indexing + key_cache[block_idx, block_offset] = key + value_cache[block_idx, block_offset] = value diff --git a/vllm_fl/dispatch/config/gcu.yaml b/vllm_fl/dispatch/config/enflame.yaml similarity index 100% rename from vllm_fl/dispatch/config/gcu.yaml rename to vllm_fl/dispatch/config/enflame.yaml