class RocmPlatform(Platform):
_enum = PlatformEnum.ROCM
device_name: str = "rocm"
device_type: str = "cuda"
dispatch_key: str = "CUDA"
ray_device_key: str = "GPU"
dist_backend: str = "nccl"
# rocm shares the same device control env var as CUDA
device_control_env_var: str = "CUDA_VISIBLE_DEVICES"
# Set in pre_register_and_update, so it exists only on the driver; Ray
# workers are separate processes and copy env vars by allowlist.
additional_env_vars: list[str] = ["GPU_PINNED_MIN_XFER_SIZE"]
ray_noset_device_env_vars: list[str] = [
"RAY_EXPERIMENTAL_NOSET_HIP_VISIBLE_DEVICES",
"RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES",
"RAY_EXPERIMENTAL_NOSET_ROCR_VISIBLE_DEVICES",
]
supported_quantization: list[str] = [
"awq",
"auto_awq",
"awq_marlin", # will be overwritten with awq
"gptq",
"auto_gptq",
"fp8",
"deepseek_v4_fp8",
"compressed-tensors",
"fbgemm_fp8",
"inc",
"quark",
"mxfp4",
"mxfp8",
"torchao",
"modelopt",
"modelopt_fp4",
"modelopt_mxfp8",
"modelopt_mixed",
"fp8_per_tensor",
"fp8_per_block",
"fp8_per_channel",
"online",
"gpt_oss_mxfp4",
]
@classmethod
def import_kernels(cls) -> None:
"""Import ROCm-specific kernels."""
super().import_kernels()
import contextlib
# Import ROCm-specific extension
with contextlib.suppress(ImportError):
import vllm._rocm_C # noqa: F401
@classmethod
def check_runner_kv_caches_multi_layer(cls) -> None:
pass
@classmethod
def is_pin_memory_available(cls) -> bool:
if in_wsl():
version = _get_wsl_kernel_version()
if version is None or version < (4, 19, 121):
# warning_once() causes a circular import on WSL, see #48397.
logger.warning(
"Using 'pin_memory=False' as WSL is detected and the "
"WSL2 kernel version is below 4.19.121. This may slow "
"down performance. Please run `wsl --update`."
)
return False
return True
@classmethod
def get_valid_backends(
cls,
device_capability: DeviceCapability,
attn_selector_config: "AttentionSelectorConfig",
num_heads: int | None = None,
) -> tuple[
list[tuple["AttentionBackendEnum", int]],
dict["AttentionBackendEnum", list[str]],
]:
valid_backends_priorities = []
invalid_reasons = {}
backend_priorities = _get_backend_priorities(
attn_selector_config.use_mla,
attn_selector_config.use_sparse,
attn_selector_config.use_kv_connector,
)
from vllm.config import get_current_vllm_config_or_none
vllm_config = get_current_vllm_config_or_none()
is_encoder_decoder = (
getattr(getattr(vllm_config, "model_config", None), "attn_type", None)
== "encoder_decoder"
)
# ROCM_ATTN still uses a legacy attention layout (KV is the outer
# dimension) that is incompatible with the encoder backend layouts. The
# encoder and decoder need the layouts to match. This is currently
# enforced implicitly.
# TODO: Make this explicit in the selector in a future PR.
if is_encoder_decoder and AttentionBackendEnum.ROCM_ATTN in backend_priorities:
backend_priorities.remove(AttentionBackendEnum.ROCM_ATTN)
is_turboquant_run = _uses_turboquant(vllm_config)
for priority, backend in enumerate(backend_priorities):
try:
invalid_reasons_i = _get_invalid_reasons(
backend.get_class(),
device_capability,
attn_selector_config,
is_turboquant_run=is_turboquant_run,
)
except ImportError:
invalid_reasons_i = ["ImportError"]
if invalid_reasons_i:
invalid_reasons[backend] = invalid_reasons_i
else:
valid_backends_priorities.append((backend, priority))
return valid_backends_priorities, invalid_reasons
@classmethod
def get_attn_backend_cls(
cls,
selected_backend: "AttentionBackendEnum",
attn_selector_config: "AttentionSelectorConfig",
num_heads: int | None = None,
) -> str:
device_capability = cls.get_device_capability()
assert device_capability is not None
# First try checking just the selected backend, if there is one.
if selected_backend is not None:
# Keep lazy: vllm.config imports current_platform during initialization.
from vllm.config import get_current_vllm_config_or_none
is_turboquant_run = _uses_turboquant(get_current_vllm_config_or_none())
try:
sel_invalid_reasons = _get_invalid_reasons(
selected_backend.get_class(),
device_capability,
attn_selector_config,
is_turboquant_run=is_turboquant_run,
)
except ImportError:
sel_invalid_reasons = ["ImportError"]
if not sel_invalid_reasons:
logger.info_once(
"Using %s backend (selected via --attention-backend).",
selected_backend.name,
)
return selected_backend.get_path()
# Only tolerate the mismatch when turboquant is in play: boundary
# layers keep the native dtype while every other layer needs
# TURBOQUANT, so no single --attention-backend can serve every layer.
# For any other dtype the selection is genuinely invalid -> fail loud.
kv_dtype = attn_selector_config.kv_cache_dtype
layer_is_turboquant = kv_dtype is not None and str(kv_dtype).startswith(
"turboquant"
)
is_turboquant_fallback = (
is_turboquant_run or layer_is_turboquant
) and sel_invalid_reasons in (
[_KV_CACHE_DTYPE_REASON],
[_TURBOQUANT_LAYOUT_REASON],
)
if not is_turboquant_fallback:
raise ValueError(
f"Selected backend {selected_backend} is not valid for "
f"this configuration. Reason: {sel_invalid_reasons}"
)
# NOTE: pass a str (not the list) -- info_once hashes its args.
logger.info_once(
"Selected backend %s is incompatible with this layer (%s) of "
"the turboquant run; using the auto-selected per-layer backend. "
"Reason: %s",
selected_backend.name,
attn_selector_config.attn_type,
str(sel_invalid_reasons),
)
# No selected backend or the selected backend is invalid,
# so we try finding a valid backend.
valid_backends_priorities, invalid_reasons = cls.get_valid_backends(
device_capability=device_capability,
attn_selector_config=attn_selector_config,
num_heads=num_heads,
)
reasons_str = (
"{"
+ ", ".join(
f"{backend.name}: [{', '.join(reasons)}]"
for backend, reasons in invalid_reasons.items()
)
+ "}"
)
config_str = attn_selector_config.__repr__()
logger.debug_once(
f"Some attention backends are not valid for {cls.device_name} with "
f"{config_str}. Reasons: {reasons_str}."
)
if len(valid_backends_priorities) == 0:
# If a backend rejected the requested kv-cache dtype, list the
# dtypes it does accept so the limitation is discoverable.
supported = sorted(
{
dt
for backend, reasons in invalid_reasons.items()
if any("kv_cache_dtype" in r for r in reasons)
for dt in backend.get_class().supported_kv_cache_dtypes
}
)
hint = (
f" Supported kv_cache_dtype values: {', '.join(supported)}."
if supported
else ""
)
raise ValueError(
f"No valid attention backend found for {cls.device_name} "
f"with {config_str}. Reasons: {reasons_str}.{hint}"
)
# We have found some valid backends. Select the one with the
# highest priority.
sorted_indices = sorted(
range(len(valid_backends_priorities)),
key=lambda i: valid_backends_priorities[i][1],
)
selected_index = sorted_indices[0]
selected_backend = valid_backends_priorities[selected_index][0]
valid_str = (
"[" + ", ".join(f"'{b[0].name}'" for b in valid_backends_priorities) + "]"
)
if invalid_reasons:
rejected_str = ", ".join(b.name for b in invalid_reasons)
logger.info(
"Found incompatible backend(s) [%s] with %s. "
"Overriding with %s out of potential backends: %s.",
rejected_str,
attn_selector_config.attn_type,
selected_backend.name,
valid_str,
)
else:
logger.info_once(
"Using %s backend out of potential backends: %s.",
selected_backend.name,
valid_str,
)
return selected_backend.get_path()
@classmethod
def get_supported_vit_attn_backends(cls) -> list["AttentionBackendEnum"]:
return [
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.ROCM_AITER_FA,
AttentionBackendEnum.TRITON_ATTN,
AttentionBackendEnum.TORCH_SDPA,
]
@classmethod
def get_vit_attn_backend(
cls,
head_size: int,
dtype: torch.dtype,
backend: "AttentionBackendEnum | None" = None,
) -> "AttentionBackendEnum":
if backend is not None:
assert backend in cls.get_supported_vit_attn_backends(), (
f"Backend {backend} is not supported for vit attention. "
f"Supported backends are: {cls.get_supported_vit_attn_backends()}"
)
logger.info_once(f"Using backend {backend} for vit attention")
return backend
from importlib.util import find_spec
from vllm._aiter_ops import rocm_aiter_ops
if rocm_aiter_ops.is_mha_enabled() and on_cdna():
logger.info_once("Using AITER Flash Attention backend for ViT model.")
return AttentionBackendEnum.ROCM_AITER_FA
if (
on_cdna()
and find_spec("flash_attn") is not None
and (dtype == torch.float16 or dtype == torch.bfloat16)
):
logger.info_once("Using Flash Attention backend for ViT model.")
return AttentionBackendEnum.FLASH_ATTN
# RDNA3/RDNA4 (gfx11xx/gfx12xx): Use Flash Attention Triton backend
if (
on_gfx1x()
and flash_attn_triton_available()
and (dtype == torch.float16 or dtype == torch.bfloat16)
):
logger.info_once(
"Using Flash Attention (Triton backend) for ViT model on RDNA."
)
return AttentionBackendEnum.FLASH_ATTN
logger.info_once("Using Torch SDPA backend for ViT model.")
return AttentionBackendEnum.TORCH_SDPA
@classmethod
def set_device(cls, device: torch.device) -> None:
"""Set the device for the current platform."""
torch.cuda.set_device(device)
@classmethod
def manual_seed_all(cls, seed: int) -> None:
torch.cuda.manual_seed_all(seed)
@classmethod
@lru_cache(maxsize=8)
def get_device_capability(cls, device_id: int = 0) -> DeviceCapability | None:
cap = _capability_from_gcn_arch(_GCN_ARCH)
if cap is not None:
return DeviceCapability(major=cap[0], minor=cap[1])
logger.warning_once(
"Could not derive device capability from GCN arch '%s', "
"falling back to torch.cuda (this will initialize CUDA).",
_GCN_ARCH,
)
major, minor = torch.cuda.get_device_capability(device_id)
return DeviceCapability(major=major, minor=minor)
@classmethod
@with_amdsmi_context
def is_fully_connected(cls, physical_device_ids: list[int]) -> bool:
"""Query if the set of gpus are fully connected by xgmi (1 hop)."""
handles = [amdsmi_get_processor_handles()[i] for i in physical_device_ids]
for i, handle in enumerate(handles):
for j, peer_handle in enumerate(handles):
if i < j:
try:
link_type = amdsmi_topo_get_link_type(handle, peer_handle)
# type is 2 for XGMI
if link_type["hops"] != 1 or link_type["type"] != 2:
return False
except AmdSmiException as error:
logger.error("AMD 1 hop XGMI detection failed.", exc_info=error)
return False
return True
@classmethod
@with_amdsmi_context
@lru_cache(maxsize=8)
def get_device_name(cls, device_id: int = 0) -> str:
physical_device_id = cls.device_id_to_physical_device_id(device_id)
handle = amdsmi_get_processor_handles()[physical_device_id]
asic_info = amdsmi_get_gpu_asic_info(handle)
asic_info_device_id: str = asic_info["device_id"]
if asic_info_device_id in _ROCM_DEVICE_ID_NAME_MAP:
return _ROCM_DEVICE_ID_NAME_MAP[asic_info_device_id]
return asic_info["market_name"]
@classmethod
@with_amdsmi_context
def get_device_uuid(cls, device_id: int = 0) -> str:
try:
device = amdsmi_get_processor_handles()[device_id]
except AmdSmiException as error:
logger.error("amdsmi device query failed ", exc_info=error)
return ""
try:
device_uuid = amdsmi_get_gpu_device_uuid(device)
except AmdSmiException as error:
logger.error("amdsmi device uuid query failed ", exc_info=error)
return device_uuid
@classmethod
def get_device_total_memory(cls, device_id: int = 0) -> int:
# Query total VRAM via amdsmi so we don't initialize a HIP context in
# the calling process. torch.cuda.get_device_properties() creates a
# HIP context, which makes vLLM fall back from `fork` to `spawn` for
# worker processes. Keeping this query context-free preserves `fork`
# where it is otherwise valid (e.g. out-of-tree models registered in
# the parent process).
try:
physical_device_id = cls.device_id_to_physical_device_id(device_id)
return _query_total_memory_from_amdsmi(physical_device_id)
except Exception as e:
logger.debug("Failed to get total memory via amdsmi: %s", e)
logger.warning_once(
"Failed to get total memory via amdsmi, falling back to "
"torch.cuda. This will initialize CUDA."
)
return torch.cuda.get_device_properties(device_id).total_memory
@classmethod
def pre_register_and_update(
cls, parser: "FlexibleArgumentParser | None" = None
) -> None:
# Keep mmap'd weight pages on the HIP staging path: above this
# threshold the runtime registers the pageable source instead, and each
# registration's MMU notifier makes KFD suspend our queues. In KB, so
# 4 GiB.
os.environ.setdefault("GPU_PINNED_MIN_XFER_SIZE", str(4 * 1024 * 1024))
@classmethod
def apply_config_platform_defaults(cls, vllm_config: "VllmConfig") -> None:
from vllm._aiter_ops import rocm_aiter_ops
compilation_config = vllm_config.compilation_config
use_aiter_fused_moe = rocm_aiter_ops.is_fused_moe_enabled()
use_aiter_fp8_linear = rocm_aiter_ops.is_linear_fp8_enabled()
use_aiter_fused_se = rocm_aiter_ops.is_fusion_moe_shared_experts_enabled()
if use_aiter_fp8_linear and "-quant_fp8" not in compilation_config.custom_ops:
compilation_config.custom_ops.append("+quant_fp8")
if use_aiter_fused_se and "-grouped_topk" in compilation_config.custom_ops:
logger.warning_once(
"VLLM_ROCM_USE_AITER_FUSION_SHARED_EXPERTS is enabled, which "
"requires the 'grouped_topk' custom op. Overriding the "
"user-provided '-grouped_topk'."
)
compilation_config.custom_ops.remove("-grouped_topk")
# Ensure grouped_topk is always enabled when using AITER if
# its not disabled by user
if (
use_aiter_fused_moe
and "+grouped_topk" not in compilation_config.custom_ops
and "-grouped_topk" not in compilation_config.custom_ops
):
compilation_config.custom_ops.append("+grouped_topk")
# Default dispatch to rocm's sparse_attn_indexer implementation
compilation_config.custom_ops.append("+sparse_attn_indexer")
@classmethod
def check_and_update_config(cls, vllm_config: "VllmConfig") -> None:
from vllm.config.compilation import CUDAGraphMode
compilation_config = vllm_config.compilation_config
parallel_config = vllm_config.parallel_config
if (
compilation_config.cudagraph_mode.has_full_cudagraphs()
and parallel_config.prefill_context_parallel_size > 1
):
# prefill context parallel do not support full cudagraphs
logger.warning_once(
"Prefill context parallel (PCP) is enabled, which is "
"incompatible with full CUDA graphs. "
"Overriding cudagraph_mode to PIECEWISE."
)
compilation_config.cudagraph_mode = CUDAGraphMode.PIECEWISE
if parallel_config.worker_cls == "auto":
parallel_config.worker_cls = "vllm.v1.worker.gpu_worker.Worker"
model_config = vllm_config.model_config
scheduler_config = vllm_config.scheduler_config
# Note: model_config may be None during testing
if (
model_config is not None
and model_config.is_mm_prefix_lm
and scheduler_config.is_multimodal_model
and not scheduler_config.disable_chunked_mm_input
):
logger.warning_once(
"Forcing --disable_chunked_mm_input for models "
"with multimodal-bidirectional attention."
)
scheduler_config.disable_chunked_mm_input = True
@classmethod
def verify_model_arch(cls, model_arch: str) -> None:
if model_arch in _ROCM_UNSUPPORTED_MODELS:
raise ValueError(
f"Model architecture '{model_arch}' is not supported by ROCm for now."
)
if model_arch in _ROCM_PARTIALLY_SUPPORTED_MODELS:
msg = _ROCM_PARTIALLY_SUPPORTED_MODELS[model_arch]
logger.warning(
"Model architecture '%s' is partially supported by ROCm: %s",
model_arch,
msg,
)
@classmethod
def verify_quantization(cls, quant: str) -> None:
super().verify_quantization(quant)
if quant == "awq" and not envs.VLLM_USE_TRITON_AWQ:
logger.warning(
"Using AWQ quantization with ROCm, but VLLM_USE_TRITON_AWQ"
" is not set, enabling VLLM_USE_TRITON_AWQ."
)
os.environ["VLLM_USE_TRITON_AWQ"] = "1"
@classmethod
def get_punica_wrapper(cls) -> str:
return "vllm.lora.punica_wrapper.punica_gpu.PunicaWrapperGPU"
@classmethod
def get_current_memory_usage(
cls, device: torch.types.Device | None = None
) -> float:
torch.cuda.empty_cache()
torch.cuda.reset_peak_memory_stats(device)
return torch.cuda.max_memory_allocated(device)
@classmethod
def get_device_communicator_cls(cls) -> str:
return (
"vllm.distributed.device_communicators.cuda_communicator.CudaCommunicator" # noqa
)
@classmethod
def supports_mx(cls) -> bool:
return any(gfx in _GCN_ARCH for gfx in ["gfx95", "gfx1250"])
@classmethod
def supports_fp8(cls) -> bool:
return on_cdna() or on_rdna4()
@classmethod
def is_fp8_fnuz(cls) -> bool:
# only device 0 is checked, this assumes MI300 platforms are homogeneous
return "gfx94" in _GCN_ARCH
@classmethod
def fp8_dtype(cls) -> torch.dtype:
if cls.is_fp8_fnuz():
return torch.float8_e4m3fnuz
else:
return torch.float8_e4m3fn
@classmethod
def use_custom_allreduce(cls) -> bool:
# We only enable custom allreduce for MI300 series
return any(gfx in _GCN_ARCH for gfx in ["gfx94", "gfx95"])
@classmethod
def opaque_attention_op(cls) -> bool:
return True
@classmethod
def is_navi(cls) -> bool:
return "gfx1" in _GCN_ARCH
@classmethod
def enable_multi_stream_overlap(
cls,
aux_stream_list: list[torch.cuda.Stream] | None,
attn_metadata: object,
) -> bool:
"""ROCm multi-stream gates: streams and capture region.
Dict metadata marks piecewise cudagraph, whose eager breaks rebuild
the attention inputs on the owning stream. Forking side streams
there would rely on runtime HIP event sync, which is unreliable in
this overlap on ROCm (event waits can hang), so multi-stream only
runs where the fork/join becomes static graph edges: inside capture,
or with non-dict metadata (full cudagraph or the profile run), which
has no eager breaks.
"""
return aux_stream_list is not None and (
torch.cuda.is_current_stream_capturing()
or not isinstance(attn_metadata, dict)
)
@classmethod
def get_static_graph_wrapper_cls(cls) -> str:
return "vllm.compilation.cuda_graph.CUDAGraphWrapper"
@classmethod
def stateless_init_device_torch_dist_pg(
cls,
backend: str,
prefix_store: PrefixStore,
group_rank: int,
group_size: int,
timeout: timedelta,
) -> ProcessGroup:
assert is_nccl_available()
pg: ProcessGroup = ProcessGroup(
prefix_store,
group_rank,
group_size,
)
from torch.distributed.distributed_c10d import ProcessGroupNCCL
backend_options = ProcessGroupNCCL.Options()
backend_options._timeout = timeout
backend_class = ProcessGroupNCCL(
prefix_store, group_rank, group_size, backend_options
)
backend_type = ProcessGroup.BackendType.NCCL
device = torch.device("cuda")
pg._set_default_backend(backend_type)
backend_class._set_sequence_number_for_group()
pg._register_backend(device, backend_type, backend_class)
return pg
@classmethod
def device_count(cls) -> int:
return _rocm_device_count_stateless(getattr(envs, cls.device_control_env_var))
@classmethod
def check_if_supports_dtype(cls, dtype: torch.dtype):
if dtype == torch.bfloat16: # noqa: SIM102
if not cls.has_device_capability(80):
capability = cls.get_device_capability()
gpu_name = cls.get_device_name()
if capability is None:
compute_str = "does not have a compute capability"
else:
version_str = capability.as_version_str()
compute_str = f"has compute capability {version_str}"
raise ValueError(
"Bfloat16 is only supported on GPUs "
"with compute capability of at least 8.0. "
f"Your {gpu_name} GPU {compute_str}. "
"You can use float16 instead by explicitly setting the "
"`dtype` flag in CLI, for example: --dtype=half."
)
@classmethod
def insert_blocks_to_device(
cls,
src_cache: torch.Tensor,
dst_cache: torch.Tensor,
src_block_indices: torch.Tensor,
dst_block_indices: torch.Tensor,
) -> None:
"""Copy blocks from src_cache to dst_cache on GPU."""
_src_cache = src_cache[src_block_indices]
dst_cache[dst_block_indices] = _src_cache.to(dst_cache.device)
@classmethod
def swap_out_blocks_to_host(
cls,
src_cache: torch.Tensor,
dst_cache: torch.Tensor,
src_block_indices: torch.Tensor,
dst_block_indices: torch.Tensor,
) -> None:
"""Copy blocks from GPU to host (CPU)."""
_src_cache = src_cache[src_block_indices]
dst_cache[dst_block_indices] = _src_cache.cpu()
@classmethod
def support_hybrid_kv_cache(cls) -> bool:
return True
@classmethod
def support_static_graph_mode(cls) -> bool:
return True
@classmethod
def num_compute_units(cls, device_id: int = 0) -> int:
return torch.cuda.get_device_properties(device_id).multi_processor_count
@classmethod
def use_custom_op_collectives(cls) -> bool:
return True
@classmethod
def get_default_ir_op_priority(
cls, vllm_config: "VllmConfig"
) -> "IrOpPriorityConfig":
from vllm.config.compilation import CompilationMode, CUDAGraphMode
from vllm.config.kernel import IrOpPriorityConfig
# Native used by default when compiling,
# use vllm_c kernels where available when no codegen
# TODO(luka/TJ) use aiter, vllm_c, native by default on ROCm
cc = vllm_config.compilation_config
using_inductor = cc.backend == "inductor" and cc.mode != CompilationMode.NONE
default = ["native"] if using_inductor else ["vllm_c", "native"]
# Aiter rms norm perform best when CUDA Graph capture is enabled.
# TODO(luka/TJ) remove env vars completely
if (
cc.cudagraph_mode != CUDAGraphMode.NONE
and envs.VLLM_ROCM_USE_AITER
and envs.VLLM_ROCM_USE_AITER_RMSNORM
and not on_rdna4()
):
rms_norm = ["aiter"] + default
else:
rms_norm = default
return IrOpPriorityConfig.with_default(
default,
rms_norm=rms_norm,
fused_add_rms_norm=rms_norm,
gelu_and_mul_sparse=["native"],
)
@classmethod
@with_amdsmi_context
def get_all_device_numa_nodes(cls) -> list[int] | None:
"""Get NUMA nodes for all visible GPU devices."""
try:
handles = amdsmi_get_processor_handles()
numa_nodes = []
for device_id in range(cls.device_count()):
physical_device_id = cls.device_id_to_physical_device_id(device_id)
try:
numa_node = amdsmi_topo_get_numa_node_number(
handles[physical_device_id]
)
except AmdSmiException as e:
logger.warning(
"Could not detect NUMA node for GPU %d, "
"disabling automatic NUMA binding: %s",
device_id,
e,
)
return None
numa_nodes.append(numa_node)
return numa_nodes
except Exception as e:
logger.warning("Failed to get NUMA nodes for GPUs: %s", e)
return None