[Attention] Update attention imports (#29540)

Signed-off-by: Matthew Bonanni <mbonanni@redhat.com>
This commit is contained in:
Matthew Bonanni
2025-11-27 11:19:09 -05:00
committed by GitHub
parent cd007a53b4
commit fc1d8be3dc
38 changed files with 63 additions and 126 deletions

View File

@@ -139,14 +139,13 @@ def test_standard_attention_backend_selection(
import importlib
import vllm.envs as envs
from vllm.attention.backends.registry import _Backend
importlib.reload(envs)
# Convert string backend to enum if provided
backend_enum = None
if selected_backend:
backend_enum = getattr(_Backend, selected_backend)
backend_enum = getattr(AttentionBackendEnum, selected_backend)
# Get the backend class path
from vllm.platforms.rocm import RocmPlatform
@@ -253,7 +252,6 @@ def test_mla_backend_selection(
import importlib
import vllm.envs as envs
from vllm.attention.backends.registry import _Backend
importlib.reload(envs)
@@ -269,7 +267,7 @@ def test_mla_backend_selection(
# Convert string backend to enum if provided
backend_enum = None
if selected_backend:
backend_enum = getattr(_Backend, selected_backend)
backend_enum = getattr(AttentionBackendEnum, selected_backend)
from vllm.platforms.rocm import RocmPlatform
@@ -301,7 +299,6 @@ def test_mla_backend_selection(
def test_aiter_fa_requires_gfx9(mock_vllm_config):
"""Test that ROCM_AITER_FA requires gfx9 architecture."""
from vllm.attention.backends.registry import _Backend
from vllm.platforms.rocm import RocmPlatform
# Mock on_gfx9 to return False
@@ -313,7 +310,7 @@ def test_aiter_fa_requires_gfx9(mock_vllm_config):
),
):
RocmPlatform.get_attn_backend_cls(
selected_backend=_Backend.ROCM_AITER_FA,
selected_backend=AttentionBackendEnum.ROCM_AITER_FA,
head_size=128,
dtype=torch.float16,
kv_cache_dtype="auto",

View File

@@ -14,6 +14,7 @@ from unittest.mock import patch
import pytest
from vllm.attention.backends.abstract import AttentionMetadata
from vllm.distributed.kv_transfer.kv_connector.factory import KVConnectorFactory
from vllm.distributed.kv_transfer.kv_connector.v1 import (
KVConnectorBase_V1,
@@ -24,7 +25,6 @@ from vllm.v1.core.sched.output import SchedulerOutput
from .utils import create_scheduler, create_vllm_config
if TYPE_CHECKING:
from vllm.attention.backends.abstract import AttentionMetadata
from vllm.config import VllmConfig
from vllm.forward_context import ForwardContext
from vllm.v1.core.kv_cache_manager import KVCacheBlocks
@@ -68,7 +68,7 @@ class OldStyleTestConnector(KVConnectorBase_V1):
self,
layer_name: str,
kv_layer,
attn_metadata: "AttentionMetadata",
attn_metadata: AttentionMetadata,
**kwargs,
) -> None:
pass
@@ -119,7 +119,7 @@ class NewStyleTestConnector(KVConnectorBase_V1):
self,
layer_name: str,
kv_layer,
attn_metadata: "AttentionMetadata",
attn_metadata: AttentionMetadata,
**kwargs,
) -> None:
pass