Name the two widths of the WORLD group (#40070)
This commit is contained in:
@@ -61,16 +61,16 @@ _SRT = _pathlib.Path(next(iter(_sglang.__path__))).resolve() / "srt"
|
||||
_PS = "sglang.srt.distributed.parallel_state"
|
||||
_DP = "sglang.srt.layers.dp_attention"
|
||||
|
||||
# Ranks and the world size read the live group: they are not implied by
|
||||
# anything, so there is nothing to derive them from. The quotients used to be
|
||||
# in this table and are not any more -- `attn_tp_size` and its siblings are
|
||||
# functions of the configured leaves, and `TestDerivedWidthsComeFromTheLeaves`
|
||||
# is what pins them. `attn_dp_rank` is not here either: no group coordinator
|
||||
# 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 are not
|
||||
# here -- `attn_tp_size` and its siblings are functions of the configured
|
||||
# leaves, and `TestDerivedWidths` 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
|
||||
# `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 = [
|
||||
("world_size", f"{_PS}.get_world_size"),
|
||||
("world_rank", f"{_PS}.get_world_rank"),
|
||||
("launch_world_size", f"{_PS}.get_world_size"),
|
||||
("launch_world_rank", f"{_PS}.get_world_rank"),
|
||||
("tp_rank", f"{_PS}.get_tensor_model_parallel_rank"),
|
||||
("dcp_rank", f"{_PS}.get_dcp_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"))
|
||||
|
||||
|
||||
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):
|
||||
"""`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["attn_dcp_size"], 1)
|
||||
|
||||
def test_the_world_size_is_not_permanently_overridden(self):
|
||||
"""It is not a quotient, and the live getter is right at every moment.
|
||||
A value fixed when the groups are built would answer with the launch
|
||||
count after `try_admit_scale_ranks` expands WORLD, and with the joining
|
||||
cohort's own width on a scale-joiner, which lays its groups out at
|
||||
`tp * pp` while WORLD spans `ep_join_rank_offset + tp * pp`."""
|
||||
def test_no_world_width_is_a_quotient_of_the_leaves(self):
|
||||
"""Deriving one would answer with the joining cohort's own width on a
|
||||
scale joiner, which lays its groups out at `tp * pp` while WORLD spans
|
||||
`ep_join_rank_offset + tp * pp`. The launch width comes off the group
|
||||
that was actually built; the ceiling is not this function's to give
|
||||
either, and `TestTheTwoWorldWidths` says where each comes from."""
|
||||
widths = derive_parallel_widths(
|
||||
tp_size=4,
|
||||
attn_cp_size=1,
|
||||
@@ -1766,11 +1815,27 @@ class TestDerivedWidths(_IsolatedOverrides):
|
||||
dcp_size=1,
|
||||
dcp_enabled=False,
|
||||
)
|
||||
self.assertNotIn("world_size", widths)
|
||||
self.assertEqual(
|
||||
{name for name in widths if "world" in name},
|
||||
set(),
|
||||
)
|
||||
parallel = get_parallel()
|
||||
parallel.override_permanently(attn_tp_size=4)
|
||||
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):
|
||||
parallel = get_parallel()
|
||||
|
||||
Reference in New Issue
Block a user