Skip to content

vllm.v1.attention.backends.rocm_segmented_attn

ROCm attention backend for token-major segmented Triton attention.

Classes:

RocmSegmentedAttentionBackend

Bases: RocmAttentionBackend

Explicit token-major ROCm backend optimized for segmented prefill.

Source code in vllm/v1/attention/backends/rocm_segmented_attn.py
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