[diffusion] fix: stop reserving nccl device buffers for single-rank groups (#35538)
This commit is contained in:
@@ -77,7 +77,9 @@ def _patch_cache_dit_similarity():
|
|||||||
tp_sp_group = getattr(self, "_sglang_tp_sp_group", None)
|
tp_sp_group = getattr(self, "_sglang_tp_sp_group", None)
|
||||||
target_group = tp_sp_group or sp_group or tp_group
|
target_group = tp_sp_group or sp_group or tp_group
|
||||||
|
|
||||||
if target_group is None:
|
# Averaging over a one-rank group returns the input, so skip the
|
||||||
|
# collective rather than pay for a round trip that cannot change it.
|
||||||
|
if target_group is None or dist.get_world_size(target_group) == 1:
|
||||||
return _original_similarity(
|
return _original_similarity(
|
||||||
self,
|
self,
|
||||||
t1,
|
t1,
|
||||||
|
|||||||
@@ -123,6 +123,18 @@ class GraphCaptureContext:
|
|||||||
stream: torch.cuda.Stream | None
|
stream: torch.cuda.Stream | None
|
||||||
|
|
||||||
|
|
||||||
|
def new_device_group(ranks, backend=None):
|
||||||
|
"""Create a process group for device collectives.
|
||||||
|
|
||||||
|
A single-rank group never runs one: every collective short-circuits on
|
||||||
|
world_size == 1. NCCL would still allocate its per-channel device buffers
|
||||||
|
for it, which costs ~390 MiB a group.
|
||||||
|
"""
|
||||||
|
return torch.distributed.new_group(
|
||||||
|
ranks, backend="gloo" if len(ranks) == 1 else backend
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class GroupCoordinator:
|
class GroupCoordinator:
|
||||||
"""
|
"""
|
||||||
PyTorch ProcessGroup wrapper for a group of processes.
|
PyTorch ProcessGroup wrapper for a group of processes.
|
||||||
@@ -169,9 +181,7 @@ class GroupCoordinator:
|
|||||||
self.cpu_group = None
|
self.cpu_group = None
|
||||||
|
|
||||||
for ranks in group_ranks:
|
for ranks in group_ranks:
|
||||||
device_group = torch.distributed.new_group(
|
device_group = new_device_group(ranks, torch_distributed_backend)
|
||||||
ranks, backend=torch_distributed_backend
|
|
||||||
)
|
|
||||||
# a group with `gloo` backend, to allow direct coordination between
|
# a group with `gloo` backend, to allow direct coordination between
|
||||||
# processes through the CPU.
|
# processes through the CPU.
|
||||||
with suppress_stdout():
|
with suppress_stdout():
|
||||||
@@ -863,9 +873,7 @@ class PipelineGroupCoordinator(GroupCoordinator):
|
|||||||
self.device_groups = []
|
self.device_groups = []
|
||||||
if len(group_ranks[0]) > 2 or len(group_ranks[0]) == 1:
|
if len(group_ranks[0]) > 2 or len(group_ranks[0]) == 1:
|
||||||
for ranks in group_ranks:
|
for ranks in group_ranks:
|
||||||
device_group = torch.distributed.new_group(
|
device_group = new_device_group(ranks, torch_distributed_backend)
|
||||||
ranks, backend=torch_distributed_backend
|
|
||||||
)
|
|
||||||
# a group with `gloo` backend, to allow direct coordination between
|
# a group with `gloo` backend, to allow direct coordination between
|
||||||
# processes through the CPU.
|
# processes through the CPU.
|
||||||
with suppress_stdout():
|
with suppress_stdout():
|
||||||
@@ -927,9 +935,7 @@ class PipelineGroupCoordinator(GroupCoordinator):
|
|||||||
] = None
|
] = None
|
||||||
self.skip_device_group = None
|
self.skip_device_group = None
|
||||||
for ranks in group_ranks:
|
for ranks in group_ranks:
|
||||||
skip_device_group = torch.distributed.new_group(
|
skip_device_group = new_device_group(ranks, torch_distributed_backend)
|
||||||
ranks, backend=torch_distributed_backend
|
|
||||||
)
|
|
||||||
if self.rank in ranks:
|
if self.rank in ranks:
|
||||||
self.skip_device_group = skip_device_group
|
self.skip_device_group = skip_device_group
|
||||||
assert self.skip_device_group is not None
|
assert self.skip_device_group is not None
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# Reference: https://github.com/feifeibear/long-context-attention/blob/main/yunchang/globals.py
|
# Reference: https://github.com/feifeibear/long-context-attention/blob/main/yunchang/globals.py
|
||||||
|
|
||||||
|
|
||||||
import torch
|
from .group_coordinator import new_device_group
|
||||||
|
|
||||||
|
|
||||||
class Singleton:
|
class Singleton:
|
||||||
@@ -56,34 +56,34 @@ def set_seq_parallel_pg_by_sp_groups(
|
|||||||
def _map_indices_to_ranks(ranks: list[int], indices: list[int]) -> list[int]:
|
def _map_indices_to_ranks(ranks: list[int], indices: list[int]) -> list[int]:
|
||||||
return [ranks[i] for i in indices]
|
return [ranks[i] for i in indices]
|
||||||
|
|
||||||
# Important: call torch.distributed.new_group in the same order on all ranks.
|
# Important: create the groups in the same order on all ranks.
|
||||||
for sp_ranks in sp_groups:
|
for sp_ranks in sp_groups:
|
||||||
if use_ulysses_low:
|
if use_ulysses_low:
|
||||||
for i in range(num_ulysses_pgs):
|
for i in range(num_ulysses_pgs):
|
||||||
idx = list(range(i * sp_ulysses_degree, (i + 1) * sp_ulysses_degree))
|
idx = list(range(i * sp_ulysses_degree, (i + 1) * sp_ulysses_degree))
|
||||||
ulysses_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
ulysses_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
||||||
group = torch.distributed.new_group(ulysses_ranks)
|
group = new_device_group(ulysses_ranks)
|
||||||
if rank in ulysses_ranks:
|
if rank in ulysses_ranks:
|
||||||
ulyssess_pg = group
|
ulyssess_pg = group
|
||||||
|
|
||||||
for i in range(num_ring_pgs):
|
for i in range(num_ring_pgs):
|
||||||
idx = list(range(i, sp_degree, num_ring_pgs))
|
idx = list(range(i, sp_degree, num_ring_pgs))
|
||||||
ring_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
ring_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
||||||
group = torch.distributed.new_group(ring_ranks)
|
group = new_device_group(ring_ranks)
|
||||||
if rank in ring_ranks:
|
if rank in ring_ranks:
|
||||||
ring_pg = group
|
ring_pg = group
|
||||||
else:
|
else:
|
||||||
for i in range(num_ring_pgs):
|
for i in range(num_ring_pgs):
|
||||||
idx = list(range(i * sp_ring_degree, (i + 1) * sp_ring_degree))
|
idx = list(range(i * sp_ring_degree, (i + 1) * sp_ring_degree))
|
||||||
ring_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
ring_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
||||||
group = torch.distributed.new_group(ring_ranks)
|
group = new_device_group(ring_ranks)
|
||||||
if rank in ring_ranks:
|
if rank in ring_ranks:
|
||||||
ring_pg = group
|
ring_pg = group
|
||||||
|
|
||||||
for i in range(num_ulysses_pgs):
|
for i in range(num_ulysses_pgs):
|
||||||
idx = list(range(i, sp_degree, num_ulysses_pgs))
|
idx = list(range(i, sp_degree, num_ulysses_pgs))
|
||||||
ulysses_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
ulysses_ranks = _map_indices_to_ranks(sp_ranks, idx)
|
||||||
group = torch.distributed.new_group(ulysses_ranks)
|
group = new_device_group(ulysses_ranks)
|
||||||
if rank in ulysses_ranks:
|
if rank in ulysses_ranks:
|
||||||
ulyssess_pg = group
|
ulyssess_pg = group
|
||||||
|
|
||||||
|
|||||||
@@ -56,6 +56,7 @@ from .group_coordinator import (
|
|||||||
PipelineGroupCoordinator,
|
PipelineGroupCoordinator,
|
||||||
SequenceParallelGroupCoordinator,
|
SequenceParallelGroupCoordinator,
|
||||||
get_local_torch_device,
|
get_local_torch_device,
|
||||||
|
new_device_group,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
@@ -999,9 +1000,7 @@ def init_dit_group(
|
|||||||
) -> None:
|
) -> None:
|
||||||
global _DIT
|
global _DIT
|
||||||
assert _DIT is None, "DIT group is already initialized"
|
assert _DIT is None, "DIT group is already initialized"
|
||||||
_DIT = torch.distributed.new_group(
|
_DIT = new_device_group(list(range(dit_parallel_size)), backend)
|
||||||
ranks=list(range(dit_parallel_size)), backend=backend
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def get_dit_group() -> ProcessGroup:
|
def get_dit_group() -> ProcessGroup:
|
||||||
@@ -1018,7 +1017,7 @@ def init_vae_group(
|
|||||||
global _VAE
|
global _VAE
|
||||||
assert _VAE is None, "VAE parallel group is already initialized"
|
assert _VAE is None, "VAE parallel group is already initialized"
|
||||||
vae_ranks = list(range(dit_parallel_size, dit_parallel_size + vae_parallel_size))
|
vae_ranks = list(range(dit_parallel_size, dit_parallel_size + vae_parallel_size))
|
||||||
_VAE = torch.distributed.new_group(ranks=vae_ranks, backend=backend)
|
_VAE = new_device_group(vae_ranks, backend)
|
||||||
|
|
||||||
|
|
||||||
def destroy_model_parallel() -> None:
|
def destroy_model_parallel() -> None:
|
||||||
|
|||||||
@@ -0,0 +1,35 @@
|
|||||||
|
"""Single-rank groups get gloo, so NCCL does not reserve device buffers for them."""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
|
||||||
|
new_device_group,
|
||||||
|
)
|
||||||
|
|
||||||
|
NEW_GROUP_PATH = "torch.distributed.new_group"
|
||||||
|
|
||||||
|
|
||||||
|
class TestSingleRankDeviceGroup(unittest.TestCase):
|
||||||
|
def test_single_rank_group_avoids_the_device_backend(self):
|
||||||
|
for ranks, requested in [([0], "nccl"), ([3], "hccl"), ([0], None)]:
|
||||||
|
with self.subTest(ranks=ranks, requested=requested):
|
||||||
|
with patch(NEW_GROUP_PATH) as new_group:
|
||||||
|
new_device_group(ranks, requested)
|
||||||
|
new_group.assert_called_once_with(ranks, backend="gloo")
|
||||||
|
|
||||||
|
def test_multi_rank_group_keeps_the_requested_backend(self):
|
||||||
|
for ranks, requested in [([0, 1], "nccl"), ([0, 1, 2, 3], None)]:
|
||||||
|
with self.subTest(ranks=ranks, requested=requested):
|
||||||
|
with patch(NEW_GROUP_PATH) as new_group:
|
||||||
|
new_device_group(ranks, requested)
|
||||||
|
new_group.assert_called_once_with(ranks, backend=requested)
|
||||||
|
|
||||||
|
def test_backend_defaults_to_none_for_multi_rank(self):
|
||||||
|
with patch(NEW_GROUP_PATH) as new_group:
|
||||||
|
new_device_group([0, 1])
|
||||||
|
new_group.assert_called_once_with([0, 1], backend=None)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user