Name the two widths of the WORLD group (#40070)

This commit is contained in:
Cheng Wan
2026-09-18 17:47:02 -07:00
committed by GitHub
parent 81421b91e9
commit fa7e83fd09
13 changed files with 178 additions and 45 deletions
@@ -65,7 +65,8 @@ class AscendTransferEngine(MooncakeTransferEngine):
# through all_gather to avoid conflicts with rdma initialization. # through all_gather to avoid conflicts with rdma initialization.
tmp_tensor = torch.zeros(1, device="npu") tmp_tensor = torch.zeros(1, device="npu")
output_tensor_list = [ output_tensor_list = [
torch.empty_like(tmp_tensor) for _ in range(get_parallel().world_size) torch.empty_like(tmp_tensor)
for _ in range(get_parallel().launch_world_size)
] ]
torch.distributed.all_gather( torch.distributed.all_gather(
output_tensor_list, output_tensor_list,
+1 -1
View File
@@ -92,7 +92,7 @@ class ElasticEPStateManager:
if get_exec().moe.elastic_ep_backend is not None: if get_exec().moe.elastic_ep_backend is not None:
world_size = torch.distributed.get_world_size() world_size = torch.distributed.get_world_size()
active_rank_capacity = get_parallel().max_ep_size or world_size active_rank_capacity = get_parallel().max_world_size
assert active_rank_capacity >= world_size, ( assert active_rank_capacity >= world_size, (
f"--max-ep-size ({active_rank_capacity}) must be >= " f"--max-ep-size ({active_rank_capacity}) must be >= "
f"world_size ({world_size})." f"world_size ({world_size})."
@@ -50,7 +50,7 @@ class ExpertBackupClient:
self.buffer_size = 0 self.buffer_size = 0
self.use_backup = False self.use_backup = False
local_ip = get_local_ip_auto() local_ip = get_local_ip_auto()
all_ips = [None] * get_parallel().world_size all_ips = [None] * get_parallel().launch_world_size
torch.distributed.all_gather_object( torch.distributed.all_gather_object(
all_ips, local_ip, group=get_parallel().world_group.cpu_group all_ips, local_ip, group=get_parallel().world_group.cpu_group
) )
@@ -59,14 +59,14 @@ class ExpertBackupClient:
for i in range(self.engine_num): for i in range(self.engine_num):
self.recv_list[i] = context.socket(zmq.SUB) self.recv_list[i] = context.socket(zmq.SUB)
self.recv_list[i].connect( self.recv_list[i].connect(
f"tcp://{all_ips[i * get_parallel().world_size // get_parallel().nnodes]}:{PORT_BASE + i * 2 + 1}" f"tcp://{all_ips[i * get_parallel().launch_world_size // get_parallel().nnodes]}:{PORT_BASE + i * 2 + 1}"
) )
self.recv_list[i].setsockopt(zmq.SUBSCRIBE, b"") self.recv_list[i].setsockopt(zmq.SUBSCRIBE, b"")
# Synchronization channel to notify the manager when this client is ready. # Synchronization channel to notify the manager when this client is ready.
self.ready_sockets[i] = context.socket(zmq.PUSH) self.ready_sockets[i] = context.socket(zmq.PUSH)
self.ready_sockets[i].connect( self.ready_sockets[i].connect(
f"tcp://{all_ips[i * get_parallel().world_size // get_parallel().nnodes]}:{PORT_BASE + i * 2}" f"tcp://{all_ips[i * get_parallel().launch_world_size // get_parallel().nnodes]}:{PORT_BASE + i * 2}"
) )
sock_send(self.ready_sockets[i], UpdateExpertBackupReq()) sock_send(self.ready_sockets[i], UpdateExpertBackupReq())
+1 -1
View File
@@ -5665,7 +5665,7 @@ class Scheduler(
old_ep_size = ElasticEPStateManager.get_effective_ep_size() old_ep_size = ElasticEPStateManager.get_effective_ep_size()
new_ep_size = recv_req.new_ep_size new_ep_size = recv_req.new_ep_size
max_ep_size = get_parallel().max_ep_size or old_ep_size max_ep_size = get_parallel().max_world_size
logger.debug( logger.debug(
"[Elastic EP][scale] request received: new_ep_size=%d " "[Elastic EP][scale] request received: new_ep_size=%d "
@@ -60,13 +60,17 @@ def _resolve_elastic_world_dp_size(
live_dp_size = get_parallel().attn_dp_size live_dp_size = get_parallel().attn_dp_size
effective_ep_size = ElasticEPStateManager.get_effective_ep_size() effective_ep_size = ElasticEPStateManager.get_effective_ep_size()
# The group's own membership, not the width it was built at: this is the
# one number an out-of-process join moves, and it is the upper bound the
# served width has to stay under.
world_size = torch.distributed.get_world_size(group) world_size = torch.distributed.get_world_size(group)
if live_dp_size != effective_ep_size: if live_dp_size != effective_ep_size:
raise RuntimeError( raise RuntimeError(
"[Elastic EP] WORLD MLP sync dp_size is out of sync: " "[Elastic EP] WORLD MLP sync dp_size is out of sync: "
f"rank={torch.distributed.get_rank(group)} " f"rank={torch.distributed.get_rank(group)} "
f"live_dp_size={live_dp_size} effective_ep_size={effective_ep_size} " f"live_dp_size={live_dp_size} "
f"effective_ep_size={effective_ep_size} "
f"world_size={world_size} server_args_dp_size={dp_size} " f"world_size={world_size} server_args_dp_size={dp_size} "
f"local_num_tokens={local_num_tokens} " f"local_num_tokens={local_num_tokens} "
f"local_forward_mode={local_forward_mode}" f"local_forward_mode={local_forward_mode}"
@@ -2195,7 +2195,7 @@ class KVCacheConfigurator:
available_gpu_memory = get_available_gpu_memory( available_gpu_memory = get_available_gpu_memory(
self.device, self.device,
self.gpu_id, self.gpu_id,
distributed=get_parallel().world_group.world_size > 1, distributed=get_parallel().launch_world_size > 1,
cpu_group=get_parallel().world_group.cpu_group, cpu_group=get_parallel().world_group.cpu_group,
) )
@@ -31,20 +31,21 @@ _WRITE_BACK_STAGING_PAGE_CHUNK = 64
def ranks_per_host() -> int: def ranks_per_host() -> int:
"""Number of ranks of this job running on the same machine as this one. """Number of ranks of this job running on the same machine as this one.
Derived as world_size // nnodes: the launcher slices ranks uniformly Derived as the launch width // nnodes: the launcher slices ranks
across nodes (resolution asserts divisibility), so no hostname collective uniformly across nodes (resolution asserts divisibility), so no hostname
is needed — a collective here would have to be issued the same number of collective is needed — a collective here would have to be issued the same
times on every rank, and ranks build different numbers of host pools. number of times on every rank, and ranks build different numbers of host
pools.
""" """
if not (torch.distributed.is_available() and torch.distributed.is_initialized()): if not (torch.distributed.is_available() and torch.distributed.is_initialized()):
return 1 return 1
try: try:
world_group = get_parallel().world_group launch_world_size = get_parallel().launch_world_size
except AssertionError: except AssertionError:
return 1 return 1
if world_group.world_size == 1: if launch_world_size == 1:
return 1 return 1
return max(world_group.world_size // get_parallel().nnodes, 1) return max(launch_world_size // get_parallel().nnodes, 1)
def host_memory_budget_bytes() -> int: def host_memory_budget_bytes() -> int:
@@ -146,7 +146,7 @@ def build_attention_backends(*, model_runner: ModelRunner) -> AttentionBackends:
model_runner.device, model_runner.device,
model_runner.gpu_id, model_runner.gpu_id,
get_parallel().world_group.rank_in_group, get_parallel().world_group.rank_in_group,
get_parallel().world_group.world_size, get_parallel().launch_world_size,
get_parallel().world_group.cpu_group, get_parallel().world_group.cpu_group,
) )
@@ -227,7 +227,7 @@ def refresh_deep_gemm_layout_memory_budget(
available_memory_gb = get_available_gpu_memory( available_memory_gb = get_available_gpu_memory(
model_runner.device, model_runner.device,
model_runner.gpu_id, model_runner.gpu_id,
distributed=world_group.world_size > 1, distributed=get_parallel().launch_world_size > 1,
cpu_group=world_group.cpu_group, cpu_group=world_group.cpu_group,
) )
budget_bytes = set_masked_standard_layout_memory_budget( budget_bytes = set_masked_standard_layout_memory_budget(
@@ -59,7 +59,7 @@ def compute_post_capture_kv_resize(
free_gb = get_available_gpu_memory( free_gb = get_available_gpu_memory(
model_runner.device, model_runner.device,
model_runner.gpu_id, model_runner.gpu_id,
distributed=get_parallel().world_group.world_size > 1, distributed=get_parallel().launch_world_size > 1,
cpu_group=get_parallel().world_group.cpu_group, cpu_group=get_parallel().world_group.cpu_group,
) )
headroom_gb = model_runner.pre_model_load_memory * ( headroom_gb = model_runner.pre_model_load_memory * (
@@ -1008,7 +1008,7 @@ def _prefetch_all_checkpoints(
if torch.distributed.is_initialized(): if torch.distributed.is_initialized():
world_group = get_parallel().world_group world_group = get_parallel().world_group
local_rank = world_group.local_rank local_rank = world_group.local_rank
local_world_size = world_group.local_size or world_group.world_size local_world_size = world_group.local_size or get_parallel().launch_world_size
else: else:
local_rank = 0 local_rank = 0
local_world_size = 1 local_world_size = 1
+72 -10
View File
@@ -117,9 +117,57 @@ def _parallel_config_leaves() -> frozenset:
# this process's attention-DP rank. # this process's attention-DP rank.
_MISSING_READ = object() _MISSING_READ = object()
class Live(msgspec.Struct, frozen=True):
"""How a rank / group / world width is answered, and what it means.
`source` is the canonical getter's name in `parallel_state`, a callable
taking the context, or `None` for a name only a stamp can answer.
Most entries in the table below are a bare getter name: a rank or a group
handle is its own explanation. This shape is for a name whose meaning is
not in its getter, and it carries the prose with the declaration rather
than in a second table keyed by the same names.
"""
source: Any = None
doc: str = ""
_LIVE_READS: dict = { _LIVE_READS: dict = {
"world_size": "get_world_size", # Two widths of the WORLD group: what it was built at, and what it has
"world_rank": "get_world_rank", # room for. Both are properties of the group itself. How much of that room
# is currently serving is elastic-EP state, owned by `ElasticEPStateManager`
# and asked of it directly -- a width that lives somewhere else does not
# become a WORLD fact by being readable from here.
"launch_world_size": Live(
source="get_world_size",
doc=(
"Width the WORLD group was built at: `len(ranks)`, frozen when the "
"coordinator was constructed. What every startup reader wants -- "
"memory accounting, KV cache sizing, graph capture, weight loading "
"-- and what a scale-up leaves behind rather than updates."
),
),
"max_world_size": Live(
source=lambda self: self.max_ep_size or self.launch_world_size,
doc=(
"Ranks the WORLD group has room for: `--max-ep-size` when it is "
"set, otherwise the launch width. This is the ceiling the process "
"group was pre-allocated to -- mooncake sizes its active-rank mask "
"to it -- which is why `init_distributed_environment` takes it "
"under this name. Whether the group can grow at all is a separate "
"question, answered by the leaf being set rather than by this width."
),
),
"launch_world_rank": Live(
source="get_world_rank",
doc=(
"This process's rank in the WORLD group as built. Frozen with the "
"coordinator, exactly like `launch_world_size`, and named for the "
"same reason: a scale-up does not renumber it."
),
),
"tp_rank": "get_tensor_model_parallel_rank", "tp_rank": "get_tensor_model_parallel_rank",
"pp_rank": "get_pipeline_model_parallel_rank", "pp_rank": "get_pipeline_model_parallel_rank",
"moe_ep_rank": "get_moe_expert_parallel_rank", "moe_ep_rank": "get_moe_expert_parallel_rank",
@@ -195,9 +243,15 @@ def derive_parallel_widths(
the arithmetic lives here rather than being read back off the group the arithmetic lives here rather than being read back off the group
coordinators. coordinators.
`world_size` is not among them: it is not a quotient, and `get_world_size()` The world widths are not among them: neither is a quotient, and they are
answers with the live WORLD group, which stays right through an elastic not one number. `launch_world_size` is what the WORLD group was built at
scale-up that a value fixed at group build would not survive. and is frozen there -- `GroupCoordinator.world_size` is `len(ranks)`, so it
does not move when mooncake admits ranks into an expandable WORLD;
`max_world_size` is what that group has room for. How much of that room is
serving right now is elastic-EP state and is asked of its owner. Deriving
either from the leaves would be wrong in a further way: on a scale joiner it
would answer with the joining cohort's own `tp * pp`, while that process's
WORLD spans `ep_join_rank_offset + tp * pp`.
""" """
return { return {
"attn_dp_size": attn_dp_size, "attn_dp_size": attn_dp_size,
@@ -335,10 +389,11 @@ class ParallelContext:
return getattr(config, name) return getattr(config, name)
live = _LIVE_READS.get(name, _MISSING_READ) live = _LIVE_READS.get(name, _MISSING_READ)
if live is not _MISSING_READ: if live is not _MISSING_READ:
if isinstance(live, str): source = live.source if isinstance(live, Live) else live
return getattr(_ps(), live)() if isinstance(source, str):
if live is not None: return getattr(_ps(), source)()
return live(self) if source is not None:
return source(self)
raise RuntimeError( raise RuntimeError(
f"parallel rank {name!r} is not available: it is computed from " f"parallel rank {name!r} is not available: it is computed from "
"this process's `tp_rank` when the attention topology is " "this process's `tp_rank` when the attention topology is "
@@ -422,8 +477,15 @@ def _install_parallel_properties() -> None:
chain rather than one per kind of name. chain rather than one per kind of name.
""" """
docs = {name: decl.doc for name, decl in _derived_widths().items()} docs = {name: decl.doc for name, decl in _derived_widths().items()}
docs.update(
{
name: live.doc
for name, live in _LIVE_READS.items()
if isinstance(live, Live) and live.doc
}
)
for name in list(docs) + list(_LIVE_READS): for name in list(_derived_widths()) + list(_LIVE_READS):
def getter(self, _name=name): def getter(self, _name=name):
return self._read(_name) return self._read(_name)
+81 -16
View File
@@ -61,16 +61,16 @@ _SRT = _pathlib.Path(next(iter(_sglang.__path__))).resolve() / "srt"
_PS = "sglang.srt.distributed.parallel_state" _PS = "sglang.srt.distributed.parallel_state"
_DP = "sglang.srt.layers.dp_attention" _DP = "sglang.srt.layers.dp_attention"
# Ranks and the world size read the live group: they are not implied by # Ranks and the launch width are asked of the group: they are not implied by
# anything, so there is nothing to derive them from. The quotients used to be # anything, so there is nothing to derive them from. The quotients are not
# in this table and are not any more -- `attn_tp_size` and its siblings are # here -- `attn_tp_size` and its siblings are functions of the configured
# functions of the configured leaves, and `TestDerivedWidthsComeFromTheLeaves` # leaves, and `TestDerivedWidths` pins them. `attn_dp_rank` is not here either: no group coordinator
# is what pins them. `attn_dp_rank` is not here either: no group coordinator
# knows it, so it is stamped when the attention topology is initialized and # knows it, so it is stamped when the attention topology is initialized and
# `TestStampedRanks` is what pins it. # `TestStampedRanks` is what pins it. The other world width is not here
# because the group does not know it; `TestTheTwoWorldWidths` pins it.
SIZE_RANK_DELEGATIONS = [ SIZE_RANK_DELEGATIONS = [
("world_size", f"{_PS}.get_world_size"), ("launch_world_size", f"{_PS}.get_world_size"),
("world_rank", f"{_PS}.get_world_rank"), ("launch_world_rank", f"{_PS}.get_world_rank"),
("tp_rank", f"{_PS}.get_tensor_model_parallel_rank"), ("tp_rank", f"{_PS}.get_tensor_model_parallel_rank"),
("dcp_rank", f"{_PS}.get_dcp_rank"), ("dcp_rank", f"{_PS}.get_dcp_rank"),
("pp_rank", f"{_PS}.get_pipeline_model_parallel_rank"), ("pp_rank", f"{_PS}.get_pipeline_model_parallel_rank"),
@@ -151,6 +151,55 @@ class TestParallelDelegation(_IsolatedOverrides):
self.assertFalse(hasattr(ParallelContext, "local_attn_dp_size")) self.assertFalse(hasattr(ParallelContext, "local_attn_dp_size"))
class TestTheTwoWorldWidths(_IsolatedOverrides):
"""Two questions about the WORLD group: what it was built at, and what it
has room for.
Neither is stored here. How much of that room is serving after a scale-up
is elastic-EP state, and is asked of the manager that owns it rather than
mirrored onto this namespace.
"""
def test_the_launch_width_is_what_the_group_was_built_at(self):
with patch(f"{_PS}.get_world_size", return_value=4):
self.assertEqual(get_parallel().launch_world_size, 4)
def test_the_ceiling_is_the_configured_one_when_there_is_one(self):
parallel = get_parallel()
with (
parallel.override(max_ep_size=32),
patch(
f"{_PS}.get_world_size",
side_effect=AssertionError("the built group must not be asked"),
),
):
self.assertEqual(parallel.max_world_size, 32)
def test_without_a_configured_ceiling_the_room_is_the_launch_width(self):
parallel = get_parallel()
with (
parallel.override(max_ep_size=None),
patch(f"{_PS}.get_world_size", return_value=8),
):
self.assertEqual(parallel.max_world_size, 8)
def test_each_width_can_be_stated_on_its_own(self):
"""Stating one must not answer for the other: they are two names."""
parallel = get_parallel()
with (
parallel.override(launch_world_size=2, max_ep_size=None),
patch(
f"{_PS}.get_world_size",
side_effect=AssertionError("the built group must not be asked"),
),
):
self.assertEqual(parallel.launch_world_size, 2)
self.assertEqual(parallel.max_world_size, 2)
with parallel.override(max_ep_size=6):
self.assertEqual(parallel.max_world_size, 6)
self.assertEqual(parallel.launch_world_size, 2)
class TestStampedRanks(_IsolatedOverrides): class TestStampedRanks(_IsolatedOverrides):
"""`attn_dp_rank` comes from the stamp, and says so when there is none. """`attn_dp_rank` comes from the stamp, and says so when there is none.
@@ -1751,12 +1800,12 @@ class TestDerivedWidths(_IsolatedOverrides):
self.assertEqual(widths["moe_tp_size"], 8 // 4 // 2) self.assertEqual(widths["moe_tp_size"], 8 // 4 // 2)
self.assertEqual(widths["attn_dcp_size"], 1) self.assertEqual(widths["attn_dcp_size"], 1)
def test_the_world_size_is_not_permanently_overridden(self): def test_no_world_width_is_a_quotient_of_the_leaves(self):
"""It is not a quotient, and the live getter is right at every moment. """Deriving one would answer with the joining cohort's own width on a
A value fixed when the groups are built would answer with the launch scale joiner, which lays its groups out at `tp * pp` while WORLD spans
count after `try_admit_scale_ranks` expands WORLD, and with the joining `ep_join_rank_offset + tp * pp`. The launch width comes off the group
cohort's own width on a scale-joiner, which lays its groups out at that was actually built; the ceiling is not this function's to give
`tp * pp` while WORLD spans `ep_join_rank_offset + tp * pp`.""" either, and `TestTheTwoWorldWidths` says where each comes from."""
widths = derive_parallel_widths( widths = derive_parallel_widths(
tp_size=4, tp_size=4,
attn_cp_size=1, attn_cp_size=1,
@@ -1766,11 +1815,27 @@ class TestDerivedWidths(_IsolatedOverrides):
dcp_size=1, dcp_size=1,
dcp_enabled=False, dcp_enabled=False,
) )
self.assertNotIn("world_size", widths) self.assertEqual(
{name for name in widths if "world" in name},
set(),
)
parallel = get_parallel() parallel = get_parallel()
parallel.override_permanently(attn_tp_size=4) parallel.override_permanently(attn_tp_size=4)
with patch(f"{_PS}.get_world_size", return_value=9): with patch(f"{_PS}.get_world_size", return_value=9):
self.assertEqual(parallel.world_size, 9) self.assertEqual(parallel.launch_world_size, 9)
def test_the_bare_name_is_gone(self):
"""It answered two questions, so every reader had to remember which.
Both spellings fail: reading it, and stating it -- the overridable set
is derived from the same declarations the read path is, so a name that
cannot be read cannot be stated either.
"""
with self.assertRaisesRegex(AttributeError, r"has no 'world_size'"):
get_parallel().world_size
with self.assertRaisesRegex(ValueError, r"unknown parallel field"):
with get_parallel().override(world_size=4):
pass
def test_a_permanently_overridden_width_is_what_the_reader_answers_with(self): def test_a_permanently_overridden_width_is_what_the_reader_answers_with(self):
parallel = get_parallel() parallel = get_parallel()