Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
207 changes: 207 additions & 0 deletions tests/distributed/test_all_reduce_into.py
Original file line number Diff line number Diff line change
@@ -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)
9 changes: 9 additions & 0 deletions vllm/distributed/communication_op.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
21 changes: 19 additions & 2 deletions vllm/distributed/device_communicators/aiter_custom_all_reduce.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
23 changes: 23 additions & 0 deletions vllm/distributed/device_communicators/cuda_communicator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
34 changes: 34 additions & 0 deletions vllm/distributed/parallel_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
23 changes: 21 additions & 2 deletions vllm/model_executor/layers/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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_
Expand All @@ -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
Expand Down
7 changes: 7 additions & 0 deletions vllm/model_executor/layers/mamba/gdn/kimi_gdn_linear_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
3 changes: 2 additions & 1 deletion vllm/model_executor/layers/mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]
Loading
Loading