forked from Karylab-cklius/vllm
[PERF] Extend NCCL symmetric memory to AllGather and ReduceScatter (#46703)
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai> Co-authored-by: snordmann <snordmann@nvidia.com>
This commit is contained in:
@@ -233,7 +233,7 @@ steps:
|
||||
num_devices: 2
|
||||
commands:
|
||||
- pytest -v -s tests/distributed/test_context_parallel.py
|
||||
- pytest -v -s tests/distributed/test_nccl_symm_mem_allreduce.py
|
||||
- pytest -v -s tests/distributed/test_nccl_symm_mem.py
|
||||
- pytest -v -s tests/v1/distributed/test_dbo.py
|
||||
- pytest -v -s tests/distributed/test_mnnvl_alltoall.py
|
||||
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import random
|
||||
import typing
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
import vllm.envs as envs
|
||||
from tests.utils import ensure_current_vllm_config
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
from vllm.distributed.device_communicators.cuda_communicator import CudaCommunicator
|
||||
from vllm.distributed.device_communicators.pynccl import register_nccl_symmetric_ops
|
||||
from vllm.distributed.device_communicators.pynccl_allocator import (
|
||||
get_nccl_mem_pool,
|
||||
is_symmetric_memory_enabled,
|
||||
)
|
||||
from vllm.distributed.parallel_state import (
|
||||
get_tp_group,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.system_utils import update_environment_variables
|
||||
|
||||
torch.manual_seed(42)
|
||||
random.seed(44)
|
||||
|
||||
test_size_elements = 4 * 1024 * 1024
|
||||
|
||||
|
||||
def nccl_symm_mem_allreduce_worker(local_rank: int, world_size: int):
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
with monkeypatch.context() as m:
|
||||
m.delenv("CUDA_VISIBLE_DEVICES", raising=False)
|
||||
dtype = torch.bfloat16
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.accelerator.set_device_index(device)
|
||||
torch.set_default_device(device)
|
||||
torch.set_default_dtype(dtype)
|
||||
update_environment_variables(
|
||||
{
|
||||
"RANK": str(local_rank),
|
||||
"LOCAL_RANK": str(local_rank),
|
||||
"WORLD_SIZE": str(world_size),
|
||||
"MASTER_ADDR": "localhost",
|
||||
"MASTER_PORT": "12345",
|
||||
}
|
||||
)
|
||||
|
||||
init_distributed_environment()
|
||||
with ensure_current_vllm_config():
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
|
||||
cuda_communicator = typing.cast(
|
||||
CudaCommunicator, get_tp_group().device_communicator
|
||||
)
|
||||
pynccl_comm = cuda_communicator.pynccl_comm
|
||||
if get_nccl_mem_pool() is None:
|
||||
pytest.skip(
|
||||
"NCCL allocator compilation failed (probably missing NCCL headers)."
|
||||
)
|
||||
if not is_symmetric_memory_enabled():
|
||||
pytest.skip("NCCL symmetric memory allreduce is disabled.")
|
||||
|
||||
register_nccl_symmetric_ops(pynccl_comm)
|
||||
input = torch.randint(1, 23, (test_size_elements,), dtype=dtype, device=device)
|
||||
input_clone = input.clone()
|
||||
output = torch.ops.vllm.all_reduce_symmetric_with_copy(input)
|
||||
assert output is not None
|
||||
|
||||
group = get_tp_group().device_group
|
||||
dist.all_reduce(input_clone, group=group)
|
||||
torch.testing.assert_close(output, input_clone, atol=2.5, rtol=0.1)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.is_cuda(),
|
||||
reason="NCCLSymmMemAllreduce is only available for CUDA platforms.",
|
||||
)
|
||||
@pytest.mark.parametrize("world_size", [2])
|
||||
@pytest.mark.skipif(envs.VLLM_TARGET_DEVICE not in ["cuda"], reason="Only test on CUDA")
|
||||
def test_nccl_symm_mem_allreduce(monkeypatch: pytest.MonkeyPatch, world_size):
|
||||
if world_size > torch.accelerator.device_count():
|
||||
pytest.skip("Not enough GPUs to run the test.")
|
||||
|
||||
# Enable SymmMemCommunicator
|
||||
monkeypatch.setenv("VLLM_USE_NCCL_SYMM_MEM", "1")
|
||||
monkeypatch.setenv("NCCL_NVLS_ENABLE", "1")
|
||||
monkeypatch.setenv("NCCL_CUMEM_ENABLE", "1")
|
||||
|
||||
mp.spawn(nccl_symm_mem_allreduce_worker, args=(world_size,), nprocs=world_size)
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def nccl_symm_mem_allgather_worker(local_rank: int, world_size: int):
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
with monkeypatch.context() as m:
|
||||
m.delenv("CUDA_VISIBLE_DEVICES", raising=False)
|
||||
dtype = torch.bfloat16
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.accelerator.set_device_index(device)
|
||||
torch.set_default_device(device)
|
||||
torch.set_default_dtype(dtype)
|
||||
update_environment_variables(
|
||||
{
|
||||
"RANK": str(local_rank),
|
||||
"LOCAL_RANK": str(local_rank),
|
||||
"WORLD_SIZE": str(world_size),
|
||||
"MASTER_ADDR": "localhost",
|
||||
"MASTER_PORT": "12346",
|
||||
}
|
||||
)
|
||||
|
||||
init_distributed_environment()
|
||||
with ensure_current_vllm_config():
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
|
||||
cuda_communicator = typing.cast(
|
||||
CudaCommunicator, get_tp_group().device_communicator
|
||||
)
|
||||
if get_nccl_mem_pool() is None:
|
||||
pytest.skip(
|
||||
"NCCL allocator compilation failed (probably missing NCCL headers)."
|
||||
)
|
||||
if not is_symmetric_memory_enabled():
|
||||
pytest.skip("NCCL symmetric memory is disabled.")
|
||||
|
||||
per_rank_size = test_size_elements // world_size
|
||||
input_tensor = torch.randint(
|
||||
1, 23, (per_rank_size,), dtype=dtype, device=device
|
||||
)
|
||||
output = cuda_communicator.all_gatherv(input_tensor, dim=0)
|
||||
|
||||
group = get_tp_group().device_group
|
||||
expected = torch.empty(test_size_elements, dtype=dtype, device=device)
|
||||
dist.all_gather_into_tensor(expected, input_tensor, group=group)
|
||||
torch.testing.assert_close(output, expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.is_cuda(),
|
||||
reason="NCCL symmetric memory is only available for CUDA platforms.",
|
||||
)
|
||||
@pytest.mark.parametrize("world_size", [2])
|
||||
@pytest.mark.skipif(envs.VLLM_TARGET_DEVICE not in ["cuda"], reason="Only test on CUDA")
|
||||
def test_nccl_symm_mem_allgather(monkeypatch: pytest.MonkeyPatch, world_size):
|
||||
if world_size > torch.accelerator.device_count():
|
||||
pytest.skip("Not enough GPUs to run the test.")
|
||||
|
||||
monkeypatch.setenv("VLLM_USE_NCCL_SYMM_MEM", "1")
|
||||
monkeypatch.setenv("NCCL_NVLS_ENABLE", "1")
|
||||
monkeypatch.setenv("NCCL_CUMEM_ENABLE", "1")
|
||||
|
||||
mp.spawn(nccl_symm_mem_allgather_worker, args=(world_size,), nprocs=world_size)
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def nccl_symm_mem_reduce_scatter_worker(local_rank: int, world_size: int):
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
with monkeypatch.context() as m:
|
||||
m.delenv("CUDA_VISIBLE_DEVICES", raising=False)
|
||||
dtype = torch.bfloat16
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.accelerator.set_device_index(device)
|
||||
torch.set_default_device(device)
|
||||
torch.set_default_dtype(dtype)
|
||||
update_environment_variables(
|
||||
{
|
||||
"RANK": str(local_rank),
|
||||
"LOCAL_RANK": str(local_rank),
|
||||
"WORLD_SIZE": str(world_size),
|
||||
"MASTER_ADDR": "localhost",
|
||||
"MASTER_PORT": "12347",
|
||||
}
|
||||
)
|
||||
|
||||
init_distributed_environment()
|
||||
with ensure_current_vllm_config():
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
|
||||
cuda_communicator = typing.cast(
|
||||
CudaCommunicator, get_tp_group().device_communicator
|
||||
)
|
||||
if get_nccl_mem_pool() is None:
|
||||
pytest.skip(
|
||||
"NCCL allocator compilation failed (probably missing NCCL headers)."
|
||||
)
|
||||
if not is_symmetric_memory_enabled():
|
||||
pytest.skip("NCCL symmetric memory is disabled.")
|
||||
|
||||
per_rank_size = test_size_elements // world_size
|
||||
input_tensor = torch.randint(
|
||||
1, 23, (test_size_elements,), dtype=dtype, device=device
|
||||
)
|
||||
input_clone = input_tensor.clone()
|
||||
output = cuda_communicator.reduce_scatter(input_tensor, dim=0)
|
||||
|
||||
group = get_tp_group().device_group
|
||||
expected = torch.empty(per_rank_size, dtype=dtype, device=device)
|
||||
dist.reduce_scatter_tensor(expected, input_clone, group=group)
|
||||
torch.testing.assert_close(output, expected, atol=2.5, rtol=0.1)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.is_cuda(),
|
||||
reason="NCCL symmetric memory is only available for CUDA platforms.",
|
||||
)
|
||||
@pytest.mark.parametrize("world_size", [2])
|
||||
@pytest.mark.skipif(envs.VLLM_TARGET_DEVICE not in ["cuda"], reason="Only test on CUDA")
|
||||
def test_nccl_symm_mem_reduce_scatter(monkeypatch: pytest.MonkeyPatch, world_size):
|
||||
if world_size > torch.accelerator.device_count():
|
||||
pytest.skip("Not enough GPUs to run the test.")
|
||||
|
||||
monkeypatch.setenv("VLLM_USE_NCCL_SYMM_MEM", "1")
|
||||
monkeypatch.setenv("NCCL_NVLS_ENABLE", "1")
|
||||
monkeypatch.setenv("NCCL_CUMEM_ENABLE", "1")
|
||||
|
||||
mp.spawn(nccl_symm_mem_reduce_scatter_worker, args=(world_size,), nprocs=world_size)
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -1,96 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import random
|
||||
import typing
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
|
||||
import vllm.envs as envs
|
||||
from tests.utils import ensure_current_vllm_config
|
||||
from vllm.distributed import cleanup_dist_env_and_memory
|
||||
from vllm.distributed.device_communicators.cuda_communicator import CudaCommunicator
|
||||
from vllm.distributed.device_communicators.pynccl import register_nccl_symmetric_ops
|
||||
from vllm.distributed.device_communicators.pynccl_allocator import (
|
||||
get_nccl_mem_pool,
|
||||
is_symmetric_memory_enabled,
|
||||
)
|
||||
from vllm.distributed.parallel_state import (
|
||||
get_tp_group,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.utils.system_utils import update_environment_variables
|
||||
|
||||
torch.manual_seed(42)
|
||||
random.seed(44)
|
||||
|
||||
test_size_elements = 4 * 1024 * 1024
|
||||
|
||||
|
||||
def nccl_symm_mem_allreduce_worker(local_rank: int, world_size: int):
|
||||
monkeypatch = pytest.MonkeyPatch()
|
||||
with monkeypatch.context() as m:
|
||||
m.delenv("CUDA_VISIBLE_DEVICES", raising=False)
|
||||
dtype = torch.bfloat16
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.accelerator.set_device_index(device)
|
||||
torch.set_default_device(device)
|
||||
torch.set_default_dtype(dtype)
|
||||
update_environment_variables(
|
||||
{
|
||||
"RANK": str(local_rank),
|
||||
"LOCAL_RANK": str(local_rank),
|
||||
"WORLD_SIZE": str(world_size),
|
||||
"MASTER_ADDR": "localhost",
|
||||
"MASTER_PORT": "12345",
|
||||
}
|
||||
)
|
||||
|
||||
init_distributed_environment()
|
||||
with ensure_current_vllm_config():
|
||||
initialize_model_parallel(tensor_model_parallel_size=world_size)
|
||||
|
||||
cuda_communicator = typing.cast(
|
||||
CudaCommunicator, get_tp_group().device_communicator
|
||||
)
|
||||
pynccl_comm = cuda_communicator.pynccl_comm
|
||||
if get_nccl_mem_pool() is None:
|
||||
pytest.skip(
|
||||
"NCCL allocator compilation failed (probably missing NCCL headers)."
|
||||
)
|
||||
if not is_symmetric_memory_enabled():
|
||||
pytest.skip("NCCL symmetric memory allreduce is disabled.")
|
||||
|
||||
register_nccl_symmetric_ops(pynccl_comm)
|
||||
input = torch.randint(1, 23, (test_size_elements,), dtype=dtype, device=device)
|
||||
input_clone = input.clone()
|
||||
output = torch.ops.vllm.all_reduce_symmetric_with_copy(input)
|
||||
assert output is not None
|
||||
|
||||
group = get_tp_group().device_group
|
||||
dist.all_reduce(input_clone, group=group)
|
||||
torch.testing.assert_close(output, input_clone, atol=2.5, rtol=0.1)
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
not current_platform.is_cuda(),
|
||||
reason="NCCLSymmMemAllreduce is only available for CUDA platforms.",
|
||||
)
|
||||
@pytest.mark.parametrize("world_size", [2])
|
||||
@pytest.mark.skipif(envs.VLLM_TARGET_DEVICE not in ["cuda"], reason="Only test on CUDA")
|
||||
def test_nccl_symm_mem_allreduce(monkeypatch: pytest.MonkeyPatch, world_size):
|
||||
if world_size > torch.accelerator.device_count():
|
||||
pytest.skip("Not enough GPUs to run the test.")
|
||||
|
||||
# Enable SymmMemCommunicator
|
||||
monkeypatch.setenv("VLLM_USE_NCCL_SYMM_MEM", "1")
|
||||
monkeypatch.setenv("NCCL_NVLS_ENABLE", "1")
|
||||
monkeypatch.setenv("NCCL_CUMEM_ENABLE", "1")
|
||||
|
||||
mp.spawn(nccl_symm_mem_allreduce_worker, args=(world_size,), nprocs=world_size)
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -134,6 +134,18 @@ def should_nccl_symm_mem_allreduce(world_size: int, input_tensor: torch.Tensor)
|
||||
return world_size > NCCL_SYMM_MEM_ALL_REDUCE_CONFIG["always_use_above_world_size"]
|
||||
|
||||
|
||||
def should_nccl_symm_mem_ag_rs() -> bool:
|
||||
"""Check whether NCCL symmetric memory should be used for
|
||||
AllGather / ReduceScatter collectives."""
|
||||
from vllm.distributed.device_communicators.pynccl_allocator import (
|
||||
is_symmetric_memory_enabled,
|
||||
)
|
||||
|
||||
if envs.VLLM_BATCH_INVARIANT:
|
||||
return False
|
||||
return is_symmetric_memory_enabled()
|
||||
|
||||
|
||||
def producer(
|
||||
batch_src: Sequence[int],
|
||||
producer_queue,
|
||||
|
||||
@@ -8,6 +8,7 @@ from torch.distributed import ProcessGroup
|
||||
import vllm.envs as envs
|
||||
from vllm.distributed.device_communicators.all_reduce_utils import (
|
||||
NCCL_SYMM_MEM_ALL_REDUCE_CONFIG,
|
||||
should_nccl_symm_mem_ag_rs,
|
||||
should_nccl_symm_mem_allreduce,
|
||||
)
|
||||
from vllm.distributed.device_communicators.pynccl import register_nccl_symmetric_ops
|
||||
@@ -310,6 +311,17 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
torch.distributed.all_reduce(out, group=self.device_group)
|
||||
return out
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
# Route uniform dim-0 all-gathers through NVLS symmetric memory when
|
||||
# enabled (mirrors reduce_scatter); otherwise fall back to the
|
||||
# base-class ring all-gather. Sequence parallelism's gather-before-GEMM
|
||||
# uses dim=0 with tp-aligned (uniform) shards.
|
||||
if dim < 0:
|
||||
dim += input_.dim()
|
||||
if dim == 0 and should_nccl_symm_mem_ag_rs():
|
||||
return self._all_gather_symm_mem(input_.contiguous())
|
||||
return super().all_gather(input_, dim)
|
||||
|
||||
def reduce_scatter(self, input_: torch.Tensor, dim: int = -1):
|
||||
world_size = self.world_size
|
||||
pynccl_comm = self.pynccl_comm
|
||||
@@ -326,11 +338,13 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
chunk_size = input_tensor.shape[0] // world_size
|
||||
output_shape = (chunk_size,) + input_tensor.shape[1:]
|
||||
|
||||
output = torch.empty(
|
||||
output_shape, dtype=input_tensor.dtype, device=input_tensor.device
|
||||
)
|
||||
|
||||
pynccl_comm.reduce_scatter(output, input_tensor)
|
||||
if should_nccl_symm_mem_ag_rs():
|
||||
output = self._reduce_scatter_symm_mem(input_tensor)
|
||||
else:
|
||||
output = torch.empty(
|
||||
output_shape, dtype=input_tensor.dtype, device=input_tensor.device
|
||||
)
|
||||
pynccl_comm.reduce_scatter(output, input_tensor)
|
||||
|
||||
# Reshape before returning
|
||||
return output.movedim(0, dim).contiguous()
|
||||
@@ -358,18 +372,97 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
chunk_size = input_tensor.shape[0] // world_size
|
||||
output_shape = (chunk_size,) + input_tensor.shape[1:]
|
||||
|
||||
output = torch.empty(
|
||||
output_shape, dtype=input_tensor.dtype, device=input_tensor.device
|
||||
)
|
||||
|
||||
if sizes is not None and sizes.count(sizes[0]) != len(sizes):
|
||||
pynccl_comm.reduce_scatterv(output, input_tensor, sizes=sizes)
|
||||
# Symmetric memory is only used when all ranks have uniform sizes.
|
||||
# ncclCommWindowRegister is collective: asymmetric pool allocations
|
||||
# from variable per-rank sizes cause deadlocks.
|
||||
use_symm_mem = sizes is None and should_nccl_symm_mem_ag_rs()
|
||||
if use_symm_mem:
|
||||
output = self._reduce_scatter_symm_mem(input_tensor)
|
||||
else:
|
||||
pynccl_comm.reduce_scatter(output, input_tensor)
|
||||
output = torch.empty(
|
||||
output_shape, dtype=input_tensor.dtype, device=input_tensor.device
|
||||
)
|
||||
if sizes is not None and sizes.count(sizes[0]) != len(sizes):
|
||||
pynccl_comm.reduce_scatterv(output, input_tensor, sizes=sizes)
|
||||
else:
|
||||
pynccl_comm.reduce_scatter(output, input_tensor)
|
||||
|
||||
# Reshape before returning
|
||||
return output.movedim(0, dim).contiguous()
|
||||
|
||||
def _get_symm_scratch(
|
||||
self,
|
||||
role: str,
|
||||
shape: tuple[int, ...],
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> torch.Tensor:
|
||||
"""Persistent, pre-registered NCCL symmetric-memory scratch buffer.
|
||||
|
||||
Allocating a fresh symm tensor per collective pays the
|
||||
``nccl_symm_mem_context`` snapshot + window-registration scan on every
|
||||
call (~0.5 ms/RS+AG pair, dwarfing the NVLS transfer itself). Instead we
|
||||
allocate once per ``(role, shape, dtype)``, register once, and reuse.
|
||||
|
||||
Safe for serial (eager) sequence parallelism: each collective's result
|
||||
is consumed on the same stream before the next same-role collective
|
||||
reuses the buffer. Distinct roles (e.g. ``rs_in`` vs ``ag_out``, both
|
||||
full-size) get distinct buffers so a reduce-scatter input copy never
|
||||
clobbers a still-live all-gather output.
|
||||
"""
|
||||
from vllm.distributed.device_communicators.pynccl_allocator import (
|
||||
nccl_symm_mem_context,
|
||||
)
|
||||
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
cache = self.__dict__.setdefault("_symm_scratch_bufs", {})
|
||||
key = (role, tuple(shape), dtype)
|
||||
buf = cache.get(key)
|
||||
if buf is None:
|
||||
with nccl_symm_mem_context(pynccl_comm):
|
||||
buf = torch.empty(shape, dtype=dtype, device=device)
|
||||
cache[key] = buf
|
||||
return buf
|
||||
|
||||
def _reduce_scatter_symm_mem(
|
||||
self,
|
||||
input_tensor: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""ReduceScatter using NCCL symmetric memory (NVLS).
|
||||
|
||||
Only called for uniform-size reduce_scatter (variable sizes are
|
||||
guarded out by the caller to avoid asymmetric ncclCommWindowRegister).
|
||||
Uses persistent pre-registered scratch (see _get_symm_scratch).
|
||||
"""
|
||||
from vllm.distributed.device_communicators.pynccl_allocator import (
|
||||
is_symmetric_memory_tensor,
|
||||
)
|
||||
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
|
||||
chunk = input_tensor.shape[0] // self.world_size
|
||||
output_shape = (chunk,) + tuple(input_tensor.shape[1:])
|
||||
|
||||
symm_output = self._get_symm_scratch(
|
||||
"rs_out", output_shape, input_tensor.dtype, input_tensor.device
|
||||
)
|
||||
# NVLS reduce-scatter (LDMC) requires the input in symmetric memory.
|
||||
if is_symmetric_memory_tensor(input_tensor):
|
||||
symm_input = input_tensor
|
||||
else:
|
||||
symm_input = self._get_symm_scratch(
|
||||
"rs_in",
|
||||
tuple(input_tensor.shape),
|
||||
input_tensor.dtype,
|
||||
input_tensor.device,
|
||||
)
|
||||
symm_input.copy_(input_tensor)
|
||||
|
||||
pynccl_comm.reduce_scatter(symm_output, symm_input)
|
||||
return symm_output
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: int | None = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
@@ -440,6 +533,14 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
if sizes is not None and all(s == sizes[0] for s in sizes):
|
||||
sizes = None
|
||||
|
||||
# Symmetric memory is only used when all ranks have uniform sizes.
|
||||
# ncclCommWindowRegister is collective: asymmetric pool allocations
|
||||
# from variable per-rank sizes cause deadlocks.
|
||||
if sizes is None and should_nccl_symm_mem_ag_rs():
|
||||
if isinstance(input_, torch.Tensor):
|
||||
return self._all_gather_symm_mem(input_)
|
||||
return self._all_gather_batched_symm_mem(input_)
|
||||
|
||||
def _all_gather_single(input_: torch.Tensor, sizes: list[int] | None = None):
|
||||
input_size = input_.size()
|
||||
if sizes is not None:
|
||||
@@ -471,6 +572,56 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
|
||||
return output_list
|
||||
|
||||
def _all_gather_symm_mem(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
"""AllGather a single tensor using NCCL symmetric memory (NVLS).
|
||||
|
||||
Only the output needs to be in symmetric memory; NCCL does not
|
||||
require the AG input to be symmetrically allocated.
|
||||
"""
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
|
||||
out_size = (input_.size(0) * self.world_size,) + tuple(input_.size()[1:])
|
||||
# Persistent pre-registered scratch avoids the per-call symm-mem context
|
||||
# snapshot/registration overhead (see _get_symm_scratch).
|
||||
symm_output = self._get_symm_scratch(
|
||||
"ag_out", out_size, input_.dtype, input_.device
|
||||
)
|
||||
pynccl_comm.all_gather(symm_output, input_)
|
||||
return symm_output
|
||||
|
||||
def _all_gather_batched_symm_mem(
|
||||
self, inputs: list[torch.Tensor]
|
||||
) -> list[torch.Tensor]:
|
||||
"""AllGather a list of tensors using NCCL symmetric memory (NVLS).
|
||||
|
||||
Uses group_start/group_end to batch the collectives.
|
||||
Only the output needs to be in symmetric memory (see
|
||||
_all_gather_symm_mem).
|
||||
"""
|
||||
from vllm.distributed.device_communicators.pynccl_allocator import (
|
||||
nccl_symm_mem_context,
|
||||
)
|
||||
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
world_size = self.world_size
|
||||
|
||||
symm_outputs = []
|
||||
with nccl_symm_mem_context(pynccl_comm):
|
||||
for inp in inputs:
|
||||
out_size = (inp.size(0) * world_size,) + inp.size()[1:]
|
||||
symm_outputs.append(
|
||||
torch.empty(out_size, dtype=inp.dtype, device=inp.device)
|
||||
)
|
||||
|
||||
pynccl_comm.group_start()
|
||||
for symm_out, inp in zip(symm_outputs, inputs):
|
||||
pynccl_comm.all_gather(symm_out, inp)
|
||||
pynccl_comm.group_end()
|
||||
|
||||
return symm_outputs
|
||||
|
||||
def dispatch_router_logits(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
|
||||
@@ -39,7 +39,7 @@ void nccl_free_plug(void* ptr, size_t size, int device, void* stream) {
|
||||
_allocator = None
|
||||
_allocator_wrapper = None
|
||||
_mem_pool = None
|
||||
_registered_base_addrs = set()
|
||||
_registered_base_addrs: dict[bytes, set] = {}
|
||||
_graph_pool_id = None
|
||||
_nccl_allocator_failed_to_compile = False
|
||||
_cached_pool_snapshot = None
|
||||
@@ -181,11 +181,14 @@ class nccl_symm_mem_context:
|
||||
assert _pool is not None
|
||||
_cached_pool_snapshot = _pool.snapshot()
|
||||
assert self.pynccl_comm is not None
|
||||
comm_key = bytes(self.pynccl_comm.unique_id.internal)
|
||||
if comm_key not in _registered_base_addrs:
|
||||
_registered_base_addrs[comm_key] = set()
|
||||
for segment in _cached_pool_snapshot:
|
||||
if segment["address"] not in _registered_base_addrs:
|
||||
if segment["address"] not in _registered_base_addrs[comm_key]:
|
||||
self.pynccl_comm.register_comm_window_raw(
|
||||
segment["address"], segment["total_size"]
|
||||
)
|
||||
_registered_base_addrs.add(segment["address"])
|
||||
_registered_base_addrs[comm_key].add(segment["address"])
|
||||
if self.is_graph_capture:
|
||||
torch._C._cuda_beginAllocateCurrentThreadToPool(self.device, _graph_pool_id)
|
||||
|
||||
Reference in New Issue
Block a user