From 740f379faefd1ddeb47c85ebd600af3ccae77b7a Mon Sep 17 00:00:00 2001 From: BadrBasowid <61441185+BadrBasowid@users.noreply.github.com> Date: Mon, 6 Jul 2026 20:16:32 +0800 Subject: [PATCH] [ROCm][AITER] Directly Implement AITER Custom All-reduce in CudaCommunicator (#46065) Signed-off-by: BadrBasowid --- .buildkite/test-amd.yaml | 6 +- tests/compile/fusions_e2e/conftest.py | 1 + .../distributed/test_fusion_all_reduce.py | 10 +- .../distributed/test_rocm_aiter_custom_ar.py | 134 ++++++++++++++++++ tests/utils.py | 40 ++++++ .../e2e/general/test_rocm_aiter_custom_ar.py | 118 +++++++++++++++ vllm/_aiter_ops.py | 129 +++++------------ .../passes/fusion/allreduce_rms_fusion.py | 49 ++----- vllm/config/vllm.py | 9 +- .../aiter_custom_all_reduce.py | 95 +++++++++++++ .../device_communicators/cuda_communicator.py | 50 +++++-- vllm/envs.py | 6 + 12 files changed, 499 insertions(+), 148 deletions(-) create mode 100644 tests/distributed/test_rocm_aiter_custom_ar.py create mode 100644 tests/v1/e2e/general/test_rocm_aiter_custom_ar.py create mode 100644 vllm/distributed/device_communicators/aiter_custom_all_reduce.py diff --git a/.buildkite/test-amd.yaml b/.buildkite/test-amd.yaml index 7f8eef67c41..eee56a3c401 100644 --- a/.buildkite/test-amd.yaml +++ b/.buildkite/test-amd.yaml @@ -2782,15 +2782,19 @@ steps: - vllm/envs.py - examples/offline_inference/data_parallel.py - tests/distributed/test_context_parallel.py + - tests/distributed/test_rocm_aiter_custom_ar.py - tests/distributed/test_rocm_quick_reduce.py - tests/distributed/test_quick_all_reduce.py + - tests/v1/e2e/general/test_rocm_aiter_custom_ar.py - tests/v1/distributed/test_dbo.py - tests/utils.py commands: - pytest -v -s tests/distributed/test_context_parallel.py - - pytest -v -s tests/v1/distributed/test_dbo.py + - pytest -v -s tests/distributed/test_rocm_aiter_custom_ar.py + - pytest -v -s tests/v1/e2e/general/test_rocm_aiter_custom_ar.py - pytest -v -s tests/distributed/test_rocm_quick_reduce.py - pytest -v -s tests/distributed/test_quick_all_reduce.py + - pytest -v -s tests/v1/distributed/test_dbo.py #-------------------------------------------------------- mi355 ยท entrypoints --------------------------------------------------------# diff --git a/tests/compile/fusions_e2e/conftest.py b/tests/compile/fusions_e2e/conftest.py index 9f34d25c46d..a4ed63ffe7b 100644 --- a/tests/compile/fusions_e2e/conftest.py +++ b/tests/compile/fusions_e2e/conftest.py @@ -79,6 +79,7 @@ def run_e2e_fusion_test(monkeypatch, caplog_mp_spawn): ): monkeypatch.setenv("VLLM_USE_DEEP_GEMM", "1" if use_deepgemm else "0") monkeypatch.setenv("VLLM_ROCM_USE_AITER", "1" if use_aiter else "0") + monkeypatch.setenv("VLLM_ROCM_USE_AITER_CUSTOM_AR", "1" if use_aiter else "0") from vllm._aiter_ops import rocm_aiter_ops rocm_aiter_ops.refresh_env_variables() diff --git a/tests/compile/passes/distributed/test_fusion_all_reduce.py b/tests/compile/passes/distributed/test_fusion_all_reduce.py index 83ce458aafd..17c4e30959b 100644 --- a/tests/compile/passes/distributed/test_fusion_all_reduce.py +++ b/tests/compile/passes/distributed/test_fusion_all_reduce.py @@ -30,6 +30,9 @@ from vllm.config import ( set_current_vllm_config, ) from vllm.distributed import tensor_model_parallel_all_reduce +from vllm.distributed.device_communicators.aiter_custom_all_reduce import ( + AiterCustomAllreduce, +) from vllm.distributed.parallel_state import ( init_distributed_environment, initialize_model_parallel, @@ -504,8 +507,12 @@ def all_reduce_fusion_pass_on_test_model( "MASTER_ADDR": "localhost", "MASTER_PORT": "12345", "VLLM_FLASHINFER_ALLREDUCE_BACKEND": flashinfer_allreduce_backend, + "VLLM_ROCM_USE_AITER": str(int(use_aiter)), + "VLLM_ROCM_USE_AITER_CUSTOM_AR": str(int(use_aiter)), } ) + if use_aiter: + rocm_aiter_ops.refresh_env_variables() init_distributed_environment() @@ -616,7 +623,7 @@ def test_rocm_aiter_all_reduce_rmsnorm_group_quant_fp8_fusion_pass_replace( m.setenv("VLLM_ROCM_USE_AITER", "1") rocm_aiter_ops.refresh_env_variables() - if not rocm_aiter_ops.has_fused_allreduce_rmsnorm_quant_per_group(): + if not AiterCustomAllreduce.build_supports_per_group_quant(): pytest.skip( "aiter build is missing 'fused_ar_rms_per_group_quant' (needs " "ROCm/aiter PR #2823); the new patterns aren't registered." @@ -671,6 +678,7 @@ def rocm_aiter_group_quant_fusion_pass_on_test_model( "MASTER_ADDR": "localhost", "MASTER_PORT": "12345", "VLLM_ROCM_USE_AITER": "1", + "VLLM_ROCM_USE_AITER_CUSTOM_AR": "1", } ) rocm_aiter_ops.refresh_env_variables() diff --git a/tests/distributed/test_rocm_aiter_custom_ar.py b/tests/distributed/test_rocm_aiter_custom_ar.py new file mode 100644 index 00000000000..0b85f36410d --- /dev/null +++ b/tests/distributed/test_rocm_aiter_custom_ar.py @@ -0,0 +1,134 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import ray +import torch +import torch.distributed as dist + +from vllm._aiter_ops import is_aiter_found, rocm_aiter_ops +from vllm.distributed.communication_op import tensor_model_parallel_all_reduce # noqa +from vllm.distributed.parallel_state import get_tp_group, graph_capture +from vllm.envs import disable_envs_cache +from vllm.platforms import current_platform + +from ..utils import ( + assert_rocm_custom_allreduce_backend_state, + ensure_model_parallel_initialized, + init_test_distributed_environment, + multi_gpu_test, + multi_process_parallel, +) + +pytestmark = pytest.mark.skipif( + not current_platform.is_rocm(), + reason="ROCm-only AITER custom allreduce tests", +) + +test_cases = [ + ((2, 7168), torch.float16), + ((2, 7168), torch.bfloat16), + ((128, 8192), torch.float16), + ((128, 8192), torch.bfloat16), +] + + +def _configure_aiter_custom_ar_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("CUDA_VISIBLE_DEVICES", raising=False) + monkeypatch.delenv("HIP_VISIBLE_DEVICES", raising=False) + monkeypatch.setenv("VLLM_ROCM_USE_AITER", "1") + monkeypatch.setenv("VLLM_ROCM_USE_AITER_CUSTOM_AR", "1") + monkeypatch.setenv("VLLM_ROCM_QUICK_REDUCE_QUANTIZATION", "NONE") + disable_envs_cache() + rocm_aiter_ops.refresh_env_variables() + + +def _assert_aiter_handles_input(inp: torch.Tensor) -> None: + aiter_ar_comm = get_tp_group().device_communicator.aiter_ar_comm + assert aiter_ar_comm is not None + assert aiter_ar_comm.should_custom_ar(inp), ( + f"AITER CustomAllreduce does not support input shape {inp.shape}." + ) + + +@ray.remote(num_gpus=1, max_calls=1) +def graph_allreduce( + monkeypatch: pytest.MonkeyPatch, + tp_size, + pp_size, + rank, + distributed_init_port, +) -> None: + with monkeypatch.context() as m: + _configure_aiter_custom_ar_env(m) + + device = torch.device(f"cuda:{rank}") + torch.accelerator.set_device_index(device) + init_test_distributed_environment(tp_size, pp_size, rank, distributed_init_port) + ensure_model_parallel_initialized(tp_size, pp_size) + assert_rocm_custom_allreduce_backend_state(True, "NONE") + group = get_tp_group().device_group + + # A small all_reduce for warmup. + # this is needed because device communicators might be created lazily + # (e.g. NCCL). This will ensure that the communicator is initialized + # before any communication happens, so that this group can be used for + # graph capture immediately. + data = torch.zeros(1) + data = data.to(device=device) + dist.all_reduce(data, group=group) + torch.accelerator.synchronize() + del data + + for shape, dtype in test_cases: + with graph_capture(device=device) as graph_capture_context: + inp = torch.ones(shape, dtype=dtype, device=device) + _assert_aiter_handles_input(inp) + expected = inp * tp_size + + torch.accelerator.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=graph_capture_context.stream): + out = tensor_model_parallel_all_reduce(inp) + + graph.replay() + torch.testing.assert_close(out, expected) + + +@ray.remote(num_gpus=1, max_calls=1) +def eager_allreduce( + monkeypatch: pytest.MonkeyPatch, + tp_size, + pp_size, + rank, + distributed_init_port, +) -> None: + with monkeypatch.context() as m: + _configure_aiter_custom_ar_env(m) + + device = torch.device(f"cuda:{rank}") + torch.accelerator.set_device_index(device) + init_test_distributed_environment(tp_size, pp_size, rank, distributed_init_port) + ensure_model_parallel_initialized(tp_size, pp_size) + assert_rocm_custom_allreduce_backend_state(True, "NONE") + + for shape, dtype in test_cases: + inp = torch.ones(shape, dtype=dtype, device=device) + _assert_aiter_handles_input(inp) + expected = inp * tp_size + out = tensor_model_parallel_all_reduce(inp) + torch.testing.assert_close(out, expected) + + +@pytest.mark.skipif(not is_aiter_found(), reason="AITER is not installed") +@multi_gpu_test(num_gpus=2) +@pytest.mark.parametrize("tp_size", [2]) +@pytest.mark.parametrize("pipeline_parallel_size", [1]) +@pytest.mark.parametrize("test_target", [eager_allreduce, graph_allreduce]) +def test_rocm_aiter_custom_allreduce( + monkeypatch: pytest.MonkeyPatch, + tp_size, + pipeline_parallel_size, + test_target, +): + multi_process_parallel(monkeypatch, tp_size, pipeline_parallel_size, test_target) diff --git a/tests/utils.py b/tests/utils.py index 08579f99e4d..2a3bdb91fe0 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -1459,6 +1459,46 @@ def multi_process_parallel( ray.shutdown() +def assert_rocm_custom_allreduce_backend_state( + use_aiter_custom_ar: bool, + quick_reduce_quantization: str, +) -> None: + from vllm.distributed.parallel_state import get_tp_group + + device_communicator = get_tp_group().device_communicator + aiter_ar_comm = device_communicator.aiter_ar_comm + if use_aiter_custom_ar: + assert aiter_ar_comm is not None, "AITER CustomAllreduce was not initialized." + assert not aiter_ar_comm.disabled, "AITER CustomAllreduce is disabled." + assert device_communicator.ca_comm is None, ( + "vLLM CustomAllreduce should not be initialized when AITER CA is used." + ) + else: + assert aiter_ar_comm is None, ( + "AITER CustomAllreduce should not be initialized when disabled." + ) + assert device_communicator.ca_comm is not None, ( + "vLLM CustomAllreduce should be initialized when AITER CA is disabled." + ) + + qr_comm = device_communicator.qr_comm + assert qr_comm is not None, "QuickReduce communicator was not initialized." + if quick_reduce_quantization == "NONE": + assert qr_comm.disabled, "QuickReduce should be disabled." + else: + assert not qr_comm.disabled, "QuickReduce should be enabled." + + +def assert_rocm_custom_allreduce_backend_state_on_worker( + _worker, + use_aiter_custom_ar: bool, + quick_reduce_quantization: str, +) -> None: + assert_rocm_custom_allreduce_backend_state( + use_aiter_custom_ar, quick_reduce_quantization + ) + + @contextmanager def error_on_warning(category: type[Warning] = Warning): """ diff --git a/tests/v1/e2e/general/test_rocm_aiter_custom_ar.py b/tests/v1/e2e/general/test_rocm_aiter_custom_ar.py new file mode 100644 index 00000000000..18e51584fa8 --- /dev/null +++ b/tests/v1/e2e/general/test_rocm_aiter_custom_ar.py @@ -0,0 +1,118 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest + +from vllm._aiter_ops import is_aiter_found, rocm_aiter_ops +from vllm.config import CompilationConfig, CompilationMode, CUDAGraphMode +from vllm.envs import disable_envs_cache +from vllm.platforms import current_platform + +from ....conftest import VllmRunner +from ....utils import ( + assert_rocm_custom_allreduce_backend_state_on_worker, + multi_gpu_test, +) + +PROMPTS = ["Hello, my name is", "The capital of France is"] + + +def _run_generation( + vllm_runner: type[VllmRunner], + monkeypatch: pytest.MonkeyPatch, + compilation_config: CompilationConfig, + *, + model: str, + max_tokens: int, + use_aiter_custom_ar: bool, + quick_reduce_quantization: str, +) -> list[tuple[list[int], str]]: + with monkeypatch.context() as m: + m.setenv("VLLM_ALLOW_INSECURE_SERIALIZATION", "1") + m.setenv("VLLM_ROCM_USE_AITER", "1") + m.setenv( + "VLLM_ROCM_USE_AITER_CUSTOM_AR", + "1" if use_aiter_custom_ar else "0", + ) + m.setenv("VLLM_ROCM_QUICK_REDUCE_QUANTIZATION", quick_reduce_quantization) + disable_envs_cache() + rocm_aiter_ops.refresh_env_variables() + + with vllm_runner( + model, + dtype="half", + tensor_parallel_size=2, + compilation_config=compilation_config, + max_model_len=256, + max_num_seqs=len(PROMPTS), + gpu_memory_utilization=0.7, + ) as llm: + llm.get_llm().collective_rpc( + assert_rocm_custom_allreduce_backend_state_on_worker, + args=(use_aiter_custom_ar, quick_reduce_quantization), + ) + + return llm.generate_greedy(PROMPTS, max_tokens) + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-only") +@pytest.mark.skipif(not is_aiter_found(), reason="AITER is not installed") +@multi_gpu_test(num_gpus=2) +@pytest.mark.parametrize( + "quick_reduce_quantization", + [ + pytest.param("FP", id="quick-reduce-on"), + pytest.param("NONE", id="quick-reduce-off"), + ], +) +@pytest.mark.parametrize( + "cudagraph_mode", + [ + pytest.param(CUDAGraphMode.NONE, id="cudagraph-none"), + pytest.param(CUDAGraphMode.FULL, id="cudagraph-full"), + ], +) +@pytest.mark.parametrize( + "model,max_tokens", + [ + pytest.param("facebook/opt-125m", 8, id="opt-125m"), + ], +) +def test_rocm_aiter_custom_ar_e2e( + vllm_runner: type[VllmRunner], + monkeypatch: pytest.MonkeyPatch, + cudagraph_mode: CUDAGraphMode, + quick_reduce_quantization: str, + model: str, + max_tokens: int, +): + compilation_mode = ( + CompilationMode.NONE + if cudagraph_mode == CUDAGraphMode.NONE + else CompilationMode.VLLM_COMPILE + ) + compilation_config = CompilationConfig( + mode=compilation_mode, + cudagraph_mode=cudagraph_mode, + ) + + baseline_generations = _run_generation( + vllm_runner, + monkeypatch, + compilation_config, + model=model, + max_tokens=max_tokens, + use_aiter_custom_ar=False, + quick_reduce_quantization=quick_reduce_quantization, + ) + aiter_custom_ar_generations = _run_generation( + vllm_runner, + monkeypatch, + compilation_config, + model=model, + max_tokens=max_tokens, + use_aiter_custom_ar=True, + quick_reduce_quantization=quick_reduce_quantization, + ) + + assert aiter_custom_ar_generations == baseline_generations diff --git a/vllm/_aiter_ops.py b/vllm/_aiter_ops.py index 40ccbcf9a28..bed4d8254a0 100644 --- a/vllm/_aiter_ops.py +++ b/vllm/_aiter_ops.py @@ -2,12 +2,9 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import functools from collections.abc import Callable -from contextlib import contextmanager -from typing import Protocol import torch from torch._ops import OpOverload -from torch.distributed import ProcessGroup import vllm.envs as envs from vllm.platforms import current_platform @@ -52,42 +49,6 @@ def is_aiter_found() -> bool: IS_AITER_FOUND = is_aiter_found() -class AiterCustomAllreduceProto(Protocol): - max_size: int - world_size: int - fully_connected: bool - - @contextmanager - def capture(self): ... - def close(self) -> None: ... - def fused_ar_rms( - self, - inp: torch.Tensor, - res_inp: torch.Tensor, - *, - w: torch.Tensor, - eps: float, - registered: bool = False, - use_1stage: bool = False, - ) -> tuple[torch.Tensor, torch.Tensor]: ... - def fused_ar_rms_per_group_quant( - self, - inp: torch.Tensor, - res_inp: torch.Tensor, - *, - w: torch.Tensor, - eps: float, - group_size: int = 128, - registered: bool = False, - use_1stage: bool = False, - emit_bf16: bool = False, - ) -> ( - tuple[torch.Tensor, torch.Tensor, torch.Tensor] - | tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] - ): ... - def should_custom_ar(self, inp: torch.Tensor) -> bool: ... - - def is_aiter_found_and_supported() -> bool: """Check if AITER library is available and platform supports it. @@ -830,6 +791,7 @@ def _rocm_aiter_fused_allreduce_rmsnorm_impl( ) -> tuple[torch.Tensor, torch.Tensor]: aiter_ar = rocm_aiter_ops.get_aiter_allreduce() assert aiter_ar is not None, "aiter allreduce must be initialized" + ca = aiter_ar.aiter_ca total_bytes = input_.numel() * input_.element_size() hidden_dim = input_.shape[-1] @@ -840,8 +802,8 @@ def _rocm_aiter_fused_allreduce_rmsnorm_impl( else: hidden_ok = False token_ok = token_num <= 80 - world_size = aiter_ar.world_size - full_nvlink = aiter_ar.fully_connected + world_size = ca.world_size + full_nvlink = ca.fully_connected if world_size == 2: size_ok = True @@ -854,12 +816,11 @@ def _rocm_aiter_fused_allreduce_rmsnorm_impl( use_1stage = hidden_ok and token_ok and size_ok - result = aiter_ar.fused_ar_rms( + result = ca.custom_fused_ar_rms( input_, residual, - w=weight, - eps=epsilon, - registered=torch.cuda.is_current_stream_capturing(), + weight, + epsilon, use_1stage=use_1stage, ) assert result is not None @@ -890,6 +851,7 @@ def _rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_impl( """ aiter_ar = rocm_aiter_ops.get_aiter_allreduce() assert aiter_ar is not None, "aiter allreduce must be initialized" + ca = aiter_ar.aiter_ca total_bytes = input_.numel() * input_.element_size() hidden_dim = input_.shape[-1] @@ -900,8 +862,8 @@ def _rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_impl( else: hidden_ok = False token_ok = token_num <= 80 - world_size = aiter_ar.world_size - full_nvlink = aiter_ar.fully_connected + world_size = ca.world_size + full_nvlink = ca.fully_connected if world_size == 2: size_ok = True @@ -914,7 +876,7 @@ def _rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_impl( use_1stage = hidden_ok and token_ok and size_ok - result = aiter_ar.fused_ar_rms_per_group_quant( + result = ca.fused_ar_rms_per_group_quant( input_, residual, w=weight, @@ -962,6 +924,7 @@ def _rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_with_bf16_norm_impl( """ aiter_ar = rocm_aiter_ops.get_aiter_allreduce() assert aiter_ar is not None, "aiter allreduce must be initialized" + ca = aiter_ar.aiter_ca total_bytes = input_.numel() * input_.element_size() hidden_dim = input_.shape[-1] @@ -972,8 +935,8 @@ def _rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_with_bf16_norm_impl( else: hidden_ok = False token_ok = token_num <= 80 - world_size = aiter_ar.world_size - full_nvlink = aiter_ar.fully_connected + world_size = ca.world_size + full_nvlink = ca.fully_connected if world_size == 2: size_ok = True @@ -986,7 +949,7 @@ def _rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_with_bf16_norm_impl( use_1stage = hidden_ok and token_ok and size_ok - result = aiter_ar.fused_ar_rms_per_group_quant( + result = ca.fused_ar_rms_per_group_quant( input_, residual, w=weight, @@ -1577,6 +1540,7 @@ class rocm_aiter_ops: # Check if the env variable is set _AITER_ENABLED = envs.VLLM_ROCM_USE_AITER + _CUSTOM_ALL_REDUCE_ENABLED = envs.VLLM_ROCM_USE_AITER_CUSTOM_AR _LINEAR_ENABLED = envs.VLLM_ROCM_USE_AITER_LINEAR _FMOE_ENABLED = envs.VLLM_ROCM_USE_AITER_MOE _MLA_ENABLED = envs.VLLM_ROCM_USE_AITER_MLA @@ -1598,9 +1562,6 @@ class rocm_aiter_ops: # num_shared_experts / shared_expert_scoring_func args (7-arg form). _TOPK_SOFTMAX_FUSED_SIGMOID: bool | None = None - _ALL_REDUCE_MAX_SIZE: int = 8192 * 1024 * 8 * 2 - _CUSTOM_ALL_REDUCE: AiterCustomAllreduceProto | None = None - @classmethod def refresh_env_variables(cls): """ @@ -1611,6 +1572,7 @@ class rocm_aiter_ops: you can call this function to reload the env variables. """ cls._AITER_ENABLED = envs.VLLM_ROCM_USE_AITER + cls._CUSTOM_ALL_REDUCE_ENABLED = envs.VLLM_ROCM_USE_AITER_CUSTOM_AR cls._LINEAR_ENABLED = envs.VLLM_ROCM_USE_AITER_LINEAR cls._FMOE_ENABLED = envs.VLLM_ROCM_USE_AITER_MOE cls._MLA_ENABLED = envs.VLLM_ROCM_USE_AITER_MLA @@ -1770,6 +1732,11 @@ class rocm_aiter_ops: def is_mha_enabled(cls) -> bool: return cls._AITER_ENABLED and cls._MHA_ENABLED + @classmethod + @if_aiter_supported + def is_custom_all_reduce_enabled(cls) -> bool: + return cls._AITER_ENABLED and cls._CUSTOM_ALL_REDUCE_ENABLED + @classmethod @if_aiter_supported def is_shuffle_kv_cache_enabled(cls) -> bool: @@ -1824,33 +1791,20 @@ class rocm_aiter_ops: return cls.is_linear_enabled() and on_gfx950() @classmethod - def initialize_aiter_allreduce( - cls, group: ProcessGroup, device: torch.device - ) -> None: - try: - from aiter.dist.device_communicators.custom_all_reduce import ( - CustomAllreduce as AiterCustomAllreduce, - ) + def get_aiter_allreduce(cls): + """Return the TP device communicator's AITER custom-allreduce if it has + one, return None otherwise + """ + from vllm.distributed.device_communicators.aiter_custom_all_reduce import ( + AiterCustomAllreduce, + ) + from vllm.distributed.parallel_state import get_tp_group - cls._CUSTOM_ALL_REDUCE = AiterCustomAllreduce(group, device) - except Exception: - cls._CUSTOM_ALL_REDUCE = None - - @classmethod - def get_aiter_allreduce(cls) -> AiterCustomAllreduceProto | None: - return cls._CUSTOM_ALL_REDUCE - - @classmethod - def destroy_aiter_allreduce(cls) -> None: - if cls._CUSTOM_ALL_REDUCE is not None: - cls._CUSTOM_ALL_REDUCE.close() - cls._CUSTOM_ALL_REDUCE = None - - @classmethod - def get_aiter_allreduce_max_size(cls) -> int | None: - # effective max input size (based on upstream aiter version: v0.1.10.post3) - # https://github.com/ROCm/aiter/blob/6a0e7b26ccf33164785531212cc2ec2cde0b9243/aiter/dist/device_communicators/custom_all_reduce.py#L272-L273 - return int(cls._ALL_REDUCE_MAX_SIZE / 2) + device_comm = get_tp_group().device_communicator + aiter_ar_comm = getattr(device_comm, "aiter_ar_comm", None) + return ( + aiter_ar_comm if isinstance(aiter_ar_comm, AiterCustomAllreduce) else None + ) @classmethod @if_aiter_supported @@ -2165,21 +2119,6 @@ class rocm_aiter_ops: def get_fused_allreduce_rmsnorm_quant_per_group_with_bf16_norm_op() -> OpOverload: # noqa: E501 return torch.ops.vllm.rocm_aiter_fused_allreduce_rmsnorm_quant_per_group_with_bf16_norm.default # noqa: E501 - # TODO(frida-andersson): drop once vLLM pins AITER >= 0.1.14 (ROCm/aiter#2823). - @classmethod - def has_fused_allreduce_rmsnorm_quant_per_group(cls) -> bool: - """True if the running AITER build exposes the per-group AR+RMS+quant - kernel (added in ROCm/aiter PR #2823). - - The pattern registration in ``RocmAiterAllReduceFusionPass`` keys off - this so vLLM degrades to the AR+RMS-only fusion when run against an - older aiter that lacks the per-group launcher. - """ - aiter_ar = cls.get_aiter_allreduce() - return aiter_ar is not None and hasattr( - aiter_ar, "fused_ar_rms_per_group_quant" - ) - @staticmethod def get_fused_mla_dual_rms_norm_op() -> OpOverload: return torch.ops.vllm.fused_mla_dual_rms_norm.default diff --git a/vllm/compilation/passes/fusion/allreduce_rms_fusion.py b/vllm/compilation/passes/fusion/allreduce_rms_fusion.py index ee706037abb..f4962424d7f 100644 --- a/vllm/compilation/passes/fusion/allreduce_rms_fusion.py +++ b/vllm/compilation/passes/fusion/allreduce_rms_fusion.py @@ -19,7 +19,6 @@ from vllm.compilation.passes.fusion.rms_quant_fusion import ( from vllm.config import VllmConfig from vllm.config.utils import Range from vllm.distributed import get_tp_group, tensor_model_parallel_all_reduce -from vllm.distributed.device_communicators.custom_all_reduce import CustomAllreduce from vllm.distributed.parallel_state import ( get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, @@ -1473,39 +1472,23 @@ class RocmAiterAllReduceFusionPass(VllmFusionPatternMatcherPass): ) return - device_comm = get_tp_group().device_communicator - if device_comm is None: - logger.warning_once("Device communicator is required.") - return - - ca_comm = getattr(device_comm, "ca_comm", None) + ca_comm = rocm_aiter_ops.get_aiter_allreduce() if ca_comm is None: - logger.warning_once("Custom Allreduce is required.") + logger.warning_once( + "AITER allreduce fusions are disabled " + "because AITER Custom All Reduce is not enabled. " + "Set VLLM_ROCM_USE_AITER_CUSTOM_AR=1 " + "to enable it." + ) return self.ca_comm = ca_comm - assert isinstance(ca_comm, CustomAllreduce) - - group = get_tp_group().cpu_group - rocm_aiter_ops.initialize_aiter_allreduce(group, self.device) hidden_dim = config.model_config.get_hidden_size() element_size = torch.tensor([], dtype=self.model_dtype).element_size() - max_size = rocm_aiter_ops.get_aiter_allreduce_max_size() - if max_size is None: - logger.warning("AITER allreduce fusion must be initialized") - return - - # Aiter's fused_allreduce_rmsnorm kernel dispatches on hidden_dim. - # Before aiter v0.1.12 the launcher was template-specialized on HIDDEN_DIM - # and silently no-op'd for sizes outside {512, 1024, 2048, 4096}. From v0.1.12 - # hidden_dim is a runtime argument. Detect the older API via the missing - # `_pool` attribute and skip fusion for unsupported sizes. - # Ref (old kernel): https://github.com/ROCm/aiter/blob/6a0e7b26ccf33164785531212cc2ec2cde0b9243/csrc/include/custom_all_reduce.cuh#L2590 - aiter_ar = rocm_aiter_ops.get_aiter_allreduce() + max_size = ca_comm.effective_max_size() _AITER_OLD_FUSED_AR_RMS_HIDDEN = (512, 1024, 2048, 4096) if ( - aiter_ar is not None - and not hasattr(aiter_ar, "_pool") + not ca_comm.supports_dynamic_hidden_dim and hidden_dim not in _AITER_OLD_FUSED_AR_RMS_HIDDEN ): logger.warning_once( @@ -1515,10 +1498,6 @@ class RocmAiterAllReduceFusionPass(VllmFusionPatternMatcherPass): _AITER_OLD_FUSED_AR_RMS_HIDDEN, hidden_dim, ) - # Tear down aiter's custom-allreduce so its IPC handles don't - # race with vllm's ca_comm on the unfused fallback path. - with contextlib.suppress(Exception): - rocm_aiter_ops.destroy_aiter_allreduce() return max_token_num = max_size // (hidden_dim * element_size) @@ -1532,9 +1511,7 @@ class RocmAiterAllReduceFusionPass(VllmFusionPatternMatcherPass): # fall back to the AR+RMS-only fusion paired with PR #41825's # standalone RMS+quant fusion -- still correct, just leaves the # post-AR quant as a standalone kernel. - supports_per_group_quant = ( - rocm_aiter_ops.has_fused_allreduce_rmsnorm_quant_per_group() - ) + supports_per_group_quant = ca_comm.supports_per_group_quant if not supports_per_group_quant: logger.warning_once( "AITER AR+RMS+per-group-FP8-quant fusion disabled: aiter " @@ -1609,9 +1586,3 @@ class RocmAiterAllReduceFusionPass(VllmFusionPatternMatcherPass): logger.debug( "%s Replaced %s patterns", self.__class__.__name__, self.matched_count ) - - def __del__(self) -> None: - if getattr(self, "disabled", True): - return - with contextlib.suppress(Exception): - rocm_aiter_ops.destroy_aiter_allreduce() diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 8f90bb23326..e8796a3db8e 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -1850,8 +1850,13 @@ class VllmConfig: tp_size = self.parallel_config.tensor_parallel_size from vllm._aiter_ops import rocm_aiter_ops - if rocm_aiter_ops.is_enabled(): - max_size = rocm_aiter_ops.get_aiter_allreduce_max_size() + max_size: int | None = None + if rocm_aiter_ops.is_custom_all_reduce_enabled(): + from vllm.distributed.device_communicators.aiter_custom_all_reduce import ( # noqa: E501 + AiterCustomAllreduce, + ) + + max_size = AiterCustomAllreduce.effective_max_size() else: max_size = compilation_config.pass_config.flashinfer_max_size(tp_size) if max_size is not None and self.model_config is not None: diff --git a/vllm/distributed/device_communicators/aiter_custom_all_reduce.py b/vllm/distributed/device_communicators/aiter_custom_all_reduce.py new file mode 100644 index 00000000000..63e06b77a77 --- /dev/null +++ b/vllm/distributed/device_communicators/aiter_custom_all_reduce.py @@ -0,0 +1,95 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""vLLM-owned wrapper over AITER's ``CustomAllreduce``. + +vLLM's ``CudaCommunicator`` stores one of these as ``aiter_ar_comm`` (when +``VLLM_ROCM_USE_AITER_CUSTOM_AR`` is set) so the plain allreduce and +the fused allreduce+RMSNorm path share a single AITER instance with its IPC buffers. + +""" + +import torch +from torch.distributed import ProcessGroup + +from vllm.logger import init_logger + +logger = init_logger(__name__) + + +class AiterCustomAllreduce: + # Default IPC buffer size for AITER's CustomAllreduce. + MAX_SIZE: int = 8192 * 1024 * 8 * 2 + + @classmethod + def effective_max_size(cls) -> int: + """ + Max input byte size eligible for AITER custom allreduce. + """ + return cls.MAX_SIZE // 2 + + def __init__( + self, + group: ProcessGroup, + device: int | str | torch.device, + max_size: int | None = None, + ): + from aiter.dist.device_communicators.custom_all_reduce import ( + CustomAllreduce as _AiterCustomAllreduce, + ) + + if max_size is None: + max_size = self.MAX_SIZE + + self._impl = _AiterCustomAllreduce(group, device, max_size=max_size) + + @property + def aiter_ca(self): + return self._impl + + @property + def disabled(self) -> bool: + return self._impl.disabled + + def should_custom_ar(self, inp: torch.Tensor) -> bool: + return self._impl.should_custom_ar(inp) + + def custom_all_reduce(self, inp: torch.Tensor) -> torch.Tensor | None: + return self._impl.custom_all_reduce(inp) + + def capture(self): + return self._impl.capture() + + def close(self) -> None: + self._impl.close() + + @property + def supports_dynamic_hidden_dim(self) -> bool: + """Aiter's fused_allreduce_rmsnorm kernel dispatches on hidden_dim. + Before aiter v0.1.12 the launcher was template-specialized on HIDDEN_DIM + and silently no-op'd for sizes outside {512, 1024, 2048, 4096}. From v0.1.12 + hidden_dim is a runtime argument. Older builds are detected via + AiterCustomAllreduce.supports_dynamic_hidden_dim; This function is used to + skip fusion for unsupported sizes on them. + Ref (old kernel): https://github.com/ROCm/aiter/blob/6a0e7b26ccf33164785531212cc2ec2cde0b9243/csrc/include/custom_all_reduce.cuh#L2590 + """ + return hasattr(self._impl, "_pool") + + @staticmethod + def build_supports_per_group_quant() -> bool: + """True if the running AITER build exposes the per-group AR+RMS+quant + kernel (added in ROCm/aiter PR #2823). + + The pattern registration in ``RocmAiterAllReduceFusionPass`` keys off + this so vLLM degrades to the AR+RMS-only fusion when run against an + older aiter that lacks the per-group launcher. + """ + from aiter.dist.device_communicators.custom_all_reduce import ( + CustomAllreduce as _AiterCustomAllreduce, + ) + + return hasattr(_AiterCustomAllreduce, "fused_ar_rms_per_group_quant") + + # TODO(frida-andersson): drop once vLLM pins AITER >= 0.1.14 (ROCm/aiter#2823). + @property + def supports_per_group_quant(self) -> bool: + return self.build_supports_per_group_quant() diff --git a/vllm/distributed/device_communicators/cuda_communicator.py b/vllm/distributed/device_communicators/cuda_communicator.py index b92015b1880..555fd0ec948 100644 --- a/vllm/distributed/device_communicators/cuda_communicator.py +++ b/vllm/distributed/device_communicators/cuda_communicator.py @@ -6,6 +6,7 @@ import torch from torch.distributed import ProcessGroup import vllm.envs as envs +from vllm._aiter_ops import rocm_aiter_ops from vllm.distributed.device_communicators.all_reduce_utils import ( NCCL_SYMM_MEM_ALL_REDUCE_CONFIG, should_nccl_symm_mem_ag_rs, @@ -19,6 +20,7 @@ from vllm.logger import init_logger from vllm.platforms import current_platform from ..utils import StatelessProcessGroup +from .aiter_custom_all_reduce import AiterCustomAllreduce from .base_device_communicator import DeviceCommunicatorBase logger = init_logger(__name__) @@ -48,16 +50,21 @@ class CudaCommunicator(DeviceCommunicatorBase): use_custom_allreduce = False use_torch_symm_mem = False use_flashinfer_allreduce = False + use_aiter_allreduce = False else: from vllm.distributed.parallel_state import _ENABLE_CUSTOM_ALL_REDUCE use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE use_torch_symm_mem = envs.VLLM_ALLREDUCE_USE_SYMM_MEM use_flashinfer_allreduce = envs.VLLM_ALLREDUCE_USE_FLASHINFER + use_aiter_allreduce = use_custom_allreduce and bool( + rocm_aiter_ops.is_custom_all_reduce_enabled() + ) self.use_custom_allreduce = use_custom_allreduce self.use_torch_symm_mem = use_torch_symm_mem self.use_flashinfer_allreduce = use_flashinfer_allreduce + self.use_aiter_allreduce = use_aiter_allreduce # lazy import to avoid documentation build error from vllm.distributed.device_communicators.custom_all_reduce import ( @@ -85,6 +92,7 @@ class CudaCommunicator(DeviceCommunicatorBase): self.qr_comm: QuickAllReduce | None = None self.symm_mem_comm: SymmMemCommunicator | None = None self.fi_ar_comm: FlashInferAllReduce | None = None + self.aiter_ar_comm: AiterCustomAllreduce | None = None if use_torch_symm_mem and current_platform.is_cuda(): self.symm_mem_comm = SymmMemCommunicator( @@ -98,7 +106,13 @@ class CudaCommunicator(DeviceCommunicatorBase): device=self.device, ) - if use_custom_allreduce and self.world_size > 1: + if self.use_aiter_allreduce and self.world_size > 1: + self.aiter_ar_comm = AiterCustomAllreduce( + group=self.cpu_group, + device=self.device, + ) + + if use_custom_allreduce and self.aiter_ar_comm is None and self.world_size > 1: # Initialize a custom fast all-reduce implementation. self.ca_comm = CustomAllreduce( group=self.cpu_group, @@ -108,13 +122,14 @@ class CudaCommunicator(DeviceCommunicatorBase): ), ) - if current_platform.is_rocm(): - # Initialize a custom quick all-reduce implementation for AMD. - # Quick reduce is designed as a complement to custom allreduce. - # Based on quickreduce (https://github.com/mk1-project/quickreduce). - # If it's a rocm, 'use_custom_allreduce==True' means it must - # currently be an MI300 series. - self.qr_comm = QuickAllReduce(group=self.cpu_group, device=self.device) + if use_custom_allreduce and self.world_size > 1 and current_platform.is_rocm(): + # Initialize a custom quick all-reduce implementation for AMD. + # Quick reduce is designed as a complement to custom allreduce + # (vLLM's or AITER's), so it is initialized for either backend. + # Based on quickreduce (https://github.com/mk1-project/quickreduce). + # On ROCm, 'use_custom_allreduce==True' means it must currently be + # an MI300 series. + self.qr_comm = QuickAllReduce(group=self.cpu_group, device=self.device) if self.world_size > 1: self._log_all_reduce_backend_selection() @@ -203,6 +218,7 @@ class CudaCommunicator(DeviceCommunicatorBase): "NCCL_SYMM_MEM", "QUICK_REDUCE", "FLASHINFER", + "AITER_CUSTOM", "CUSTOM", "SYMM_MEM", "PYNCCL", @@ -236,6 +252,8 @@ class CudaCommunicator(DeviceCommunicatorBase): enabled_ar_backends.append("QUICK_REDUCE") if self.fi_ar_comm is not None and not self.fi_ar_comm.disabled: enabled_ar_backends.append("FLASHINFER") + if self.aiter_ar_comm is not None and not self.aiter_ar_comm.disabled: + enabled_ar_backends.append("AITER_CUSTOM") if self.ca_comm is not None and not self.ca_comm.disabled: enabled_ar_backends.append("CUSTOM") if self.symm_mem_comm is not None and not self.symm_mem_comm.disabled: @@ -261,8 +279,8 @@ class CudaCommunicator(DeviceCommunicatorBase): out = torch.ops.vllm.all_reduce_symmetric_with_copy(input_) if out is not None: return out - # always try quick reduce first, then flashinfer, then custom allreduce, - # and then pynccl. (quick reduce just for ROCM MI3*) + # always try quick reduce first, then flashinfer, then the AITER or vLLM + # custom allreduce, and then pynccl. (quick reduce just for ROCM MI3*) qr_comm = self.qr_comm if ( qr_comm is not None @@ -281,6 +299,15 @@ class CudaCommunicator(DeviceCommunicatorBase): out = fi_ar_comm.all_reduce(input_) assert out is not None return out + aiter_ar_comm = self.aiter_ar_comm + if ( + aiter_ar_comm is not None + and not aiter_ar_comm.disabled + and aiter_ar_comm.should_custom_ar(input_) + ): + out = aiter_ar_comm.custom_all_reduce(input_) + assert out is not None + return out ca_comm = self.ca_comm if ( ca_comm is not None @@ -509,6 +536,9 @@ class CudaCommunicator(DeviceCommunicatorBase): self.pynccl_comm = None if self.ca_comm is not None: self.ca_comm = None + if self.aiter_ar_comm is not None: + self.aiter_ar_comm.close() + self.aiter_ar_comm = None if self.fi_ar_comm is not None: self.fi_ar_comm.destroy() self.fi_ar_comm = None diff --git a/vllm/envs.py b/vllm/envs.py index 1f94be8ac5f..13a19b86c0a 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -119,6 +119,7 @@ if TYPE_CHECKING: VLLM_USE_OINK_OPS: bool = False VLLM_MXFP8_EMULATION_DEQUANT_AT_LOAD: bool = True VLLM_ROCM_USE_AITER: bool = False + VLLM_ROCM_USE_AITER_CUSTOM_AR: bool = True VLLM_ROCM_USE_AITER_PAGED_ATTN: bool = False VLLM_ROCM_USE_AITER_LINEAR: bool = True VLLM_ROCM_USE_AITER_LINEAR_HIPBMM: bool = False @@ -1146,6 +1147,11 @@ environment_variables: dict[str, Callable[[], Any]] = { "VLLM_ROCM_USE_AITER": lambda: ( os.getenv("VLLM_ROCM_USE_AITER", "False").lower() in ("true", "1") ), + # Use AITER's CustomAllreduce as the custom-allreduce backend inside vLLM's + # CudaCommunicator on ROCm. + "VLLM_ROCM_USE_AITER_CUSTOM_AR": lambda: ( + os.getenv("VLLM_ROCM_USE_AITER_CUSTOM_AR", "True").lower() in ("true", "1") + ), # Whether to use aiter paged attention. # By default is disabled. "VLLM_ROCM_USE_AITER_PAGED_ATTN": lambda: (