class RocmSegmentedAttentionBackend(RocmAttentionBackend):
"""Explicit token-major ROCm backend optimized for segmented prefill."""
supported_dtypes: ClassVar[list[torch.dtype]] = [
torch.float16,
torch.bfloat16,
]
supported_kv_cache_dtypes: ClassVar[list[CacheDType]] = [
"auto",
"float16",
"bfloat16",
"fp8",
"fp8_e4m3",
]
@staticmethod
def get_name() -> str:
return "ROCM_SEGMENTED_ATTN"
@staticmethod
def get_impl_cls() -> type["RocmSegmentedAttentionImpl"]:
return RocmSegmentedAttentionImpl
@staticmethod
def get_builder_cls() -> type["RocmSegmentedAttentionMetadataBuilder"]:
return RocmSegmentedAttentionMetadataBuilder
@classmethod
def get_supported_head_sizes(cls) -> list[int]:
return [64, 128, 256]
@classmethod
def supports_attn_type(cls, attn_type: str) -> bool:
return attn_type == AttentionType.DECODER
@classmethod
def supports_non_causal(cls) -> bool:
return True
@classmethod
def supports_sliding_window(cls) -> bool:
return True
@classmethod
def supports_sink(cls) -> bool:
return True
@classmethod
def supports_mm_prefix(cls) -> bool:
return False
@classmethod
def supports_combination(
cls,
head_size: int,
dtype: torch.dtype,
kv_cache_dtype: CacheDType | None,
block_size: int | None,
use_mla: bool,
has_sink: bool,
use_sparse: bool,
use_mm_prefix: bool,
device_capability: "DeviceCapability",
) -> str | None:
del (
head_size,
dtype,
block_size,
use_mla,
has_sink,
use_sparse,
use_mm_prefix,
device_capability,
)
if not is_rdna():
return "ROCM_SEGMENTED_ATTN requires AMD RDNA GPUs on ROCm"
from vllm.platforms.rocm import on_gfx12x
if kv_cache_dtype in ("fp8", "fp8_e4m3") and not on_gfx12x():
return "FP8 segmented attention requires gfx12"
return None