From b0ce38e14c8d1162e2cde34613d1b10bc89ea5a4 Mon Sep 17 00:00:00 2001 From: Yanyuan Qin Date: Wed, 29 Jul 2026 21:56:11 +0000 Subject: [PATCH] perf(rocm): reduce Kimi-K3 attention into stable outputs Add a fail-closed caller-owned all-reduce destination and thread it through row-parallel output projection so AMD Kimi-K3 KDA and MLA write directly into graph-stable decoder buffers. Default callers and NVIDIA call sites retain the existing allocation path; older AITER builds fall back to all-reduce plus copy. Assisted-by: OpenAI Codex Signed-off-by: Yanyuan Qin --- tests/distributed/test_all_reduce_into.py | 207 ++++++++++++++++++ vllm/distributed/communication_op.py | 9 + .../aiter_custom_all_reduce.py | 21 +- .../device_communicators/cuda_communicator.py | 23 ++ vllm/distributed/parallel_state.py | 34 +++ vllm/model_executor/layers/linear.py | 23 +- .../layers/mamba/gdn/kimi_gdn_linear_attn.py | 7 + vllm/model_executor/layers/mla.py | 3 +- vllm/models/kimi_k3/amd/kda.py | 7 + vllm/models/kimi_k3/amd/linear.py | 2 +- 10 files changed, 330 insertions(+), 6 deletions(-) create mode 100644 tests/distributed/test_all_reduce_into.py diff --git a/tests/distributed/test_all_reduce_into.py b/tests/distributed/test_all_reduce_into.py new file mode 100644 index 000000000000..859429b2e905 --- /dev/null +++ b/tests/distributed/test_all_reduce_into.py @@ -0,0 +1,207 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import pytest +import torch + +from vllm.distributed.device_communicators.aiter_custom_all_reduce import ( + AiterCustomAllreduce, +) +from vllm.distributed.device_communicators.cuda_communicator import CudaCommunicator +from vllm.distributed.parallel_state import GroupCoordinator +from vllm.model_executor.layers import linear + + +def test_aiter_adapter_detects_caller_output_support(): + supported = object.__new__(AiterCustomAllreduce) + supported._impl = SimpleNamespace(supports_custom_all_reduce_out=True) + unsupported = object.__new__(AiterCustomAllreduce) + unsupported._impl = SimpleNamespace() + + assert supported.supports_custom_all_reduce_out + assert not unsupported.supports_custom_all_reduce_out + + +def test_aiter_adapter_forwards_caller_owned_output(): + calls = [] + output = torch.empty(8) + + def custom_all_reduce(input_, **kwargs): + calls.append((input_, kwargs)) + return kwargs["out"] + + implementation = SimpleNamespace(custom_all_reduce=custom_all_reduce) + communicator = SimpleNamespace(_impl=implementation) + input_ = torch.empty_like(output) + + result = AiterCustomAllreduce.custom_all_reduce( + communicator, + input_, + out=output, + ) + + assert result is output + assert calls == [(input_, {"out": output})] + + +def test_aiter_adapter_does_not_pass_output_to_older_aiter(): + calls = [] + expected = torch.empty(8) + + def custom_all_reduce(input_): + calls.append(input_) + return expected + + implementation = SimpleNamespace(custom_all_reduce=custom_all_reduce) + communicator = SimpleNamespace(_impl=implementation) + input_ = torch.empty_like(expected) + + result = AiterCustomAllreduce.custom_all_reduce(communicator, input_) + + assert result is expected + assert calls == [input_] + + +def test_cuda_communicator_uses_aiter_output_contract(): + calls = [] + output = torch.empty(8) + + def custom_all_reduce(input_, **kwargs): + calls.append((input_, kwargs)) + return kwargs["out"] + + aiter = SimpleNamespace( + disabled=False, + supports_custom_all_reduce_out=True, + should_custom_ar=lambda input_: True, + custom_all_reduce=custom_all_reduce, + ) + communicator = SimpleNamespace( + aiter_ar_comm=aiter, + all_reduce=lambda input_: (_ for _ in ()).throw( + AssertionError("fallback should not run") + ), + ) + input_ = torch.empty_like(output) + + result = CudaCommunicator.all_reduce_into(communicator, input_, output) + + assert result is output + assert calls == [(input_, {"out": output})] + + +def test_cuda_communicator_fallback_copies_result(): + input_ = torch.arange(8, dtype=torch.float32) + output = torch.empty_like(input_) + communicator = SimpleNamespace( + aiter_ar_comm=None, + all_reduce=lambda input_: input_ + 1, + ) + + result = CudaCommunicator.all_reduce_into(communicator, input_, output) + + assert result is output + torch.testing.assert_close(output, input_ + 1) + + +def test_cuda_communicator_falls_back_for_older_aiter(): + input_ = torch.arange(8, dtype=torch.float32) + output = torch.empty_like(input_) + aiter = SimpleNamespace( + disabled=False, + supports_custom_all_reduce_out=False, + should_custom_ar=lambda input_: True, + custom_all_reduce=lambda *args, **kwargs: (_ for _ in ()).throw( + AssertionError("older AITER must not receive an output argument") + ), + ) + communicator = SimpleNamespace( + aiter_ar_comm=aiter, + all_reduce=lambda input_: input_ + 1, + ) + + result = CudaCommunicator.all_reduce_into(communicator, input_, output) + + assert result is output + torch.testing.assert_close(output, input_ + 1) + + +def test_group_coordinator_world_size_one_preserves_output_identity(): + input_ = torch.arange(8, dtype=torch.float32) + output = torch.empty_like(input_) + coordinator = SimpleNamespace(world_size=1) + + result = GroupCoordinator.all_reduce_into(coordinator, input_, output) + + assert result is output + torch.testing.assert_close(output, input_) + + +def test_group_coordinator_rejects_mismatched_output(): + coordinator = SimpleNamespace(world_size=1) + + with pytest.raises(ValueError, match="all-reduce output"): + GroupCoordinator.all_reduce_into( + coordinator, + torch.empty(8), + torch.empty(7), + ) + + +def test_row_parallel_linear_passes_stable_output(monkeypatch): + calls = [] + input_ = torch.arange(8, dtype=torch.float32).reshape(1, 8) + partial = input_ + 1 + output = torch.empty_like(partial) + layer = SimpleNamespace( + input_is_parallel=True, + tp_rank=0, + skip_bias_add=False, + bias=None, + quant_method=SimpleNamespace( + apply=lambda layer_, input_parallel, bias: partial + ), + reduce_results=True, + tp_size=8, + return_bias=False, + ) + + def all_reduce_into(input_parallel, destination): + calls.append((input_parallel, destination)) + destination.copy_(input_parallel) + return destination + + monkeypatch.setattr( + linear, + "tensor_model_parallel_all_reduce_into", + all_reduce_into, + ) + + result = linear.RowParallelLinear.forward(layer, input_, output=output) + + assert result is output + assert calls == [(partial, output)] + torch.testing.assert_close(output, partial) + + +def test_row_parallel_linear_rejects_mismatched_output(): + input_ = torch.arange(8, dtype=torch.float32).reshape(1, 8) + layer = SimpleNamespace( + input_is_parallel=True, + tp_rank=0, + skip_bias_add=False, + bias=None, + quant_method=SimpleNamespace( + apply=lambda layer_, input_parallel, bias: input_parallel + ), + reduce_results=True, + tp_size=8, + return_bias=False, + ) + + mismatched = torch.empty((1, 7), dtype=torch.float32) + + with pytest.raises(ValueError, match="row-parallel output"): + linear.RowParallelLinear.forward(layer, input_, output=mismatched) diff --git a/vllm/distributed/communication_op.py b/vllm/distributed/communication_op.py index 5ad99e4e1592..0978a3beacc2 100644 --- a/vllm/distributed/communication_op.py +++ b/vllm/distributed/communication_op.py @@ -14,6 +14,15 @@ def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor: return get_tp_group().all_reduce(input_) +def tensor_model_parallel_all_reduce_into( + input_: torch.Tensor, + output: torch.Tensor, +) -> torch.Tensor: + """All-reduce into a caller-owned model-parallel output tensor.""" + + return get_tp_group().all_reduce_into(input_, output) + + def tensor_model_parallel_all_gather( input_: torch.Tensor, dim: int = -1 ) -> torch.Tensor: diff --git a/vllm/distributed/device_communicators/aiter_custom_all_reduce.py b/vllm/distributed/device_communicators/aiter_custom_all_reduce.py index 63e06b77a77e..e29e33c6bcbe 100644 --- a/vllm/distributed/device_communicators/aiter_custom_all_reduce.py +++ b/vllm/distributed/device_communicators/aiter_custom_all_reduce.py @@ -53,8 +53,25 @@ def disabled(self) -> bool: 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) + @property + def supports_custom_all_reduce_out(self) -> bool: + return bool( + getattr( + self._impl, + "supports_custom_all_reduce_out", + False, + ) + ) + + def custom_all_reduce( + self, + inp: torch.Tensor, + *, + out: torch.Tensor | None = None, + ) -> torch.Tensor | None: + if out is None: + return self._impl.custom_all_reduce(inp) + return self._impl.custom_all_reduce(inp, out=out) def capture(self): return self._impl.capture() diff --git a/vllm/distributed/device_communicators/cuda_communicator.py b/vllm/distributed/device_communicators/cuda_communicator.py index 06b441c5a416..71f17999f379 100644 --- a/vllm/distributed/device_communicators/cuda_communicator.py +++ b/vllm/distributed/device_communicators/cuda_communicator.py @@ -340,6 +340,29 @@ def all_reduce(self, input_): torch.distributed.all_reduce(out, group=self.device_group) return out + def all_reduce_into( + self, + input_: torch.Tensor, + output: torch.Tensor, + ) -> torch.Tensor: + """All-reduce into a stable caller-owned destination when supported.""" + + aiter_ar_comm = self.aiter_ar_comm + if ( + aiter_ar_comm is not None + and not aiter_ar_comm.disabled + and aiter_ar_comm.supports_custom_all_reduce_out + and aiter_ar_comm.should_custom_ar(input_) + ): + result = aiter_ar_comm.custom_all_reduce(input_, out=output) + if result is not None: + return result + + result = self.all_reduce(input_) + if result.data_ptr() != output.data_ptr(): + output.copy_(result) + return output + def custom_all_gather(self, input_: torch.Tensor) -> torch.Tensor | None: ca_comm = self.ca_comm if ca_comm is None: diff --git a/vllm/distributed/parallel_state.py b/vllm/distributed/parallel_state.py index a90e8acbcad8..605820434414 100644 --- a/vllm/distributed/parallel_state.py +++ b/vllm/distributed/parallel_state.py @@ -688,6 +688,40 @@ def _all_reduce_out_place(self, input_: torch.Tensor) -> torch.Tensor: raise ValueError("No device communicator found") return self.device_communicator.all_reduce(input_) + def all_reduce_into( + self, + input_: torch.Tensor, + output: torch.Tensor, + ) -> torch.Tensor: + """All-reduce into a stable output tensor without a custom-op call.""" + + if ( + output.shape != input_.shape + or output.dtype != input_.dtype + or output.device != input_.device + or not output.is_contiguous() + ): + raise ValueError( + "all-reduce output must be a contiguous tensor matching " + "the input's shape, dtype, and device" + ) + if self.world_size == 1: + output.copy_(input_) + return output + if self.device_communicator is None: + raise ValueError("No device communicator found") + + implementation = getattr( + self.device_communicator, + "all_reduce_into", + None, + ) + if implementation is not None: + return implementation(input_, output) + + output.copy_(self.device_communicator.all_reduce(input_)) + return output + def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor: world_size = self.world_size # Bypass the function if we are using only 1 GPU. diff --git a/vllm/model_executor/layers/linear.py b/vllm/model_executor/layers/linear.py index e4662148ff75..550e4349bd28 100644 --- a/vllm/model_executor/layers/linear.py +++ b/vllm/model_executor/layers/linear.py @@ -19,6 +19,7 @@ split_tensor_along_last_dim, tensor_model_parallel_all_gather, tensor_model_parallel_all_reduce, + tensor_model_parallel_all_reduce_into, ) from vllm.logger import init_logger from vllm.model_executor.custom_op import PluggableLayer @@ -1748,6 +1749,7 @@ def weight_loader_v2(self, param: BasevLLMParameter, loaded_weight: torch.Tensor def forward( self, input_, + output: torch.Tensor | None = None, ) -> torch.Tensor | tuple[torch.Tensor, Parameter | None]: if self.input_is_parallel: input_parallel = input_ @@ -1763,10 +1765,27 @@ def forward( bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias output_parallel = self.quant_method.apply(self, input_parallel, bias_) + if output is not None and ( + output.shape != output_parallel.shape + or output.dtype != output_parallel.dtype + or output.device != output_parallel.device + or not output.is_contiguous() + ): + raise ValueError( + "row-parallel output must be a contiguous tensor matching " + "the projection output's shape, dtype, and device" + ) + if self.reduce_results and self.tp_size > 1: - output = tensor_model_parallel_all_reduce(output_parallel) + if output is None: + output = tensor_model_parallel_all_reduce(output_parallel) + else: + tensor_model_parallel_all_reduce_into(output_parallel, output) else: - output = output_parallel + if output is None: + output = output_parallel + else: + output.copy_(output_parallel) if not self.return_bias: return output diff --git a/vllm/model_executor/layers/mamba/gdn/kimi_gdn_linear_attn.py b/vllm/model_executor/layers/mamba/gdn/kimi_gdn_linear_attn.py index 76fd8c695446..03fa5ab26ea6 100644 --- a/vllm/model_executor/layers/mamba/gdn/kimi_gdn_linear_attn.py +++ b/vllm/model_executor/layers/mamba/gdn/kimi_gdn_linear_attn.py @@ -375,6 +375,13 @@ def forward( core_attn_out=core_attn_out, ) core_attn_out = rearrange(core_attn_out, "1 n h d -> n (h d)") + self._project_output(core_attn_out, output) + + def _project_output( + self, + core_attn_out: torch.Tensor, + output: torch.Tensor, + ) -> None: output[:] = self.o_proj(core_attn_out)[0] def _run_core( diff --git a/vllm/model_executor/layers/mla.py b/vllm/model_executor/layers/mla.py index 6c0e8e069fc7..3536417725fe 100644 --- a/vllm/model_executor/layers/mla.py +++ b/vllm/model_executor/layers/mla.py @@ -152,6 +152,7 @@ def forward( positions: torch.Tensor, hidden_states: torch.Tensor, llama_4_scaling: torch.Tensor | None = None, + output: torch.Tensor | None = None, ) -> torch.Tensor: q_c = None kv_lora = None @@ -223,4 +224,4 @@ def forward( if self.g_proj is not None: attn_out = attn_out * self.g_proj(hidden_states)[0].sigmoid() - return self.o_proj(attn_out)[0] + return self.o_proj(attn_out, output=output)[0] diff --git a/vllm/models/kimi_k3/amd/kda.py b/vllm/models/kimi_k3/amd/kda.py index 068e4ebe1c7a..4c7b08cbdcea 100644 --- a/vllm/models/kimi_k3/amd/kda.py +++ b/vllm/models/kimi_k3/amd/kda.py @@ -266,6 +266,13 @@ def _run_core( core_attn_out=core_attn_out, ) + def _project_output( + self, + core_attn_out: torch.Tensor, + output: torch.Tensor, + ) -> None: + self.o_proj(core_attn_out, output=output) + @eager_break_during_capture def _try_aiter_kda_fb_decode( self, diff --git a/vllm/models/kimi_k3/amd/linear.py b/vllm/models/kimi_k3/amd/linear.py index a8f2f9403bc2..8bfa6923e2f2 100644 --- a/vllm/models/kimi_k3/amd/linear.py +++ b/vllm/models/kimi_k3/amd/linear.py @@ -446,7 +446,7 @@ def forward( hidden_states: torch.Tensor, output: torch.Tensor, ) -> None: - output[:] = self.mla_attn(positions, hidden_states) + self.mla_attn(positions, hidden_states, output=output) class KimiDecoderLayer(nn.Module):