From 74e23b9085d1b609ffb019eac3aad06428a830dc Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Fri, 29 Aug 2025 03:30:01 -0400 Subject: [PATCH 01/28] feat: Unify ViT attention backend selection across vision models - Refactor 6 ViT attention classes to use unified backend selection - Add support for Flash Attention, xFormers, ROCm AITer FA, and PyTorch SDPA - Implement get_vit_attn_backend() for automatic hardware-aware backend selection - Maintain model-specific features (QK normalization, dummy heads, etc.) Modified models: - Idefics2VisionAttention: Complete backend unification - InternSdpaAttention (both intern_vit.py and interns1_vit.py): Added unified backend selection - MllamaVisionSdpaAttention: Replaced fixed SDPA with dynamic backend selection - PixtralHFAttention: Migrated from USE_XFORMERS_OPS to unified backend selection - Step3VisionAttention: Added complete backend support Addresses GitHub issue #23880 for ViT attention performance optimization. --- .../models/idefics2_vision_model.py | 48 +++++++++++++++- vllm/model_executor/models/intern_vit.py | 43 ++++++++++++-- vllm/model_executor/models/interns1_vit.py | 43 ++++++++++++-- vllm/model_executor/models/mllama.py | 57 +++++++++++++++---- vllm/model_executor/models/pixtral.py | 33 ++++++++++- vllm/model_executor/models/step3_vl.py | 52 +++++++++++++---- 6 files changed, 238 insertions(+), 38 deletions(-) diff --git a/vllm/model_executor/models/idefics2_vision_model.py b/vllm/model_executor/models/idefics2_vision_model.py index 88b2a295905b..3b2d83bf9c5d 100644 --- a/vllm/model_executor/models/idefics2_vision_model.py +++ b/vllm/model_executor/models/idefics2_vision_model.py @@ -27,6 +27,8 @@ Idefics2Config, Idefics2VisionConfig) from vllm.attention.layer import MultiHeadAttention +from vllm.model_executor.models.vision import get_vit_attn_backend +from vllm.platforms import _Backend from vllm.distributed import get_tensor_model_parallel_world_size from vllm.model_executor.layers.activation import get_act_fn from vllm.model_executor.layers.linear import (ColumnParallelLinear, @@ -172,6 +174,17 @@ def __init__( ) self.attn = MultiHeadAttention(self.num_heads_per_partition, self.head_dim, self.scale) + + # Detect attention backend at initialization time + self.attn_backend = get_vit_attn_backend(support_fa=True) + + # Validate supported backends + if self.attn_backend not in { + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + }: + raise RuntimeError( + f"Vision attention does not support {self.attn_backend} backend now." + ) def forward( self, @@ -181,8 +194,39 @@ def forward( hidden_states ) # batch_size, q_len, 3 * num_heads_per_partition * head_dim query_states, key_states, value_states = qkv.chunk(3, dim=-1) - out = self.attn(query_states, key_states, value_states) - attn_output, _ = self.out_proj(out) + + batch_size, seq_len, _ = query_states.shape + + # Reshape for attention computation + q = query_states.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) + k = key_states.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) + v = value_states.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) + + # Apply attention using the pre-selected backend + if self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + attn_output = flash_attn_func(q, k, v, softmax_scale=self.scale) + attn_output = attn_output.reshape(batch_size, seq_len, -1) + elif self.attn_backend == _Backend.XFORMERS: + from xformers import ops as xops + # xFormers expects (batch, seq, heads, head_dim) + attn_output = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) + attn_output = attn_output.reshape(batch_size, seq_len, -1) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + attn_output = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) + attn_output = attn_output.reshape(batch_size, seq_len, -1) + else: + # PyTorch SDPA (default and fallback) + q = q.transpose(1, 2) # (batch, heads, seq, head_dim) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + attn_output = torch.nn.functional.scaled_dot_product_attention(q, k, v, scale=self.scale) + attn_output = attn_output.transpose(1, 2).reshape(batch_size, seq_len, -1) + + attn_output, _ = self.out_proj(attn_output) return attn_output diff --git a/vllm/model_executor/models/intern_vit.py b/vllm/model_executor/models/intern_vit.py index 58e8163e0b26..df5e502c3856 100644 --- a/vllm/model_executor/models/intern_vit.py +++ b/vllm/model_executor/models/intern_vit.py @@ -28,6 +28,8 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.model_executor.models.vision import get_vit_attn_backend +from vllm.platforms import _Backend NORM2FN = { 'rms_norm': RMSNorm, @@ -254,6 +256,17 @@ def __init__( var_hidden_size=self.embed_dim) self.proj = nn.Linear(self.dummy_dim, self.embed_dim) + + # Detect attention backend at initialization time + self.attn_backend = get_vit_attn_backend(support_fa=True) + + # Validate supported backends + if self.attn_backend not in { + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + }: + raise RuntimeError( + f"Vision attention does not support {self.attn_backend} backend now." + ) def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape @@ -268,12 +281,30 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: B_, N_, H_, D_ = q.shape q = self.q_norm(q.flatten(-2, -1)).view(B_, N_, H_, D_) k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - - x = F.scaled_dot_product_attention(q, k, v, scale=self.scale) - x = x.transpose(1, 2).reshape(B, N, -1) + + # Apply attention using the pre-selected backend + if self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + x = flash_attn_func(q, k, v, softmax_scale=self.scale) + x = x.reshape(B, N, -1) + elif self.attn_backend == _Backend.XFORMERS: + from xformers import ops as xops + # xFormers expects (batch, seq, heads, head_dim) + x = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) + x = x.reshape(B, N, -1) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + x = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) + x = x.reshape(B, N, -1) + else: + # PyTorch SDPA (default and fallback) + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + x = F.scaled_dot_product_attention(q, k, v, scale=self.scale) + x = x.transpose(1, 2).reshape(B, N, -1) x = self.proj(x) return x diff --git a/vllm/model_executor/models/interns1_vit.py b/vllm/model_executor/models/interns1_vit.py index 300ed17ecaab..740e6d516e29 100644 --- a/vllm/model_executor/models/interns1_vit.py +++ b/vllm/model_executor/models/interns1_vit.py @@ -22,6 +22,8 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.model_executor.models.vision import get_vit_attn_backend +from vllm.platforms import _Backend NORM2FN = { 'rms_norm': RMSNorm, @@ -205,6 +207,17 @@ def __init__( var_hidden_size=self.embed_dim) self.projection_layer = nn.Linear(self.dummy_dim, self.embed_dim) + + # Detect attention backend at initialization time + self.attn_backend = get_vit_attn_backend(support_fa=True) + + # Validate supported backends + if self.attn_backend not in { + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + }: + raise RuntimeError( + f"Vision attention does not support {self.attn_backend} backend now." + ) def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape @@ -221,12 +234,30 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: B_, N_, H_, D_ = q.shape q = self.q_norm(q.flatten(-2, -1)).view(B_, N_, H_, D_) k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - - x = F.scaled_dot_product_attention(q, k, v, scale=self.scale) - x = x.transpose(1, 2).reshape(B, N, -1) + + # Apply attention using the pre-selected backend + if self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + x = flash_attn_func(q, k, v, softmax_scale=self.scale) + x = x.reshape(B, N, -1) + elif self.attn_backend == _Backend.XFORMERS: + from xformers import ops as xops + # xFormers expects (batch, seq, heads, head_dim) + x = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) + x = x.reshape(B, N, -1) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + x = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) + x = x.reshape(B, N, -1) + else: + # PyTorch SDPA (default and fallback) + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + x = F.scaled_dot_product_attention(q, k, v, scale=self.scale) + x = x.transpose(1, 2).reshape(B, N, -1) x = self.projection_layer(x) return x diff --git a/vllm/model_executor/models/mllama.py b/vllm/model_executor/models/mllama.py index cc2216996f03..d5835bebc5b4 100644 --- a/vllm/model_executor/models/mllama.py +++ b/vllm/model_executor/models/mllama.py @@ -53,6 +53,7 @@ from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name) from vllm.model_executor.models.module_mapping import MultiModelKeys +from vllm.model_executor.models.vision import get_vit_attn_backend from vllm.model_executor.sampling_metadata import SamplingMetadata from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalEncDecInputs, @@ -516,6 +517,17 @@ def __init__(self, quant_config=quant_config, prefix=f"{prefix}.o_proj", ) + + # Detect attention backend at initialization time + self.attn_backend = get_vit_attn_backend(support_fa=True) + + # Validate supported backends + if self.attn_backend not in { + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + }: + raise RuntimeError( + f"Vision attention does not support {self.attn_backend} backend now." + ) def forward( self, @@ -525,20 +537,41 @@ def forward( qkv, _ = self.qkv_proj(hidden_state) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) q = q.view(q.shape[0], q.shape[1], self.num_local_heads, - self.head_dim).transpose(1, 2) + self.head_dim) k = k.view(k.shape[0], k.shape[1], self.num_local_heads, - self.head_dim).transpose(1, 2) + self.head_dim) v = v.view(v.shape[0], v.shape[1], self.num_local_heads, - self.head_dim).transpose(1, 2) - - # TODO: remove padding in image encoder - attn_output = F.scaled_dot_product_attention(q, - k, - v, - attn_mask=attention_mask, - dropout_p=0.0) - - attn_output = attn_output.transpose(1, 2).contiguous() + self.head_dim) + + # Apply attention using the pre-selected backend + if self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + # Note: attention_mask is not supported in flash attention + attn_output = flash_attn_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + elif self.attn_backend == _Backend.XFORMERS: + from xformers import ops as xops + # xFormers expects (batch, seq, heads, head_dim) + attn_output = xops.memory_efficient_attention_forward( + q, k, v, attn_bias=attention_mask, scale=1.0 / math.sqrt(self.head_dim) + ) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + attn_output = flash_attn_varlen_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + else: + # PyTorch SDPA (default and fallback) + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + attn_output = F.scaled_dot_product_attention(q, + k, + v, + attn_mask=attention_mask, + dropout_p=0.0) + attn_output = attn_output.transpose(1, 2) + + attn_output = attn_output.contiguous() attn_output = attn_output.reshape(attn_output.shape[0], attn_output.shape[1], -1) output, _ = self.o_proj(attn_output) diff --git a/vllm/model_executor/models/pixtral.py b/vllm/model_executor/models/pixtral.py index a74e01a59697..deac41658a27 100644 --- a/vllm/model_executor/models/pixtral.py +++ b/vllm/model_executor/models/pixtral.py @@ -32,7 +32,9 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader +from vllm.model_executor.models.vision import get_vit_attn_backend from vllm.model_executor.sampling_metadata import SamplingMetadata +from vllm.platforms import _Backend from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalKwargsItems from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalFieldConfig, NestedTensors) @@ -1077,6 +1079,17 @@ def __init__( quant_config=quant_config, prefix=f"{prefix}.o_proj", ) + + # Detect attention backend at initialization time + self.attn_backend = get_vit_attn_backend(support_fa=True) + + # Validate supported backends + if self.attn_backend not in { + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + }: + raise RuntimeError( + f"Vision attention does not support {self.attn_backend} backend now." + ) def forward( self, @@ -1096,15 +1109,31 @@ def forward( cos, sin = position_embeddings q, k = apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=0) - if USE_XFORMERS_OPS: - # Transpose q and k back for attention + # Apply attention using the pre-selected backend + if self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + q = q.transpose(1, 2).contiguous() + k = k.transpose(1, 2).contiguous() + # Note: attention_mask is not supported in flash attention + out = flash_attn_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + elif self.attn_backend == _Backend.XFORMERS: + from xformers import ops as xops + # xFormers expects (batch, seq, heads, head_dim) q = q.transpose(1, 2).contiguous() k = k.transpose(1, 2).contiguous() out = xops.memory_efficient_attention(q, k, v, attn_bias=attention_mask) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + q = q.transpose(1, 2).contiguous() + k = k.transpose(1, 2).contiguous() + out = flash_attn_varlen_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) else: + # PyTorch SDPA (default and fallback) v = v.transpose(1, 2) out = nn.functional.scaled_dot_product_attention( q, k, v, attn_mask=attention_mask) diff --git a/vllm/model_executor/models/step3_vl.py b/vllm/model_executor/models/step3_vl.py index f379d2c15fb6..86200575a369 100644 --- a/vllm/model_executor/models/step3_vl.py +++ b/vllm/model_executor/models/step3_vl.py @@ -25,7 +25,9 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.sampler import SamplerOutput, get_sampler +from vllm.model_executor.models.vision import get_vit_attn_backend from vllm.model_executor.sampling_metadata import SamplingMetadata +from vllm.platforms import _Backend from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalFieldConfig, MultiModalKwargsItems, NestedTensors) @@ -696,6 +698,17 @@ def __init__(self, bias=True, quant_config=quant_config, prefix=prefix) + + # Detect attention backend at initialization time + self.attn_backend = get_vit_attn_backend(support_fa=True) + + # Validate supported backends + if self.attn_backend not in { + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + }: + raise RuntimeError( + f"Vision attention does not support {self.attn_backend} backend now." + ) def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int): return tensor.view(bsz, seq_len, self.num_heads, @@ -714,16 +727,35 @@ def forward( q = q.view(bsz, tgt_len, self.num_heads, self.head_dim) k = k.view(bsz, tgt_len, self.num_heads, self.head_dim) v = v.view(bsz, tgt_len, self.num_heads, self.head_dim) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - attn_output = F.scaled_dot_product_attention(q, - k, - v, - scale=self.scale, - is_causal=False) - attn_output = attn_output.transpose(1, 2).reshape( - bsz, tgt_len, self.num_heads * self.head_dim) + + # Apply attention using the pre-selected backend + if self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + attn_output = flash_attn_func(q, k, v, softmax_scale=self.scale) + attn_output = attn_output.reshape(bsz, tgt_len, self.num_heads * self.head_dim) + elif self.attn_backend == _Backend.XFORMERS: + from xformers import ops as xops + # xFormers expects (batch, seq, heads, head_dim) + attn_output = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) + attn_output = attn_output.reshape(bsz, tgt_len, self.num_heads * self.head_dim) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + attn_output = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) + attn_output = attn_output.reshape(bsz, tgt_len, self.num_heads * self.head_dim) + else: + # PyTorch SDPA (default and fallback) + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) + attn_output = F.scaled_dot_product_attention(q, + k, + v, + scale=self.scale, + is_causal=False) + attn_output = attn_output.transpose(1, 2).reshape( + bsz, tgt_len, self.num_heads * self.head_dim) attn_output, _ = self.out_proj(attn_output) From 624bcdae312c39a12f81d97e5f10a4f42cb5de6e Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Fri, 29 Aug 2025 20:59:22 -0400 Subject: [PATCH 02/28] feat: Add unified VisionAttention interface for automatic backend selection --- vllm/model_executor/models/vision.py | 141 +++++++++++++++++++++++++++ 1 file changed, 141 insertions(+) diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index de30509b1ccb..52e939594276 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -5,6 +5,8 @@ from typing import Final, Generic, Optional, Protocol, TypeVar, Union import torch +import torch.nn as nn +import torch.nn.functional as F from transformers import PretrainedConfig from vllm.attention.selector import get_env_variable_attn_backend @@ -123,3 +125,142 @@ def resolve_visual_encoder_outputs( if post_layer_norm is not None and uses_last_layer: hs_pool[-1] = post_layer_norm(encoder_outputs) return torch.cat(hs_pool, dim=-1) + + +class VisionAttention(torch.nn.Module): + """ + Unified Vision Transformer attention module that automatically selects + the optimal backend based on hardware, compute capability, head size, etc. + + This allows model developers to focus on model architecture without + worrying about attention implementation details. + """ + + def __init__( + self, + embed_dim: int, + num_heads: int, + head_dim: Optional[int] = None, + dropout: float = 0.0, + bias: bool = True, + use_rotary: bool = False, + rotary_dim: Optional[int] = None, + ) -> None: + super().__init__() + + self.embed_dim = embed_dim + self.num_heads = num_heads + self.head_dim = head_dim or (embed_dim // num_heads) + self.dropout = dropout + self.bias = bias + self.use_rotary = use_rotary + self.rotary_dim = rotary_dim or self.head_dim + + # Auto-select optimal backend + self.backend = self._select_backend() + + # Initialize QKV projection + self.qkv = torch.nn.Linear(embed_dim, embed_dim * 3, bias=bias) + self.proj = torch.nn.Linear(embed_dim, embed_dim, bias=bias) + + # Rotary embeddings if needed + if use_rotary: + self.rotary_emb = self._create_rotary_embeddings() + + def _select_backend(self) -> _Backend: + """Automatically select the optimal attention backend.""" + # Check environment override first + env_backend = get_env_variable_attn_backend() + if env_backend is not None: + return env_backend + + # Use existing logic with support for FA + return get_vit_attn_backend(support_fa=True) + + def _create_rotary_embeddings(self): + """Create rotary position embeddings if needed.""" + # This would be implemented based on the specific rotary embedding + # requirements of the model + pass + + def _apply_rotary_embeddings(self, q: torch.Tensor, k: torch.Tensor, + positions: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Apply rotary position embeddings to Q and K.""" + if not self.use_rotary: + return q, k + + # Implementation would depend on the specific rotary embedding method + # For now, return as-is + return q, k + + def _flash_attention_forward(self, q: torch.Tensor, k: torch.Tensor, + v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: + """Forward pass using FlashAttention.""" + try: + from flash_attn import flash_attn_func + return flash_attn_func(q, k, v, dropout_p=self.dropout, causal=False) + except ImportError: + # Fallback to torch SDPA + return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) + + def _torch_sdpa_forward(self, q: torch.Tensor, k: torch.Tensor, + v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: + """Forward pass using torch scaled_dot_product_attention.""" + return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) + + def _xformers_forward(self, q: torch.Tensor, k: torch.Tensor, + v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: + """Forward pass using xFormers.""" + try: + from xformers import ops as xops + return xops.memory_efficient_attention_forward(q, k, v, p=self.dropout) + except ImportError: + # Fallback to torch SDPA + return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) + + def forward( + self, + x: torch.Tensor, + mask: Optional[torch.Tensor] = None, + positions: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """ + Forward pass with automatic backend selection. + + Args: + x: Input tensor of shape (batch_size, seq_len, embed_dim) + mask: Optional attention mask + positions: Optional position indices for rotary embeddings + + Returns: + Output tensor of shape (batch_size, seq_len, embed_dim) + """ + batch_size, seq_len, _ = x.shape + + # Project to QKV + qkv = self.qkv(x) + qkv = qkv.view(batch_size, seq_len, 3, self.num_heads, self.head_dim) + qkv = qkv.permute(2, 0, 3, 1, 4) # (3, batch, heads, seq, head_dim) + q, k, v = qkv[0], qkv[1], qkv[2] + + # Apply rotary embeddings if needed + if positions is not None: + q, k = self._apply_rotary_embeddings(q, k, positions) + + # Select attention implementation based on backend + if self.backend == _Backend.FLASH_ATTN: + attn_output = self._flash_attention_forward(q, k, v, mask) + elif self.backend == _Backend.TORCH_SDPA: + attn_output = self._torch_sdpa_forward(q, k, v, mask) + elif self.backend == _Backend.XFORMERS: + attn_output = self._xformers_forward(q, k, v, mask) + else: + # Fallback to torch SDPA + attn_output = self._torch_sdpa_forward(q, k, v, mask) + + # Project output + attn_output = attn_output.transpose(1, 2).contiguous() + attn_output = attn_output.view(batch_size, seq_len, self.embed_dim) + output = self.proj(attn_output) + + return output From 54160c57073287ec9a93f6f04e24325b17cab129 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Fri, 29 Aug 2025 22:00:40 -0400 Subject: [PATCH 03/28] fix: Remove trailing whitespace in vision.py --- vllm/model_executor/models/vision.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index 52e939594276..b769ebaf2058 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -263,4 +263,4 @@ def forward( attn_output = attn_output.view(batch_size, seq_len, self.embed_dim) output = self.proj(attn_output) - return output + return output \ No newline at end of file From 5e5c80c7280d71461f42c7a786f921056c153a65 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Fri, 29 Aug 2025 22:22:44 -0400 Subject: [PATCH 04/28] feat: Enhance VisionAttention with ROCM_AITER_FA support and proper rotary embeddings --- vllm/model_executor/models/vision.py | 58 ++++++++++++++++++++++++---- 1 file changed, 50 insertions(+), 8 deletions(-) diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index b769ebaf2058..4818f4dd115f 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -179,25 +179,37 @@ def _select_backend(self) -> _Backend: def _create_rotary_embeddings(self): """Create rotary position embeddings if needed.""" - # This would be implemented based on the specific rotary embedding - # requirements of the model - pass + if not self.use_rotary: + return None + + # Create rotary embeddings based on head dimension + try: + from vllm.model_executor.layers.rotary_embedding import get_rope + return get_rope(head_size=self.rotary_dim) + except ImportError: + # Fallback to basic implementation + return None def _apply_rotary_embeddings(self, q: torch.Tensor, k: torch.Tensor, positions: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: """Apply rotary position embeddings to Q and K.""" - if not self.use_rotary: + if not self.use_rotary or self.rotary_emb is None: return q, k - # Implementation would depend on the specific rotary embedding method - # For now, return as-is - return q, k + try: + # Apply rotary embeddings using vLLM's implementation + q = self.rotary_emb(q, positions) + k = self.rotary_emb(k, positions) + return q, k + except Exception: + # Fallback: return as-is if rotary embedding fails + return q, k def _flash_attention_forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: """Forward pass using FlashAttention.""" try: - from flash_attn import flash_attn_func + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func return flash_attn_func(q, k, v, dropout_p=self.dropout, causal=False) except ImportError: # Fallback to torch SDPA @@ -218,6 +230,34 @@ def _xformers_forward(self, q: torch.Tensor, k: torch.Tensor, # Fallback to torch SDPA return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) + def _rocm_aiter_fa_forward(self, q: torch.Tensor, k: torch.Tensor, + v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: + """Forward pass using ROCm Aiter FlashAttention.""" + try: + from aiter import flash_attn_varlen_func + # Reshape for varlen function + batch_size, seq_len, num_heads, head_dim = q.shape + q_flat = q.view(-1, num_heads, head_dim) + k_flat = k.view(-1, num_heads, head_dim) + v_flat = v.view(-1, num_heads, head_dim) + + # Create cu_seqlens for varlen function + cu_seqlens = torch.arange(0, (batch_size + 1) * seq_len, seq_len, + dtype=torch.int32, device=q.device) + + output = flash_attn_varlen_func(q_flat, k_flat, v_flat, + cu_seqlens_q=cu_seqlens, + cu_seqlens_k=cu_seqlens, + max_seqlen_q=seq_len, + max_seqlen_k=seq_len, + dropout_p=self.dropout, + causal=False) + + return output.view(batch_size, seq_len, num_heads, head_dim) + except ImportError: + # Fallback to torch SDPA + return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) + def forward( self, x: torch.Tensor, @@ -250,6 +290,8 @@ def forward( # Select attention implementation based on backend if self.backend == _Backend.FLASH_ATTN: attn_output = self._flash_attention_forward(q, k, v, mask) + elif self.backend == _Backend.ROCM_AITER_FA: + attn_output = self._rocm_aiter_fa_forward(q, k, v, mask) elif self.backend == _Backend.TORCH_SDPA: attn_output = self._torch_sdpa_forward(q, k, v, mask) elif self.backend == _Backend.XFORMERS: From 3537659efa5f5b0e7d676eb5ad9717fb0850bd90 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 14:32:18 -0400 Subject: [PATCH 05/28] feat: unify vision attention implementations using MultiHeadAttention - Add Flash Attention support to MultiHeadAttention class - Simplify Idefics2VisionAttention to use unified MultiHeadAttention - Simplify VisionAttention to use unified MultiHeadAttention - Remove duplicate attention implementations - Maintain backward compatibility while reducing code duplication --- vllm/attention/layer.py | 9 +- .../models/idefics2_vision_model.py | 46 +-------- vllm/model_executor/models/vision.py | 95 +++---------------- 3 files changed, 24 insertions(+), 126 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 237802afccde..1c49a2a0dfe1 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -365,7 +365,8 @@ def __init__( backend = _Backend.XFORMERS self.attn_backend = backend if backend in { - _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1 + _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, + _Backend.FLASH_ATTN, _Backend.ROCM_AITER_FA } else _Backend.TORCH_SDPA if (self.attn_backend == _Backend.XFORMERS @@ -413,6 +414,12 @@ def forward( from torch_xla.experimental.custom_kernel import flash_attention out = flash_attention(query, key, value, sm_scale=self.scale) out = out.transpose(1, 2) + elif self.attn_backend == _Backend.FLASH_ATTN: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + out = flash_attn_func(query, key, value, softmax_scale=self.scale) + elif self.attn_backend == _Backend.ROCM_AITER_FA: + from aiter import flash_attn_varlen_func + out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) return out.reshape(bsz, q_len, -1) diff --git a/vllm/model_executor/models/idefics2_vision_model.py b/vllm/model_executor/models/idefics2_vision_model.py index 7ca23409b899..fde1ab2696a8 100644 --- a/vllm/model_executor/models/idefics2_vision_model.py +++ b/vllm/model_executor/models/idefics2_vision_model.py @@ -27,8 +27,6 @@ Idefics2Config, Idefics2VisionConfig) from vllm.attention.layer import MultiHeadAttention -from vllm.model_executor.models.vision import get_vit_attn_backend -from vllm.platforms import _Backend from vllm.distributed import get_tensor_model_parallel_world_size from vllm.model_executor.layers.activation import get_act_fn from vllm.model_executor.layers.linear import (ColumnParallelLinear, @@ -172,19 +170,9 @@ def __init__( quant_config=quant_config, prefix=f"{prefix}.out_proj", ) + # Use unified MultiHeadAttention with Flash Attention support self.attn = MultiHeadAttention(self.num_heads_per_partition, self.head_dim, self.scale) - - # Detect attention backend at initialization time - self.attn_backend = get_vit_attn_backend(support_fa=True) - - # Validate supported backends - if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA - }: - raise RuntimeError( - f"Vision attention does not support {self.attn_backend} backend now." - ) def forward( self, @@ -195,36 +183,8 @@ def forward( ) # batch_size, q_len, 3 * num_heads_per_partition * head_dim query_states, key_states, value_states = qkv.chunk(3, dim=-1) - batch_size, seq_len, _ = query_states.shape - - # Reshape for attention computation - q = query_states.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) - k = key_states.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) - v = value_states.view(batch_size, seq_len, self.num_heads_per_partition, self.head_dim) - - # Apply attention using the pre-selected backend - if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func - # Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_func(q, k, v, softmax_scale=self.scale) - attn_output = attn_output.reshape(batch_size, seq_len, -1) - elif self.attn_backend == _Backend.XFORMERS: - from xformers import ops as xops - # xFormers expects (batch, seq, heads, head_dim) - attn_output = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) - attn_output = attn_output.reshape(batch_size, seq_len, -1) - elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - # ROCm Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) - attn_output = attn_output.reshape(batch_size, seq_len, -1) - else: - # PyTorch SDPA (default and fallback) - q = q.transpose(1, 2) # (batch, heads, seq, head_dim) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - attn_output = torch.nn.functional.scaled_dot_product_attention(q, k, v, scale=self.scale) - attn_output = attn_output.transpose(1, 2).reshape(batch_size, seq_len, -1) + # Use unified MultiHeadAttention implementation + attn_output = self.attn(query_states, key_states, value_states) attn_output, _ = self.out_proj(attn_output) return attn_output diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index 4818f4dd115f..01d90303487a 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -9,6 +9,7 @@ import torch.nn.functional as F from transformers import PretrainedConfig +from vllm.attention.layer import MultiHeadAttention from vllm.attention.selector import get_env_variable_attn_backend from vllm.logger import init_logger from vllm.platforms import _Backend, current_platform @@ -129,11 +130,10 @@ def resolve_visual_encoder_outputs( class VisionAttention(torch.nn.Module): """ - Unified Vision Transformer attention module that automatically selects - the optimal backend based on hardware, compute capability, head size, etc. + Unified Vision Transformer attention module using MultiHeadAttention. - This allows model developers to focus on model architecture without - worrying about attention implementation details. + This simplified version uses the unified MultiHeadAttention implementation + while maintaining the same interface for backward compatibility. """ def __init__( @@ -156,26 +156,18 @@ def __init__( self.use_rotary = use_rotary self.rotary_dim = rotary_dim or self.head_dim - # Auto-select optimal backend - self.backend = self._select_backend() - # Initialize QKV projection self.qkv = torch.nn.Linear(embed_dim, embed_dim * 3, bias=bias) self.proj = torch.nn.Linear(embed_dim, embed_dim, bias=bias) + # Use unified MultiHeadAttention with Flash Attention support + self.attn = MultiHeadAttention(num_heads, self.head_dim, self.head_dim**-0.5) + # Rotary embeddings if needed if use_rotary: self.rotary_emb = self._create_rotary_embeddings() - def _select_backend(self) -> _Backend: - """Automatically select the optimal attention backend.""" - # Check environment override first - env_backend = get_env_variable_attn_backend() - if env_backend is not None: - return env_backend - - # Use existing logic with support for FA - return get_vit_attn_backend(support_fa=True) + def _create_rotary_embeddings(self): """Create rotary position embeddings if needed.""" @@ -205,58 +197,7 @@ def _apply_rotary_embeddings(self, q: torch.Tensor, k: torch.Tensor, # Fallback: return as-is if rotary embedding fails return q, k - def _flash_attention_forward(self, q: torch.Tensor, k: torch.Tensor, - v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: - """Forward pass using FlashAttention.""" - try: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func - return flash_attn_func(q, k, v, dropout_p=self.dropout, causal=False) - except ImportError: - # Fallback to torch SDPA - return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) - - def _torch_sdpa_forward(self, q: torch.Tensor, k: torch.Tensor, - v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: - """Forward pass using torch scaled_dot_product_attention.""" - return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) - - def _xformers_forward(self, q: torch.Tensor, k: torch.Tensor, - v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: - """Forward pass using xFormers.""" - try: - from xformers import ops as xops - return xops.memory_efficient_attention_forward(q, k, v, p=self.dropout) - except ImportError: - # Fallback to torch SDPA - return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) - - def _rocm_aiter_fa_forward(self, q: torch.Tensor, k: torch.Tensor, - v: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: - """Forward pass using ROCm Aiter FlashAttention.""" - try: - from aiter import flash_attn_varlen_func - # Reshape for varlen function - batch_size, seq_len, num_heads, head_dim = q.shape - q_flat = q.view(-1, num_heads, head_dim) - k_flat = k.view(-1, num_heads, head_dim) - v_flat = v.view(-1, num_heads, head_dim) - - # Create cu_seqlens for varlen function - cu_seqlens = torch.arange(0, (batch_size + 1) * seq_len, seq_len, - dtype=torch.int32, device=q.device) - - output = flash_attn_varlen_func(q_flat, k_flat, v_flat, - cu_seqlens_q=cu_seqlens, - cu_seqlens_k=cu_seqlens, - max_seqlen_q=seq_len, - max_seqlen_k=seq_len, - dropout_p=self.dropout, - causal=False) - - return output.view(batch_size, seq_len, num_heads, head_dim) - except ImportError: - # Fallback to torch SDPA - return torch.nn.functional.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout) + def forward( self, @@ -265,11 +206,11 @@ def forward( positions: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ - Forward pass with automatic backend selection. + Forward pass using unified MultiHeadAttention. Args: x: Input tensor of shape (batch_size, seq_len, embed_dim) - mask: Optional attention mask + mask: Optional attention mask (not used in current implementation) positions: Optional position indices for rotary embeddings Returns: @@ -287,18 +228,8 @@ def forward( if positions is not None: q, k = self._apply_rotary_embeddings(q, k, positions) - # Select attention implementation based on backend - if self.backend == _Backend.FLASH_ATTN: - attn_output = self._flash_attention_forward(q, k, v, mask) - elif self.backend == _Backend.ROCM_AITER_FA: - attn_output = self._rocm_aiter_fa_forward(q, k, v, mask) - elif self.backend == _Backend.TORCH_SDPA: - attn_output = self._torch_sdpa_forward(q, k, v, mask) - elif self.backend == _Backend.XFORMERS: - attn_output = self._xformers_forward(q, k, v, mask) - else: - # Fallback to torch SDPA - attn_output = self._torch_sdpa_forward(q, k, v, mask) + # Use unified MultiHeadAttention + attn_output = self.attn(q, k, v) # Project output attn_output = attn_output.transpose(1, 2).contiguous() From ad3d9c347efe3b8df97aecb583db5a59b5afd2ea Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 16:32:55 -0400 Subject: [PATCH 06/28] fix: add fallback for MultiHeadAttention backend detection - Add try-catch block to handle backend detection failures - Fallback to TORCH_SDPA when platform detection fails - Ensures MultiHeadAttention works without full vLLM installation --- vllm/attention/layer.py | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 1c49a2a0dfe1..d1fe9243cc70 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -350,12 +350,16 @@ def __init__( self.num_queries_per_kv = self.num_heads // self.num_kv_heads dtype = torch.get_default_dtype() - attn_backend = get_attn_backend(head_size, - dtype, - kv_cache_dtype=None, - block_size=16, - is_attention_free=False) - backend = backend_name_to_enum(attn_backend.get_name()) + try: + attn_backend = get_attn_backend(head_size, + dtype, + kv_cache_dtype=None, + block_size=16, + is_attention_free=False) + backend = backend_name_to_enum(attn_backend.get_name()) + except (ValueError, AttributeError): + # Fallback to TORCH_SDPA if backend detection fails + backend = _Backend.TORCH_SDPA if current_platform.is_rocm(): # currently, only torch_sdpa is supported on rocm self.attn_backend = _Backend.TORCH_SDPA From 83981554331fa0f42abc56378b3c419e1760778f Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 16:39:26 -0400 Subject: [PATCH 07/28] fix: correct tensor dimensions in VisionAttention forward method - Fix tensor reshaping for MultiHeadAttention compatibility - Ensure proper (batch, seq, hidden_size) format for attention input - Resolve dimension mismatch error in forward pass --- vllm/model_executor/models/vision.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index 01d90303487a..3128f8f4f418 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -206,7 +206,7 @@ def forward( positions: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ - Forward pass using unified MultiHeadAttention. +拉取 Forward pass using unified MultiHeadAttention. Args: x: Input tensor of shape (batch_size, seq_len, embed_dim) @@ -228,8 +228,13 @@ def forward( if positions is not None: q, k = self._apply_rotary_embeddings(q, k, positions) + # Reshape for MultiHeadAttention: (batch, seq, hidden_size) + q_reshaped = q.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) + k_reshaped = k.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) + v_reshaped = v.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) + # Use unified MultiHeadAttention - attn_output = self.attn(q, k, v) + attn_output = self.attn(q_reshaped, k_reshaped, v_reshaped) # Project output attn_output = attn_output.transpose(1, 2).contiguous() From b13f5ba9b4628c028663754e0ae6aff1c3ca9467 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 17:35:38 -0400 Subject: [PATCH 08/28] Fix MultiHeadAttention backend selection logic, remove forced conversion to xFormers --- vllm/attention/layer.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index d1fe9243cc70..5cd8c2a12a09 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -364,13 +364,9 @@ def __init__( # currently, only torch_sdpa is supported on rocm self.attn_backend = _Backend.TORCH_SDPA else: - if backend in (_Backend.FLASH_ATTN, _Backend.FLASH_ATTN_VLLM_V1, - _Backend.FLEX_ATTENTION): - backend = _Backend.XFORMERS - self.attn_backend = backend if backend in { _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, - _Backend.FLASH_ATTN, _Backend.ROCM_AITER_FA + _Backend.FLASH_ATTN, _Backend.ROCM_AITER_FA, _Backend.FLEX_ATTENTION } else _Backend.TORCH_SDPA if (self.attn_backend == _Backend.XFORMERS From f3a28c2ad594ae649085687579fabab3365851ee Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 17:49:20 -0400 Subject: [PATCH 09/28] Add FLEX_ATTENTION backend support to MultiHeadAttention.forward --- vllm/attention/layer.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 5cd8c2a12a09..71515c529db1 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -420,6 +420,18 @@ def forward( elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) + elif self.attn_backend == _Backend.FLEX_ATTENTION: + # FlexAttention requires specific tensor format + query, key, value = (x.transpose(1, 2) for x in (query, key, value)) + from vllm.attention.backends.flex_attention import FlexAttentionBackend + # Use a simple fallback to torch SDPA for now + out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) + out = out.transpose(1, 2) + else: + # Fallback to torch SDPA for unsupported backends + query, key, value = (x.transpose(1, 2) for x in (query, key, value)) + out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) + out = out.transpose(1, 2) return out.reshape(bsz, q_len, -1) From 0491cd532693a82f5db22b73fe210d7326676dea Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 17:50:45 -0400 Subject: [PATCH 10/28] Remove unnecessary FlexAttention import in MultiHeadAttention --- vllm/attention/layer.py | 1 - 1 file changed, 1 deletion(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 71515c529db1..956c2ec1dfb3 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -423,7 +423,6 @@ def forward( elif self.attn_backend == _Backend.FLEX_ATTENTION: # FlexAttention requires specific tensor format query, key, value = (x.transpose(1, 2) for x in (query, key, value)) - from vllm.attention.backends.flex_attention import FlexAttentionBackend # Use a simple fallback to torch SDPA for now out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) out = out.transpose(1, 2) From 39b4002106fff359a791b3c225189c87b2eca01f Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 17:59:14 -0400 Subject: [PATCH 11/28] Add FLASH_ATTN_VLLM_V1 support to MultiHeadAttention --- vllm/attention/layer.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 956c2ec1dfb3..4e7a9bf115a4 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -366,7 +366,7 @@ def __init__( else: self.attn_backend = backend if backend in { _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, - _Backend.FLASH_ATTN, _Backend.ROCM_AITER_FA, _Backend.FLEX_ATTENTION + _Backend.FLASH_ATTN, _Backend.FLASH_ATTN_VLLM_V1, _Backend.ROCM_AITER_FA, _Backend.FLEX_ATTENTION } else _Backend.TORCH_SDPA if (self.attn_backend == _Backend.XFORMERS @@ -417,6 +417,9 @@ def forward( elif self.attn_backend == _Backend.FLASH_ATTN: from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func out = flash_attn_func(query, key, value, softmax_scale=self.scale) + elif self.attn_backend == _Backend.FLASH_ATTN_VLLM_V1: + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + out = flash_attn_func(query, key, value, softmax_scale=self.scale) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) From 3fd3c9f82ef1838a1ff6b68c13071ac760735c48 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 18:01:23 -0400 Subject: [PATCH 12/28] Fix Flash Attention function import name --- vllm/attention/layer.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 4e7a9bf115a4..fadb0ab4e484 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -415,11 +415,11 @@ def forward( out = flash_attention(query, key, value, sm_scale=self.scale) out = out.transpose(1, 2) elif self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func - out = flash_attn_func(query, key, value, softmax_scale=self.scale) + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_varlen_func + out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) elif self.attn_backend == _Backend.FLASH_ATTN_VLLM_V1: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func - out = flash_attn_func(query, key, value, softmax_scale=self.scale) + from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_varlen_func + out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) From b4bb47b21ae4c86b53d1eb6f222d1ba8643a89bc Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 30 Aug 2025 18:07:47 -0400 Subject: [PATCH 13/28] Fix Flash Attention function call with proper parameters --- vllm/attention/layer.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index fadb0ab4e484..ab8a1d5283e6 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -366,7 +366,7 @@ def __init__( else: self.attn_backend = backend if backend in { _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, - _Backend.FLASH_ATTN, _Backend.FLASH_ATTN_VLLM_V1, _Backend.ROCM_AITER_FA, _Backend.FLEX_ATTENTION + _Backend.ROCM_AITER_FA, _Backend.FLEX_ATTENTION } else _Backend.TORCH_SDPA if (self.attn_backend == _Backend.XFORMERS @@ -414,12 +414,11 @@ def forward( from torch_xla.experimental.custom_kernel import flash_attention out = flash_attention(query, key, value, sm_scale=self.scale) out = out.transpose(1, 2) - elif self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_varlen_func - out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) - elif self.attn_backend == _Backend.FLASH_ATTN_VLLM_V1: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_varlen_func - out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) + elif self.attn_backend in (_Backend.FLASH_ATTN, _Backend.FLASH_ATTN_VLLM_V1): + # Flash Attention for ViT is not essential, fallback to PyTorch SDPA + query, key, value = (x.transpose(1, 2) for x in (query, key, value)) + out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) + out = out.transpose(1, 2) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) From aa0c158dc42cac98b24f731806551a363d22dea9 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sun, 31 Aug 2025 01:56:45 -0400 Subject: [PATCH 14/28] Fix ruff linting errors: line length and unused imports --- vllm/attention/layer.py | 15 +++++++++----- vllm/model_executor/models/intern_vit.py | 12 +++++++---- vllm/model_executor/models/interns1_vit.py | 12 +++++++---- vllm/model_executor/models/mllama.py | 18 ++++++++++------ vllm/model_executor/models/pixtral.py | 15 +++++++++----- vllm/model_executor/models/step3_vl.py | 24 ++++++++++++++-------- vllm/model_executor/models/vision.py | 20 ++++++++++-------- 7 files changed, 76 insertions(+), 40 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index ab8a1d5283e6..7e79fc97b3c1 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -414,24 +414,29 @@ def forward( from torch_xla.experimental.custom_kernel import flash_attention out = flash_attention(query, key, value, sm_scale=self.scale) out = out.transpose(1, 2) - elif self.attn_backend in (_Backend.FLASH_ATTN, _Backend.FLASH_ATTN_VLLM_V1): + elif self.attn_backend in (_Backend.FLASH_ATTN, + _Backend.FLASH_ATTN_VLLM_V1): # Flash Attention for ViT is not essential, fallback to PyTorch SDPA query, key, value = (x.transpose(1, 2) for x in (query, key, value)) - out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) + out = F.scaled_dot_product_attention( + query, key, value, scale=self.scale) out = out.transpose(1, 2) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func - out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) + out = flash_attn_varlen_func( + query, key, value, softmax_scale=self.scale) elif self.attn_backend == _Backend.FLEX_ATTENTION: # FlexAttention requires specific tensor format query, key, value = (x.transpose(1, 2) for x in (query, key, value)) # Use a simple fallback to torch SDPA for now - out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) + out = F.scaled_dot_product_attention( + query, key, value, scale=self.scale) out = out.transpose(1, 2) else: # Fallback to torch SDPA for unsupported backends query, key, value = (x.transpose(1, 2) for x in (query, key, value)) - out = F.scaled_dot_product_attention(query, key, value, scale=self.scale) + out = F.scaled_dot_product_attention( + query, key, value, scale=self.scale) out = out.transpose(1, 2) return out.reshape(bsz, q_len, -1) diff --git a/vllm/model_executor/models/intern_vit.py b/vllm/model_executor/models/intern_vit.py index df5e502c3856..ac0be3f6a6f9 100644 --- a/vllm/model_executor/models/intern_vit.py +++ b/vllm/model_executor/models/intern_vit.py @@ -262,10 +262,12 @@ def __init__( # Validate supported backends if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, + _Backend.ROCM_AITER_FA }: raise RuntimeError( - f"Vision attention does not support {self.attn_backend} backend now." + f"Vision attention does not support {self.attn_backend} " + f"backend now." ) def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -284,14 +286,16 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: # Apply attention using the pre-selected backend if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + from vllm.vllm_flash_attn.flash_attn_interface import ( + flash_attn_func) # Flash Attention expects (batch, seq, heads, head_dim) x = flash_attn_func(q, k, v, softmax_scale=self.scale) x = x.reshape(B, N, -1) elif self.attn_backend == _Backend.XFORMERS: from xformers import ops as xops # xFormers expects (batch, seq, heads, head_dim) - x = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) + x = xops.memory_efficient_attention_forward( + q, k, v, scale=self.scale) x = x.reshape(B, N, -1) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func diff --git a/vllm/model_executor/models/interns1_vit.py b/vllm/model_executor/models/interns1_vit.py index 740e6d516e29..e8bbda0d576d 100644 --- a/vllm/model_executor/models/interns1_vit.py +++ b/vllm/model_executor/models/interns1_vit.py @@ -213,10 +213,12 @@ def __init__( # Validate supported backends if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, + _Backend.ROCM_AITER_FA }: raise RuntimeError( - f"Vision attention does not support {self.attn_backend} backend now." + f"Vision attention does not support {self.attn_backend} " + f"backend now." ) def forward(self, x: torch.Tensor) -> torch.Tensor: @@ -237,14 +239,16 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: # Apply attention using the pre-selected backend if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + from vllm.vllm_flash_attn.flash_attn_interface import ( + flash_attn_func) # Flash Attention expects (batch, seq, heads, head_dim) x = flash_attn_func(q, k, v, softmax_scale=self.scale) x = x.reshape(B, N, -1) elif self.attn_backend == _Backend.XFORMERS: from xformers import ops as xops # xFormers expects (batch, seq, heads, head_dim) - x = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) + x = xops.memory_efficient_attention_forward( + q, k, v, scale=self.scale) x = x.reshape(B, N, -1) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func diff --git a/vllm/model_executor/models/mllama.py b/vllm/model_executor/models/mllama.py index d5835bebc5b4..58df34a61399 100644 --- a/vllm/model_executor/models/mllama.py +++ b/vllm/model_executor/models/mllama.py @@ -523,10 +523,12 @@ def __init__(self, # Validate supported backends if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, + _Backend.ROCM_AITER_FA }: raise RuntimeError( - f"Vision attention does not support {self.attn_backend} backend now." + f"Vision attention does not support {self.attn_backend} " + f"backend now." ) def forward( @@ -545,20 +547,24 @@ def forward( # Apply attention using the pre-selected backend if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + from vllm.vllm_flash_attn.flash_attn_interface import ( + flash_attn_func) # Flash Attention expects (batch, seq, heads, head_dim) # Note: attention_mask is not supported in flash attention - attn_output = flash_attn_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + attn_output = flash_attn_func( + q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) elif self.attn_backend == _Backend.XFORMERS: from xformers import ops as xops # xFormers expects (batch, seq, heads, head_dim) attn_output = xops.memory_efficient_attention_forward( - q, k, v, attn_bias=attention_mask, scale=1.0 / math.sqrt(self.head_dim) + q, k, v, attn_bias=attention_mask, + scale=1.0 / math.sqrt(self.head_dim) ) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func # ROCm Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_varlen_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + attn_output = flash_attn_varlen_func( + q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) else: # PyTorch SDPA (default and fallback) q = q.transpose(1, 2) diff --git a/vllm/model_executor/models/pixtral.py b/vllm/model_executor/models/pixtral.py index 3bd6fbd9d2b9..56fbb214cf59 100644 --- a/vllm/model_executor/models/pixtral.py +++ b/vllm/model_executor/models/pixtral.py @@ -1087,10 +1087,12 @@ def __init__( # Validate supported backends if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, + _Backend.ROCM_AITER_FA }: raise RuntimeError( - f"Vision attention does not support {self.attn_backend} backend now." + f"Vision attention does not support {self.attn_backend} " + f"backend now." ) def forward( @@ -1113,12 +1115,14 @@ def forward( # Apply attention using the pre-selected backend if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + from vllm.vllm_flash_attn.flash_attn_interface import ( + flash_attn_func) # Flash Attention expects (batch, seq, heads, head_dim) q = q.transpose(1, 2).contiguous() k = k.transpose(1, 2).contiguous() # Note: attention_mask is not supported in flash attention - out = flash_attn_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + out = flash_attn_func( + q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) elif self.attn_backend == _Backend.XFORMERS: from xformers import ops as xops # xFormers expects (batch, seq, heads, head_dim) @@ -1133,7 +1137,8 @@ def forward( # ROCm Flash Attention expects (batch, seq, heads, head_dim) q = q.transpose(1, 2).contiguous() k = k.transpose(1, 2).contiguous() - out = flash_attn_varlen_func(q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) + out = flash_attn_varlen_func( + q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) else: # PyTorch SDPA (default and fallback) v = v.transpose(1, 2) diff --git a/vllm/model_executor/models/step3_vl.py b/vllm/model_executor/models/step3_vl.py index 86200575a369..dffaea2a7dd5 100644 --- a/vllm/model_executor/models/step3_vl.py +++ b/vllm/model_executor/models/step3_vl.py @@ -704,10 +704,12 @@ def __init__(self, # Validate supported backends if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.ROCM_AITER_FA + _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, + _Backend.ROCM_AITER_FA }: raise RuntimeError( - f"Vision attention does not support {self.attn_backend} backend now." + f"Vision attention does not support {self.attn_backend} " + f"backend now." ) def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int): @@ -730,20 +732,26 @@ def forward( # Apply attention using the pre-selected backend if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import flash_attn_func + from vllm.vllm_flash_attn.flash_attn_interface import ( + flash_attn_func) # Flash Attention expects (batch, seq, heads, head_dim) attn_output = flash_attn_func(q, k, v, softmax_scale=self.scale) - attn_output = attn_output.reshape(bsz, tgt_len, self.num_heads * self.head_dim) + attn_output = attn_output.reshape( + bsz, tgt_len, self.num_heads * self.head_dim) elif self.attn_backend == _Backend.XFORMERS: from xformers import ops as xops # xFormers expects (batch, seq, heads, head_dim) - attn_output = xops.memory_efficient_attention_forward(q, k, v, scale=self.scale) - attn_output = attn_output.reshape(bsz, tgt_len, self.num_heads * self.head_dim) + attn_output = xops.memory_efficient_attention_forward( + q, k, v, scale=self.scale) + attn_output = attn_output.reshape( + bsz, tgt_len, self.num_heads * self.head_dim) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func # ROCm Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) - attn_output = attn_output.reshape(bsz, tgt_len, self.num_heads * self.head_dim) + attn_output = flash_attn_varlen_func( + q, k, v, softmax_scale=self.scale) + attn_output = attn_output.reshape( + bsz, tgt_len, self.num_heads * self.head_dim) else: # PyTorch SDPA (default and fallback) q = q.transpose(1, 2) diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index 3128f8f4f418..1fcb687116bb 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -5,8 +5,6 @@ from typing import Final, Generic, Optional, Protocol, TypeVar, Union import torch -import torch.nn as nn -import torch.nn.functional as F from transformers import PretrainedConfig from vllm.attention.layer import MultiHeadAttention @@ -161,7 +159,8 @@ def __init__( self.proj = torch.nn.Linear(embed_dim, embed_dim, bias=bias) # Use unified MultiHeadAttention with Flash Attention support - self.attn = MultiHeadAttention(num_heads, self.head_dim, self.head_dim**-0.5) + self.attn = MultiHeadAttention( + num_heads, self.head_dim, self.head_dim**-0.5) # Rotary embeddings if needed if use_rotary: @@ -182,8 +181,10 @@ def _create_rotary_embeddings(self): # Fallback to basic implementation return None - def _apply_rotary_embeddings(self, q: torch.Tensor, k: torch.Tensor, - positions: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + def _apply_rotary_embeddings( + self, q: torch.Tensor, k: torch.Tensor, + positions: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: """Apply rotary position embeddings to Q and K.""" if not self.use_rotary or self.rotary_emb is None: return q, k @@ -229,9 +230,12 @@ def forward( q, k = self._apply_rotary_embeddings(q, k, positions) # Reshape for MultiHeadAttention: (batch, seq, hidden_size) - q_reshaped = q.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) - k_reshaped = k.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) - v_reshaped = v.transpose(1, 2).contiguous().view(batch_size, seq_len, -1) + q_reshaped = q.transpose(1, 2).contiguous().view( + batch_size, seq_len, -1) + k_reshaped = k.transpose(1, 2).contiguous().view( + batch_size, seq_len, -1) + v_reshaped = v.transpose(1, 2).contiguous().view( + batch_size, seq_len, -1) # Use unified MultiHeadAttention attn_output = self.attn(q_reshaped, k_reshaped, v_reshaped) From bb1b8c83fc50756f5d4195e87d91462fbaffe6fb Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sun, 31 Aug 2025 03:52:40 -0400 Subject: [PATCH 15/28] Simplify vision attention modules to use unified MultiHeadAttention --- vllm/model_executor/models/intern_vit.py | 46 ++------ vllm/model_executor/models/interns1_vit.py | 51 ++------- vllm/model_executor/models/mllama.py | 59 ++-------- vllm/model_executor/models/pixtral.py | 58 ++-------- vllm/model_executor/models/step3_vl.py | 66 ++--------- vllm/model_executor/models/vision.py | 125 +-------------------- 6 files changed, 46 insertions(+), 359 deletions(-) diff --git a/vllm/model_executor/models/intern_vit.py b/vllm/model_executor/models/intern_vit.py index ac0be3f6a6f9..f8a2ac2c9243 100644 --- a/vllm/model_executor/models/intern_vit.py +++ b/vllm/model_executor/models/intern_vit.py @@ -28,8 +28,7 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader -from vllm.model_executor.models.vision import get_vit_attn_backend -from vllm.platforms import _Backend + NORM2FN = { 'rms_norm': RMSNorm, @@ -257,18 +256,9 @@ def __init__( self.proj = nn.Linear(self.dummy_dim, self.embed_dim) - # Detect attention backend at initialization time - self.attn_backend = get_vit_attn_backend(support_fa=True) - - # Validate supported backends - if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, - _Backend.ROCM_AITER_FA - }: - raise RuntimeError( - f"Vision attention does not support {self.attn_backend} " - f"backend now." - ) + # Use unified MultiHeadAttention with automatic backend selection + self.attn = MultiHeadAttention( + self.num_heads, self.head_dim, self.scale) def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape @@ -284,31 +274,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: q = self.q_norm(q.flatten(-2, -1)).view(B_, N_, H_, D_) k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) - # Apply attention using the pre-selected backend - if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import ( - flash_attn_func) - # Flash Attention expects (batch, seq, heads, head_dim) - x = flash_attn_func(q, k, v, softmax_scale=self.scale) - x = x.reshape(B, N, -1) - elif self.attn_backend == _Backend.XFORMERS: - from xformers import ops as xops - # xFormers expects (batch, seq, heads, head_dim) - x = xops.memory_efficient_attention_forward( - q, k, v, scale=self.scale) - x = x.reshape(B, N, -1) - elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - # ROCm Flash Attention expects (batch, seq, heads, head_dim) - x = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) - x = x.reshape(B, N, -1) - else: - # PyTorch SDPA (default and fallback) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - x = F.scaled_dot_product_attention(q, k, v, scale=self.scale) - x = x.transpose(1, 2).reshape(B, N, -1) + # Use unified MultiHeadAttention with automatic backend selection + x = self.attn(q, k, v) + x = x.reshape(B, N, -1) x = self.proj(x) return x diff --git a/vllm/model_executor/models/interns1_vit.py b/vllm/model_executor/models/interns1_vit.py index e8bbda0d576d..7ea8d85b8717 100644 --- a/vllm/model_executor/models/interns1_vit.py +++ b/vllm/model_executor/models/interns1_vit.py @@ -12,7 +12,6 @@ import torch import torch.nn as nn -import torch.nn.functional as F from transformers import PretrainedConfig from transformers.utils import torch_int @@ -22,8 +21,7 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader -from vllm.model_executor.models.vision import get_vit_attn_backend -from vllm.platforms import _Backend +from vllm.attention.layer import MultiHeadAttention NORM2FN = { 'rms_norm': RMSNorm, @@ -208,18 +206,9 @@ def __init__( self.projection_layer = nn.Linear(self.dummy_dim, self.embed_dim) - # Detect attention backend at initialization time - self.attn_backend = get_vit_attn_backend(support_fa=True) - - # Validate supported backends - if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, - _Backend.ROCM_AITER_FA - }: - raise RuntimeError( - f"Vision attention does not support {self.attn_backend} " - f"backend now." - ) + # Use unified MultiHeadAttention with automatic backend selection + self.attn = MultiHeadAttention( + self.num_heads, self.head_dim, self.scale) def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape @@ -228,40 +217,14 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: k = self.k_proj(x) v = self.v_proj(x) - q = q.view(B, N, self.num_heads, self.head_dim) - k = k.view(B, N, self.num_heads, self.head_dim) - v = v.view(B, N, self.num_heads, self.head_dim) - if self.qk_normalization: B_, N_, H_, D_ = q.shape q = self.q_norm(q.flatten(-2, -1)).view(B_, N_, H_, D_) k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) - # Apply attention using the pre-selected backend - if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import ( - flash_attn_func) - # Flash Attention expects (batch, seq, heads, head_dim) - x = flash_attn_func(q, k, v, softmax_scale=self.scale) - x = x.reshape(B, N, -1) - elif self.attn_backend == _Backend.XFORMERS: - from xformers import ops as xops - # xFormers expects (batch, seq, heads, head_dim) - x = xops.memory_efficient_attention_forward( - q, k, v, scale=self.scale) - x = x.reshape(B, N, -1) - elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - # ROCm Flash Attention expects (batch, seq, heads, head_dim) - x = flash_attn_varlen_func(q, k, v, softmax_scale=self.scale) - x = x.reshape(B, N, -1) - else: - # PyTorch SDPA (default and fallback) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - x = F.scaled_dot_product_attention(q, k, v, scale=self.scale) - x = x.transpose(1, 2).reshape(B, N, -1) + # Use unified MultiHeadAttention with automatic backend selection + x = self.attn(q, k, v) + x = x.reshape(B, N, -1) x = self.projection_layer(x) return x diff --git a/vllm/model_executor/models/mllama.py b/vllm/model_executor/models/mllama.py index 58df34a61399..89d00fda0166 100644 --- a/vllm/model_executor/models/mllama.py +++ b/vllm/model_executor/models/mllama.py @@ -53,7 +53,7 @@ from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name) from vllm.model_executor.models.module_mapping import MultiModelKeys -from vllm.model_executor.models.vision import get_vit_attn_backend + from vllm.model_executor.sampling_metadata import SamplingMetadata from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalEncDecInputs, @@ -71,6 +71,7 @@ from .interfaces import SupportsMultiModal, SupportsV0Only from .llama import LlamaDecoderLayer, LlamaMLP from .utils import AutoWeightsLoader, WeightsMapper, maybe_prefix +from vllm.attention.layer import MultiHeadAttention logger = init_logger(__name__) @@ -518,18 +519,9 @@ def __init__(self, prefix=f"{prefix}.o_proj", ) - # Detect attention backend at initialization time - self.attn_backend = get_vit_attn_backend(support_fa=True) - - # Validate supported backends - if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, - _Backend.ROCM_AITER_FA - }: - raise RuntimeError( - f"Vision attention does not support {self.attn_backend} " - f"backend now." - ) + # Use unified MultiHeadAttention with automatic backend selection + self.attn = MultiHeadAttention( + self.num_local_heads, self.head_dim, 1.0 / math.sqrt(self.head_dim)) def forward( self, @@ -538,44 +530,9 @@ def forward( ) -> torch.Tensor: qkv, _ = self.qkv_proj(hidden_state) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - q = q.view(q.shape[0], q.shape[1], self.num_local_heads, - self.head_dim) - k = k.view(k.shape[0], k.shape[1], self.num_local_heads, - self.head_dim) - v = v.view(v.shape[0], v.shape[1], self.num_local_heads, - self.head_dim) - - # Apply attention using the pre-selected backend - if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import ( - flash_attn_func) - # Flash Attention expects (batch, seq, heads, head_dim) - # Note: attention_mask is not supported in flash attention - attn_output = flash_attn_func( - q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) - elif self.attn_backend == _Backend.XFORMERS: - from xformers import ops as xops - # xFormers expects (batch, seq, heads, head_dim) - attn_output = xops.memory_efficient_attention_forward( - q, k, v, attn_bias=attention_mask, - scale=1.0 / math.sqrt(self.head_dim) - ) - elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - # ROCm Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_varlen_func( - q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) - else: - # PyTorch SDPA (default and fallback) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - attn_output = F.scaled_dot_product_attention(q, - k, - v, - attn_mask=attention_mask, - dropout_p=0.0) - attn_output = attn_output.transpose(1, 2) + + # Use unified MultiHeadAttention with automatic backend selection + attn_output = self.attn(q, k, v) attn_output = attn_output.contiguous() attn_output = attn_output.reshape(attn_output.shape[0], diff --git a/vllm/model_executor/models/pixtral.py b/vllm/model_executor/models/pixtral.py index 56fbb214cf59..49aead0fbb31 100644 --- a/vllm/model_executor/models/pixtral.py +++ b/vllm/model_executor/models/pixtral.py @@ -32,9 +32,9 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader -from vllm.model_executor.models.vision import get_vit_attn_backend +from vllm.attention.layer import MultiHeadAttention from vllm.model_executor.sampling_metadata import SamplingMetadata -from vllm.platforms import _Backend + from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalKwargsItems from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalFieldConfig, NestedTensors) @@ -1082,18 +1082,9 @@ def __init__( prefix=f"{prefix}.o_proj", ) - # Detect attention backend at initialization time - self.attn_backend = get_vit_attn_backend(support_fa=True) - - # Validate supported backends - if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, - _Backend.ROCM_AITER_FA - }: - raise RuntimeError( - f"Vision attention does not support {self.attn_backend} " - f"backend now." - ) + # Use unified MultiHeadAttention with automatic backend selection + self.attn = MultiHeadAttention( + self.n_heads, self.head_dim, 1.0 / math.sqrt(self.head_dim)) def forward( self, @@ -1113,38 +1104,13 @@ def forward( cos, sin = position_embeddings q, k = apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=0) - # Apply attention using the pre-selected backend - if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import ( - flash_attn_func) - # Flash Attention expects (batch, seq, heads, head_dim) - q = q.transpose(1, 2).contiguous() - k = k.transpose(1, 2).contiguous() - # Note: attention_mask is not supported in flash attention - out = flash_attn_func( - q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) - elif self.attn_backend == _Backend.XFORMERS: - from xformers import ops as xops - # xFormers expects (batch, seq, heads, head_dim) - q = q.transpose(1, 2).contiguous() - k = k.transpose(1, 2).contiguous() - out = xops.memory_efficient_attention(q, - k, - v, - attn_bias=attention_mask) - elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - # ROCm Flash Attention expects (batch, seq, heads, head_dim) - q = q.transpose(1, 2).contiguous() - k = k.transpose(1, 2).contiguous() - out = flash_attn_varlen_func( - q, k, v, softmax_scale=1.0 / math.sqrt(self.head_dim)) - else: - # PyTorch SDPA (default and fallback) - v = v.transpose(1, 2) - out = nn.functional.scaled_dot_product_attention( - q, k, v, attn_mask=attention_mask) - out = out.transpose(1, 2) + # Reshape for MultiHeadAttention: (batch, seq, hidden_size) + q_reshaped = q.transpose(1, 2).contiguous().view(batch, patches, -1) + k_reshaped = k.transpose(1, 2).contiguous().view(batch, patches, -1) + v_reshaped = v.transpose(1, 2).contiguous().view(batch, patches, -1) + + # Use unified MultiHeadAttention with automatic backend selection + out = self.attn(q_reshaped, k_reshaped, v_reshaped) out = out.view(batch, patches, self.n_heads * self.head_dim) attn_output, _ = self.o_proj(out) diff --git a/vllm/model_executor/models/step3_vl.py b/vllm/model_executor/models/step3_vl.py index dffaea2a7dd5..8df80940c268 100644 --- a/vllm/model_executor/models/step3_vl.py +++ b/vllm/model_executor/models/step3_vl.py @@ -25,9 +25,8 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.sampler import SamplerOutput, get_sampler -from vllm.model_executor.models.vision import get_vit_attn_backend +from vllm.attention.layer import MultiHeadAttention from vllm.model_executor.sampling_metadata import SamplingMetadata -from vllm.platforms import _Backend from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalFieldConfig, MultiModalKwargsItems, NestedTensors) @@ -699,22 +698,9 @@ def __init__(self, quant_config=quant_config, prefix=prefix) - # Detect attention backend at initialization time - self.attn_backend = get_vit_attn_backend(support_fa=True) - - # Validate supported backends - if self.attn_backend not in { - _Backend.FLASH_ATTN, _Backend.TORCH_SDPA, _Backend.XFORMERS, - _Backend.ROCM_AITER_FA - }: - raise RuntimeError( - f"Vision attention does not support {self.attn_backend} " - f"backend now." - ) - - def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int): - return tensor.view(bsz, seq_len, self.num_heads, - self.head_dim).transpose(1, 2).contiguous() + # Use unified MultiHeadAttention with automatic backend selection + self.attn = MultiHeadAttention( + self.num_heads, self.head_dim, self.scale) def forward( self, @@ -726,44 +712,14 @@ def forward( # get query proj qkv, _ = self.qkv_proj(hidden_states) q, k, v = qkv.chunk(chunks=3, dim=-1) - q = q.view(bsz, tgt_len, self.num_heads, self.head_dim) - k = k.view(bsz, tgt_len, self.num_heads, self.head_dim) - v = v.view(bsz, tgt_len, self.num_heads, self.head_dim) - # Apply attention using the pre-selected backend - if self.attn_backend == _Backend.FLASH_ATTN: - from vllm.vllm_flash_attn.flash_attn_interface import ( - flash_attn_func) - # Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_func(q, k, v, softmax_scale=self.scale) - attn_output = attn_output.reshape( - bsz, tgt_len, self.num_heads * self.head_dim) - elif self.attn_backend == _Backend.XFORMERS: - from xformers import ops as xops - # xFormers expects (batch, seq, heads, head_dim) - attn_output = xops.memory_efficient_attention_forward( - q, k, v, scale=self.scale) - attn_output = attn_output.reshape( - bsz, tgt_len, self.num_heads * self.head_dim) - elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - # ROCm Flash Attention expects (batch, seq, heads, head_dim) - attn_output = flash_attn_varlen_func( - q, k, v, softmax_scale=self.scale) - attn_output = attn_output.reshape( - bsz, tgt_len, self.num_heads * self.head_dim) - else: - # PyTorch SDPA (default and fallback) - q = q.transpose(1, 2) - k = k.transpose(1, 2) - v = v.transpose(1, 2) - attn_output = F.scaled_dot_product_attention(q, - k, - v, - scale=self.scale, - is_causal=False) - attn_output = attn_output.transpose(1, 2).reshape( - bsz, tgt_len, self.num_heads * self.head_dim) + # Reshape for MultiHeadAttention: (batch, seq, hidden_size) + q_reshaped = q.view(bsz, tgt_len, -1) + k_reshaped = k.view(bsz, tgt_len, -1) + v_reshaped = v.view(bsz, tgt_len, -1) + + # Use unified MultiHeadAttention with automatic backend selection + attn_output = self.attn(q_reshaped, k_reshaped, v_reshaped) attn_output, _ = self.out_proj(attn_output) diff --git a/vllm/model_executor/models/vision.py b/vllm/model_executor/models/vision.py index 1fcb687116bb..c16aa5ac608f 100644 --- a/vllm/model_executor/models/vision.py +++ b/vllm/model_executor/models/vision.py @@ -7,7 +7,6 @@ import torch from transformers import PretrainedConfig -from vllm.attention.layer import MultiHeadAttention from vllm.attention.selector import get_env_variable_attn_backend from vllm.logger import init_logger from vllm.platforms import _Backend, current_platform @@ -123,126 +122,4 @@ def resolve_visual_encoder_outputs( uses_last_layer = feature_sample_layers[-1] in (len(hs_pool) - 1, -1) if post_layer_norm is not None and uses_last_layer: hs_pool[-1] = post_layer_norm(encoder_outputs) - return torch.cat(hs_pool, dim=-1) - - -class VisionAttention(torch.nn.Module): - """ - Unified Vision Transformer attention module using MultiHeadAttention. - - This simplified version uses the unified MultiHeadAttention implementation - while maintaining the same interface for backward compatibility. - """ - - def __init__( - self, - embed_dim: int, - num_heads: int, - head_dim: Optional[int] = None, - dropout: float = 0.0, - bias: bool = True, - use_rotary: bool = False, - rotary_dim: Optional[int] = None, - ) -> None: - super().__init__() - - self.embed_dim = embed_dim - self.num_heads = num_heads - self.head_dim = head_dim or (embed_dim // num_heads) - self.dropout = dropout - self.bias = bias - self.use_rotary = use_rotary - self.rotary_dim = rotary_dim or self.head_dim - - # Initialize QKV projection - self.qkv = torch.nn.Linear(embed_dim, embed_dim * 3, bias=bias) - self.proj = torch.nn.Linear(embed_dim, embed_dim, bias=bias) - - # Use unified MultiHeadAttention with Flash Attention support - self.attn = MultiHeadAttention( - num_heads, self.head_dim, self.head_dim**-0.5) - - # Rotary embeddings if needed - if use_rotary: - self.rotary_emb = self._create_rotary_embeddings() - - - - def _create_rotary_embeddings(self): - """Create rotary position embeddings if needed.""" - if not self.use_rotary: - return None - - # Create rotary embeddings based on head dimension - try: - from vllm.model_executor.layers.rotary_embedding import get_rope - return get_rope(head_size=self.rotary_dim) - except ImportError: - # Fallback to basic implementation - return None - - def _apply_rotary_embeddings( - self, q: torch.Tensor, k: torch.Tensor, - positions: torch.Tensor - ) -> tuple[torch.Tensor, torch.Tensor]: - """Apply rotary position embeddings to Q and K.""" - if not self.use_rotary or self.rotary_emb is None: - return q, k - - try: - # Apply rotary embeddings using vLLM's implementation - q = self.rotary_emb(q, positions) - k = self.rotary_emb(k, positions) - return q, k - except Exception: - # Fallback: return as-is if rotary embedding fails - return q, k - - - - def forward( - self, - x: torch.Tensor, - mask: Optional[torch.Tensor] = None, - positions: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - """ -拉取 Forward pass using unified MultiHeadAttention. - - Args: - x: Input tensor of shape (batch_size, seq_len, embed_dim) - mask: Optional attention mask (not used in current implementation) - positions: Optional position indices for rotary embeddings - - Returns: - Output tensor of shape (batch_size, seq_len, embed_dim) - """ - batch_size, seq_len, _ = x.shape - - # Project to QKV - qkv = self.qkv(x) - qkv = qkv.view(batch_size, seq_len, 3, self.num_heads, self.head_dim) - qkv = qkv.permute(2, 0, 3, 1, 4) # (3, batch, heads, seq, head_dim) - q, k, v = qkv[0], qkv[1], qkv[2] - - # Apply rotary embeddings if needed - if positions is not None: - q, k = self._apply_rotary_embeddings(q, k, positions) - - # Reshape for MultiHeadAttention: (batch, seq, hidden_size) - q_reshaped = q.transpose(1, 2).contiguous().view( - batch_size, seq_len, -1) - k_reshaped = k.transpose(1, 2).contiguous().view( - batch_size, seq_len, -1) - v_reshaped = v.transpose(1, 2).contiguous().view( - batch_size, seq_len, -1) - - # Use unified MultiHeadAttention - attn_output = self.attn(q_reshaped, k_reshaped, v_reshaped) - - # Project output - attn_output = attn_output.transpose(1, 2).contiguous() - attn_output = attn_output.view(batch_size, seq_len, self.embed_dim) - output = self.proj(attn_output) - - return output \ No newline at end of file + return torch.cat(hs_pool, dim=-1) \ No newline at end of file From b4c4037b815acf8537f36816f0f91dd224dcda51 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 1 Sep 2025 20:55:37 -0400 Subject: [PATCH 16/28] Fix MultiHeadAttention issues based on review feedback - Implement real FlashAttention support instead of fallback to SDPA - Remove redundant reshape operations in Step3VL forward method - Remove unnecessary try-except block in backend detection - Fix ROCm FlashAttention implementation with correct API - Apply MultiHeadAttention to idefics2 vision model Addresses review comments from @Isotr0py --- vllm/attention/layer.py | 33 +++++++++---------- .../models/idefics2_vision_model.py | 5 ++- vllm/model_executor/models/step3_vl.py | 7 +--- 3 files changed, 18 insertions(+), 27 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 7e79fc97b3c1..611c2d3c6474 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -350,16 +350,12 @@ def __init__( self.num_queries_per_kv = self.num_heads // self.num_kv_heads dtype = torch.get_default_dtype() - try: - attn_backend = get_attn_backend(head_size, - dtype, - kv_cache_dtype=None, - block_size=16, - is_attention_free=False) - backend = backend_name_to_enum(attn_backend.get_name()) - except (ValueError, AttributeError): - # Fallback to TORCH_SDPA if backend detection fails - backend = _Backend.TORCH_SDPA + attn_backend = get_attn_backend(head_size, + dtype, + kv_cache_dtype=None, + block_size=16, + is_attention_free=False) + backend = backend_name_to_enum(attn_backend.get_name()) if current_platform.is_rocm(): # currently, only torch_sdpa is supported on rocm self.attn_backend = _Backend.TORCH_SDPA @@ -416,15 +412,16 @@ def forward( out = out.transpose(1, 2) elif self.attn_backend in (_Backend.FLASH_ATTN, _Backend.FLASH_ATTN_VLLM_V1): - # Flash Attention for ViT is not essential, fallback to PyTorch SDPA - query, key, value = (x.transpose(1, 2) for x in (query, key, value)) - out = F.scaled_dot_product_attention( - query, key, value, scale=self.scale) - out = out.transpose(1, 2) + # Use real Flash Attention implementation + from flash_attn import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) + out = flash_attn_func(query, key, value, softmax_scale=self.scale) + out = out.reshape(bsz, q_len, -1) elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_varlen_func - out = flash_attn_varlen_func( - query, key, value, softmax_scale=self.scale) + from aiter import flash_attn_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) + out = flash_attn_func(query, key, value, softmax_scale=self.scale) + out = out.reshape(bsz, q_len, -1) elif self.attn_backend == _Backend.FLEX_ATTENTION: # FlexAttention requires specific tensor format query, key, value = (x.transpose(1, 2) for x in (query, key, value)) diff --git a/vllm/model_executor/models/idefics2_vision_model.py b/vllm/model_executor/models/idefics2_vision_model.py index fde1ab2696a8..7c77ed6dfc56 100644 --- a/vllm/model_executor/models/idefics2_vision_model.py +++ b/vllm/model_executor/models/idefics2_vision_model.py @@ -184,9 +184,8 @@ def forward( query_states, key_states, value_states = qkv.chunk(3, dim=-1) # Use unified MultiHeadAttention implementation - attn_output = self.attn(query_states, key_states, value_states) - - attn_output, _ = self.out_proj(attn_output) + out = self.attn(query_states, key_states, value_states) + attn_output, _ = self.out_proj(out) return attn_output diff --git a/vllm/model_executor/models/step3_vl.py b/vllm/model_executor/models/step3_vl.py index 8df80940c268..977867866dae 100644 --- a/vllm/model_executor/models/step3_vl.py +++ b/vllm/model_executor/models/step3_vl.py @@ -713,13 +713,8 @@ def forward( qkv, _ = self.qkv_proj(hidden_states) q, k, v = qkv.chunk(chunks=3, dim=-1) - # Reshape for MultiHeadAttention: (batch, seq, hidden_size) - q_reshaped = q.view(bsz, tgt_len, -1) - k_reshaped = k.view(bsz, tgt_len, -1) - v_reshaped = v.view(bsz, tgt_len, -1) - # Use unified MultiHeadAttention with automatic backend selection - attn_output = self.attn(q_reshaped, k_reshaped, v_reshaped) + attn_output = self.attn(q, k, v) attn_output, _ = self.out_proj(attn_output) From 0e7cc3c65c8294a16bb9d628b33c05288caf3ea6 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 1 Sep 2025 21:30:32 -0400 Subject: [PATCH 17/28] Fix MultiHeadAttention FlashAttention implementation and ROCm AITer FA support --- vllm/attention/layer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 611c2d3c6474..b582fda32e2b 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -418,9 +418,9 @@ def forward( out = flash_attn_func(query, key, value, softmax_scale=self.scale) out = out.reshape(bsz, q_len, -1) elif self.attn_backend == _Backend.ROCM_AITER_FA: - from aiter import flash_attn_func + from aiter import flash_attn_varlen_func # ROCm Flash Attention expects (batch, seq, heads, head_dim) - out = flash_attn_func(query, key, value, softmax_scale=self.scale) + out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) out = out.reshape(bsz, q_len, -1) elif self.attn_backend == _Backend.FLEX_ATTENTION: # FlexAttention requires specific tensor format From 4b80a176e294b52a62a80677eade6689b9a04d52 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Wed, 3 Sep 2025 16:40:16 -0400 Subject: [PATCH 18/28] refactor: address Isotr0py's review comments - Use vllm.vllm_flash_attn instead of external flash_attn package - Remove unnecessary .contiguous(), .reshape(), and .view() calls - Remove unused FLEX_ATTENTION backend - Fix import ordering with isort - Optimize code formatting with yapf Addresses review comments from Isotr0py on: - Flash Attention import optimization - Removal of redundant tensor operations - Backend cleanup and code structure improvements --- vllm/attention/layer.py | 40 ++++++++++++---------- vllm/model_executor/models/intern_vit.py | 3 +- vllm/model_executor/models/interns1_vit.py | 11 +++--- vllm/model_executor/models/mllama.py | 12 +++---- vllm/model_executor/models/pixtral.py | 12 +++---- vllm/model_executor/models/step3_vl.py | 10 +++--- 6 files changed, 42 insertions(+), 46 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index b582fda32e2b..23032b7d043d 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -361,9 +361,11 @@ def __init__( self.attn_backend = _Backend.TORCH_SDPA else: self.attn_backend = backend if backend in { - _Backend.TORCH_SDPA, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, - _Backend.ROCM_AITER_FA, _Backend.FLEX_ATTENTION - } else _Backend.TORCH_SDPA + _Backend.TORCH_SDPA, + _Backend.XFORMERS, + _Backend.PALLAS_VLLM_V1, + _Backend.ROCM_AITER_FA, + } else current_platform.get_vit_attn_backend() if (self.attn_backend == _Backend.XFORMERS and not check_xformers_availability()): @@ -410,30 +412,30 @@ def forward( from torch_xla.experimental.custom_kernel import flash_attention out = flash_attention(query, key, value, sm_scale=self.scale) out = out.transpose(1, 2) - elif self.attn_backend in (_Backend.FLASH_ATTN, - _Backend.FLASH_ATTN_VLLM_V1): - # Use real Flash Attention implementation - from flash_attn import flash_attn_func + elif self.attn_backend in (_Backend.FLASH_ATTN, + _Backend.FLASH_ATTN_VLLM_V1): + # Use vLLM's Flash Attention implementation + from vllm.vllm_flash_attn import flash_attn_func + # Flash Attention expects (batch, seq, heads, head_dim) out = flash_attn_func(query, key, value, softmax_scale=self.scale) out = out.reshape(bsz, q_len, -1) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func + # ROCm Flash Attention expects (batch, seq, heads, head_dim) - out = flash_attn_varlen_func(query, key, value, softmax_scale=self.scale) - out = out.reshape(bsz, q_len, -1) - elif self.attn_backend == _Backend.FLEX_ATTENTION: - # FlexAttention requires specific tensor format - query, key, value = (x.transpose(1, 2) for x in (query, key, value)) - # Use a simple fallback to torch SDPA for now - out = F.scaled_dot_product_attention( - query, key, value, scale=self.scale) - out = out.transpose(1, 2) + out = flash_attn_varlen_func(query, + key, + value, + softmax_scale=self.scale) else: # Fallback to torch SDPA for unsupported backends - query, key, value = (x.transpose(1, 2) for x in (query, key, value)) - out = F.scaled_dot_product_attention( - query, key, value, scale=self.scale) + query, key, value = (x.transpose(1, 2) + for x in (query, key, value)) + out = F.scaled_dot_product_attention(query, + key, + value, + scale=self.scale) out = out.transpose(1, 2) return out.reshape(bsz, q_len, -1) diff --git a/vllm/model_executor/models/intern_vit.py b/vllm/model_executor/models/intern_vit.py index f8a2ac2c9243..643624686e5c 100644 --- a/vllm/model_executor/models/intern_vit.py +++ b/vllm/model_executor/models/intern_vit.py @@ -275,8 +275,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) # Use unified MultiHeadAttention with automatic backend selection - x = self.attn(q, k, v) - x = x.reshape(B, N, -1) + x = self.attn(q, k, v) x = self.proj(x) return x diff --git a/vllm/model_executor/models/interns1_vit.py b/vllm/model_executor/models/interns1_vit.py index 7ea8d85b8717..eb6b685d03dc 100644 --- a/vllm/model_executor/models/interns1_vit.py +++ b/vllm/model_executor/models/interns1_vit.py @@ -15,13 +15,13 @@ from transformers import PretrainedConfig from transformers.utils import torch_int +from vllm.attention.layer import MultiHeadAttention from vllm.model_executor.layers.activation import get_act_fn from vllm.model_executor.layers.layernorm import RMSNorm from vllm.model_executor.layers.linear import (ColumnParallelLinear, RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader -from vllm.attention.layer import MultiHeadAttention NORM2FN = { 'rms_norm': RMSNorm, @@ -205,10 +205,10 @@ def __init__( var_hidden_size=self.embed_dim) self.projection_layer = nn.Linear(self.dummy_dim, self.embed_dim) - + # Use unified MultiHeadAttention with automatic backend selection - self.attn = MultiHeadAttention( - self.num_heads, self.head_dim, self.scale) + self.attn = MultiHeadAttention(self.num_heads, self.head_dim, + self.scale) def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape @@ -221,10 +221,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: B_, N_, H_, D_ = q.shape q = self.q_norm(q.flatten(-2, -1)).view(B_, N_, H_, D_) k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) - + # Use unified MultiHeadAttention with automatic backend selection x = self.attn(q, k, v) - x = x.reshape(B, N, -1) x = self.projection_layer(x) return x diff --git a/vllm/model_executor/models/mllama.py b/vllm/model_executor/models/mllama.py index 7ec8e29dd93c..b628d781af9f 100644 --- a/vllm/model_executor/models/mllama.py +++ b/vllm/model_executor/models/mllama.py @@ -35,6 +35,7 @@ import vllm.distributed.parallel_state as ps from vllm.attention import Attention, AttentionMetadata, AttentionType +from vllm.attention.layer import MultiHeadAttention from vllm.attention.ops.paged_attn import PagedAttention from vllm.attention.selector import _Backend from vllm.config import VllmConfig @@ -53,7 +54,6 @@ from vllm.model_executor.model_loader.weight_utils import ( default_weight_loader, maybe_remap_kv_scale_name) from vllm.model_executor.models.module_mapping import MultiModelKeys - from vllm.model_executor.sampling_metadata import SamplingMetadata from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalEncDecInputs, @@ -71,7 +71,6 @@ from .interfaces import SupportsMultiModal, SupportsV0Only from .llama import LlamaDecoderLayer, LlamaMLP from .utils import AutoWeightsLoader, WeightsMapper, maybe_prefix -from vllm.attention.layer import MultiHeadAttention logger = init_logger(__name__) @@ -518,10 +517,10 @@ def __init__(self, quant_config=quant_config, prefix=f"{prefix}.o_proj", ) - + # Use unified MultiHeadAttention with automatic backend selection - self.attn = MultiHeadAttention( - self.num_local_heads, self.head_dim, 1.0 / math.sqrt(self.head_dim)) + self.attn = MultiHeadAttention(self.num_local_heads, self.head_dim, + 1.0 / math.sqrt(self.head_dim)) def forward( self, @@ -530,11 +529,10 @@ def forward( ) -> torch.Tensor: qkv, _ = self.qkv_proj(hidden_state) q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - + # Use unified MultiHeadAttention with automatic backend selection attn_output = self.attn(q, k, v) - attn_output = attn_output.contiguous() attn_output = attn_output.reshape(attn_output.shape[0], attn_output.shape[1], -1) output, _ = self.o_proj(attn_output) diff --git a/vllm/model_executor/models/pixtral.py b/vllm/model_executor/models/pixtral.py index 49aead0fbb31..997267870262 100644 --- a/vllm/model_executor/models/pixtral.py +++ b/vllm/model_executor/models/pixtral.py @@ -23,6 +23,7 @@ PixtralRotaryEmbedding, apply_rotary_pos_emb, position_ids_in_meshgrid) from transformers.tokenization_utils_base import TextInput +from vllm.attention.layer import MultiHeadAttention from vllm.config import VllmConfig from vllm.distributed import divide, get_tensor_model_parallel_world_size from vllm.model_executor.layers.activation import get_act_and_mul_fn @@ -32,9 +33,7 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader -from vllm.attention.layer import MultiHeadAttention from vllm.model_executor.sampling_metadata import SamplingMetadata - from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalKwargsItems from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalFieldConfig, NestedTensors) @@ -1081,10 +1080,10 @@ def __init__( quant_config=quant_config, prefix=f"{prefix}.o_proj", ) - + # Use unified MultiHeadAttention with automatic backend selection - self.attn = MultiHeadAttention( - self.n_heads, self.head_dim, 1.0 / math.sqrt(self.head_dim)) + self.attn = MultiHeadAttention(self.n_heads, self.head_dim, + 1.0 / math.sqrt(self.head_dim)) def forward( self, @@ -1108,11 +1107,10 @@ def forward( q_reshaped = q.transpose(1, 2).contiguous().view(batch, patches, -1) k_reshaped = k.transpose(1, 2).contiguous().view(batch, patches, -1) v_reshaped = v.transpose(1, 2).contiguous().view(batch, patches, -1) - + # Use unified MultiHeadAttention with automatic backend selection out = self.attn(q_reshaped, k_reshaped, v_reshaped) - out = out.view(batch, patches, self.n_heads * self.head_dim) attn_output, _ = self.o_proj(out) return attn_output, None diff --git a/vllm/model_executor/models/step3_vl.py b/vllm/model_executor/models/step3_vl.py index 977867866dae..a51bb783174a 100644 --- a/vllm/model_executor/models/step3_vl.py +++ b/vllm/model_executor/models/step3_vl.py @@ -16,6 +16,7 @@ from torchvision.transforms.functional import InterpolationMode from transformers import BatchFeature, PretrainedConfig, TensorType +from vllm.attention.layer import MultiHeadAttention from vllm.config import VllmConfig from vllm.distributed import get_tensor_model_parallel_world_size from vllm.model_executor.layers.activation import get_act_fn @@ -25,7 +26,6 @@ RowParallelLinear) from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.sampler import SamplerOutput, get_sampler -from vllm.attention.layer import MultiHeadAttention from vllm.model_executor.sampling_metadata import SamplingMetadata from vllm.multimodal import MULTIMODAL_REGISTRY from vllm.multimodal.inputs import (MultiModalDataDict, MultiModalFieldConfig, @@ -697,10 +697,10 @@ def __init__(self, bias=True, quant_config=quant_config, prefix=prefix) - + # Use unified MultiHeadAttention with automatic backend selection - self.attn = MultiHeadAttention( - self.num_heads, self.head_dim, self.scale) + self.attn = MultiHeadAttention(self.num_heads, self.head_dim, + self.scale) def forward( self, @@ -712,7 +712,7 @@ def forward( # get query proj qkv, _ = self.qkv_proj(hidden_states) q, k, v = qkv.chunk(chunks=3, dim=-1) - + # Use unified MultiHeadAttention with automatic backend selection attn_output = self.attn(q, k, v) From b877af8a49296cf18de765ce39157eea4444d6bb Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Wed, 3 Sep 2025 16:51:42 -0400 Subject: [PATCH 19/28] Delete layer.py out = out.reshape(bsz, q_len, -1) --- vllm/attention/layer.py | 1 - 1 file changed, 1 deletion(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 23032b7d043d..c9d04c31ef62 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -419,7 +419,6 @@ def forward( # Flash Attention expects (batch, seq, heads, head_dim) out = flash_attn_func(query, key, value, softmax_scale=self.scale) - out = out.reshape(bsz, q_len, -1) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func From c861dbab207e36848f3575eac19f5f2357a02af3 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Sat, 6 Sep 2025 04:31:21 -0400 Subject: [PATCH 20/28] Fix reviewer comments: remove unnecessary transpose and improve fallback logic - Fix fallback logic in MultiHeadAttention to raise NotImplementedError for unsupported backends - Remove unnecessary transpose operation for value_states in pixtral.py - Simplify qkv processing logic by removing redundant transpose operations - Apply yapf formatting to ensure code quality --- vllm/attention/layer.py | 12 ++++-------- vllm/model_executor/models/pixtral.py | 8 ++++---- 2 files changed, 8 insertions(+), 12 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index c9d04c31ef62..30cc7b5984be 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -428,14 +428,10 @@ def forward( value, softmax_scale=self.scale) else: - # Fallback to torch SDPA for unsupported backends - query, key, value = (x.transpose(1, 2) - for x in (query, key, value)) - out = F.scaled_dot_product_attention(query, - key, - value, - scale=self.scale) - out = out.transpose(1, 2) + # ViT attention hasn't supported this backend yet + raise NotImplementedError( + f"ViT attention hasn't supported {self.attn_backend} " + f"backend yet.") return out.reshape(bsz, q_len, -1) diff --git a/vllm/model_executor/models/pixtral.py b/vllm/model_executor/models/pixtral.py index 997267870262..03c60721fedb 100644 --- a/vllm/model_executor/models/pixtral.py +++ b/vllm/model_executor/models/pixtral.py @@ -1104,12 +1104,12 @@ def forward( q, k = apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=0) # Reshape for MultiHeadAttention: (batch, seq, hidden_size) - q_reshaped = q.transpose(1, 2).contiguous().view(batch, patches, -1) - k_reshaped = k.transpose(1, 2).contiguous().view(batch, patches, -1) - v_reshaped = v.transpose(1, 2).contiguous().view(batch, patches, -1) + q = q.contiguous().view(batch, patches, -1) + k = k.contiguous().view(batch, patches, -1) + v = v.contiguous().view(batch, patches, -1) # Use unified MultiHeadAttention with automatic backend selection - out = self.attn(q_reshaped, k_reshaped, v_reshaped) + out = self.attn(q, k, v) attn_output, _ = self.o_proj(out) From 88ddf98196b0a72e3572644b81e2bcb968262eb3 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Mon, 8 Sep 2025 16:41:24 +0800 Subject: [PATCH 21/28] revert pixtral and code format Signed-off-by: Isotr0py --- .../models/idefics2_vision_model.py | 2 +- vllm/model_executor/models/intern_vit.py | 11 ++++---- vllm/model_executor/models/pixtral.py | 26 ++++++++++--------- 3 files changed, 20 insertions(+), 19 deletions(-) diff --git a/vllm/model_executor/models/idefics2_vision_model.py b/vllm/model_executor/models/idefics2_vision_model.py index 7c77ed6dfc56..ea5d6f29f6cf 100644 --- a/vllm/model_executor/models/idefics2_vision_model.py +++ b/vllm/model_executor/models/idefics2_vision_model.py @@ -182,7 +182,7 @@ def forward( hidden_states ) # batch_size, q_len, 3 * num_heads_per_partition * head_dim query_states, key_states, value_states = qkv.chunk(3, dim=-1) - + # Use unified MultiHeadAttention implementation out = self.attn(query_states, key_states, value_states) attn_output, _ = self.out_proj(out) diff --git a/vllm/model_executor/models/intern_vit.py b/vllm/model_executor/models/intern_vit.py index 643624686e5c..8e9ab9649bd4 100644 --- a/vllm/model_executor/models/intern_vit.py +++ b/vllm/model_executor/models/intern_vit.py @@ -29,7 +29,6 @@ from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.model_loader.weight_utils import default_weight_loader - NORM2FN = { 'rms_norm': RMSNorm, 'layer_norm': nn.LayerNorm, @@ -255,10 +254,10 @@ def __init__( var_hidden_size=self.embed_dim) self.proj = nn.Linear(self.dummy_dim, self.embed_dim) - + # Use unified MultiHeadAttention with automatic backend selection - self.attn = MultiHeadAttention( - self.num_heads, self.head_dim, self.scale) + self.attn = MultiHeadAttention(self.num_heads, self.head_dim, + self.scale) def forward(self, x: torch.Tensor) -> torch.Tensor: B, N, C = x.shape @@ -273,9 +272,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: B_, N_, H_, D_ = q.shape q = self.q_norm(q.flatten(-2, -1)).view(B_, N_, H_, D_) k = self.k_norm(k.flatten(-2, -1)).view(B_, N_, H_, D_) - + # Use unified MultiHeadAttention with automatic backend selection - x = self.attn(q, k, v) + x = self.attn(q, k, v) x = self.proj(x) return x diff --git a/vllm/model_executor/models/pixtral.py b/vllm/model_executor/models/pixtral.py index 03c60721fedb..e7f5799a8006 100644 --- a/vllm/model_executor/models/pixtral.py +++ b/vllm/model_executor/models/pixtral.py @@ -23,7 +23,6 @@ PixtralRotaryEmbedding, apply_rotary_pos_emb, position_ids_in_meshgrid) from transformers.tokenization_utils_base import TextInput -from vllm.attention.layer import MultiHeadAttention from vllm.config import VllmConfig from vllm.distributed import divide, get_tensor_model_parallel_world_size from vllm.model_executor.layers.activation import get_act_and_mul_fn @@ -1081,10 +1080,6 @@ def __init__( prefix=f"{prefix}.o_proj", ) - # Use unified MultiHeadAttention with automatic backend selection - self.attn = MultiHeadAttention(self.n_heads, self.head_dim, - 1.0 / math.sqrt(self.head_dim)) - def forward( self, hidden_states: torch.Tensor, @@ -1103,14 +1098,21 @@ def forward( cos, sin = position_embeddings q, k = apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=0) - # Reshape for MultiHeadAttention: (batch, seq, hidden_size) - q = q.contiguous().view(batch, patches, -1) - k = k.contiguous().view(batch, patches, -1) - v = v.contiguous().view(batch, patches, -1) - - # Use unified MultiHeadAttention with automatic backend selection - out = self.attn(q, k, v) + if USE_XFORMERS_OPS: + # Transpose q and k back for attention + q = q.transpose(1, 2).contiguous() + k = k.transpose(1, 2).contiguous() + out = xops.memory_efficient_attention(q, + k, + v, + attn_bias=attention_mask) + else: + v = v.transpose(1, 2) + out = nn.functional.scaled_dot_product_attention( + q, k, v, attn_mask=attention_mask) + out = out.transpose(1, 2) + out = out.view(batch, patches, self.n_heads * self.head_dim) attn_output, _ = self.o_proj(out) return attn_output, None From dcb457b54477542b1afca07f8461d659685b0e73 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 8 Sep 2025 22:49:34 -0400 Subject: [PATCH 22/28] retry Fix test_mha_attn.py --- tests/kernels/attention/test_mha_attn.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/kernels/attention/test_mha_attn.py b/tests/kernels/attention/test_mha_attn.py index 53c37554b15a..8a18d5e67def 100644 --- a/tests/kernels/attention/test_mha_attn.py +++ b/tests/kernels/attention/test_mha_attn.py @@ -33,19 +33,19 @@ def test_mha_attn_platform(device: str): torch.set_default_dtype(torch.float16) if device == "cpu": - with patch("vllm.attention.selector.current_platform", CpuPlatform()): + with patch("vllm.platforms.current_platform", CpuPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.TORCH_SDPA elif device == "hip": - with patch("vllm.attention.selector.current_platform", RocmPlatform()): + with patch("vllm.platforms.current_platform", RocmPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.TORCH_SDPA else: - with patch("vllm.attention.selector.current_platform", CudaPlatform()): + with patch("vllm.platforms.current_platform", CudaPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.XFORMERS - with patch("vllm.attention.selector.current_platform", CudaPlatform()): + with patch("vllm.platforms.current_platform", CudaPlatform()): attn = MultiHeadAttention(16, 72, scale=1) assert attn.attn_backend == _Backend.XFORMERS From ad360e5c819a7d0544cdcba8f47a9e2a63b1fbe4 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 8 Sep 2025 22:52:26 -0400 Subject: [PATCH 23/28] Fix test_mha_attn.py --- tests/kernels/attention/test_mha_attn.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/kernels/attention/test_mha_attn.py b/tests/kernels/attention/test_mha_attn.py index 8a18d5e67def..aee172404845 100644 --- a/tests/kernels/attention/test_mha_attn.py +++ b/tests/kernels/attention/test_mha_attn.py @@ -23,6 +23,9 @@ def clear_cache(): """Clear lru cache to ensure each test case runs without caching. """ _cached_get_attn_backend.cache_clear() + # Clear xformers availability cache + import vllm.attention.layer as layer_module + layer_module.USE_XFORMERS_OPS = None @pytest.mark.parametrize("device", ["cpu", "hip", "cuda"]) From 86f1b87cd0474e83a97ebce53f5b1b080272622d Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 8 Sep 2025 22:54:55 -0400 Subject: [PATCH 24/28] Fix --- tests/kernels/attention/test_mha_attn.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/kernels/attention/test_mha_attn.py b/tests/kernels/attention/test_mha_attn.py index aee172404845..3b588059549c 100644 --- a/tests/kernels/attention/test_mha_attn.py +++ b/tests/kernels/attention/test_mha_attn.py @@ -36,19 +36,27 @@ def test_mha_attn_platform(device: str): torch.set_default_dtype(torch.float16) if device == "cpu": - with patch("vllm.platforms.current_platform", CpuPlatform()): + with patch("vllm.attention.selector.current_platform", + CpuPlatform()), \ + patch("vllm.platforms.current_platform", CpuPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.TORCH_SDPA elif device == "hip": - with patch("vllm.platforms.current_platform", RocmPlatform()): + with patch("vllm.attention.selector.current_platform", + RocmPlatform()), \ + patch("vllm.platforms.current_platform", RocmPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.TORCH_SDPA else: - with patch("vllm.platforms.current_platform", CudaPlatform()): + with patch("vllm.attention.selector.current_platform", + CudaPlatform()), \ + patch("vllm.platforms.current_platform", CudaPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.XFORMERS - with patch("vllm.platforms.current_platform", CudaPlatform()): + with patch("vllm.attention.selector.current_platform", + CudaPlatform()), \ + patch("vllm.platforms.current_platform", CudaPlatform()): attn = MultiHeadAttention(16, 72, scale=1) assert attn.attn_backend == _Backend.XFORMERS From 01a5410cba39867ec662e49ccb6241de4cb7fa00 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 8 Sep 2025 23:03:58 -0400 Subject: [PATCH 25/28] Fix: --- tests/kernels/attention/test_mha_attn.py | 4 ++-- vllm/attention/layer.py | 1 + vllm/platforms/interface.py | 1 + 3 files changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/kernels/attention/test_mha_attn.py b/tests/kernels/attention/test_mha_attn.py index 3b588059549c..c39b818eb958 100644 --- a/tests/kernels/attention/test_mha_attn.py +++ b/tests/kernels/attention/test_mha_attn.py @@ -40,13 +40,13 @@ def test_mha_attn_platform(device: str): CpuPlatform()), \ patch("vllm.platforms.current_platform", CpuPlatform()): attn = MultiHeadAttention(16, 64, scale=1) - assert attn.attn_backend == _Backend.TORCH_SDPA + assert attn.attn_backend == _Backend.TORCH_SDPA_VLLM_V1 elif device == "hip": with patch("vllm.attention.selector.current_platform", RocmPlatform()), \ patch("vllm.platforms.current_platform", RocmPlatform()): attn = MultiHeadAttention(16, 64, scale=1) - assert attn.attn_backend == _Backend.TORCH_SDPA + assert attn.attn_backend == _Backend.TORCH_SDPA_VLLM_V1 else: with patch("vllm.attention.selector.current_platform", CudaPlatform()), \ diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 30cc7b5984be..030d165ef6c6 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -362,6 +362,7 @@ def __init__( else: self.attn_backend = backend if backend in { _Backend.TORCH_SDPA, + _Backend.TORCH_SDPA_VLLM_V1, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, _Backend.ROCM_AITER_FA, diff --git a/vllm/platforms/interface.py b/vllm/platforms/interface.py index fdd3764d2c35..0cea49eece42 100644 --- a/vllm/platforms/interface.py +++ b/vllm/platforms/interface.py @@ -48,6 +48,7 @@ class _Backend(enum.Enum): ROCM_AITER_MLA_VLLM_V1 = enum.auto() ROCM_AITER_FA = enum.auto() # used for ViT attn backend TORCH_SDPA = enum.auto() + TORCH_SDPA_VLLM_V1 = enum.auto() FLASHINFER = enum.auto() FLASHINFER_VLLM_V1 = enum.auto() TRITON_MLA = enum.auto() # Supported by V1 From d03c4da97c3baa03f102bf2132d623eb534bcd6f Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 8 Sep 2025 23:26:36 -0400 Subject: [PATCH 26/28] Fix --- tests/kernels/attention/test_mha_attn.py | 2 +- vllm/attention/layer.py | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/kernels/attention/test_mha_attn.py b/tests/kernels/attention/test_mha_attn.py index c39b818eb958..c0462aad1a1c 100644 --- a/tests/kernels/attention/test_mha_attn.py +++ b/tests/kernels/attention/test_mha_attn.py @@ -46,7 +46,7 @@ def test_mha_attn_platform(device: str): RocmPlatform()), \ patch("vllm.platforms.current_platform", RocmPlatform()): attn = MultiHeadAttention(16, 64, scale=1) - assert attn.attn_backend == _Backend.TORCH_SDPA_VLLM_V1 + assert attn.attn_backend == _Backend.TORCH_SDPA else: with patch("vllm.attention.selector.current_platform", CudaPlatform()), \ diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 030d165ef6c6..9c2ded1597f8 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -363,6 +363,7 @@ def __init__( self.attn_backend = backend if backend in { _Backend.TORCH_SDPA, _Backend.TORCH_SDPA_VLLM_V1, + _Backend.TRITON_MLA_VLLM_V1, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, _Backend.ROCM_AITER_FA, From 77f76547c7b299393b8082880a009cb950945c95 Mon Sep 17 00:00:00 2001 From: baonudesifeizhai Date: Mon, 8 Sep 2025 23:32:29 -0400 Subject: [PATCH 27/28] Fix --- tests/kernels/attention/test_mha_attn.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/kernels/attention/test_mha_attn.py b/tests/kernels/attention/test_mha_attn.py index c0462aad1a1c..c01ea32994da 100644 --- a/tests/kernels/attention/test_mha_attn.py +++ b/tests/kernels/attention/test_mha_attn.py @@ -44,7 +44,8 @@ def test_mha_attn_platform(device: str): elif device == "hip": with patch("vllm.attention.selector.current_platform", RocmPlatform()), \ - patch("vllm.platforms.current_platform", RocmPlatform()): + patch("vllm.platforms.current_platform", RocmPlatform()), \ + patch("vllm.attention.layer.current_platform", RocmPlatform()): attn = MultiHeadAttention(16, 64, scale=1) assert attn.attn_backend == _Backend.TORCH_SDPA else: From 8f16d0209ebc078069c7a2e74965be6b2e6029f7 Mon Sep 17 00:00:00 2001 From: Isotr0py Date: Wed, 10 Sep 2025 13:19:31 +0800 Subject: [PATCH 28/28] remove never used ViT FA Signed-off-by: Isotr0py --- vllm/attention/layer.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/vllm/attention/layer.py b/vllm/attention/layer.py index 9c2ded1597f8..be4dc3eb3c0d 100644 --- a/vllm/attention/layer.py +++ b/vllm/attention/layer.py @@ -363,7 +363,6 @@ def __init__( self.attn_backend = backend if backend in { _Backend.TORCH_SDPA, _Backend.TORCH_SDPA_VLLM_V1, - _Backend.TRITON_MLA_VLLM_V1, _Backend.XFORMERS, _Backend.PALLAS_VLLM_V1, _Backend.ROCM_AITER_FA, @@ -414,13 +413,6 @@ def forward( from torch_xla.experimental.custom_kernel import flash_attention out = flash_attention(query, key, value, sm_scale=self.scale) out = out.transpose(1, 2) - elif self.attn_backend in (_Backend.FLASH_ATTN, - _Backend.FLASH_ATTN_VLLM_V1): - # Use vLLM's Flash Attention implementation - from vllm.vllm_flash_attn import flash_attn_func - - # Flash Attention expects (batch, seq, heads, head_dim) - out = flash_attn_func(query, key, value, softmax_scale=self.scale) elif self.attn_backend == _Backend.ROCM_AITER_FA: from aiter import flash_attn_varlen_func