Check the topology identities where the layout is written, and build at the published widths (#40340)

This commit is contained in:
Cheng Wan
2026-09-21 12:22:59 -07:00
committed by GitHub
parent d5fdab7022
commit 2d0e94e3a3
43 changed files with 843 additions and 261 deletions
+392 -65
View File
@@ -41,6 +41,7 @@ from sglang.srt.runtime_context import (
RuntimeContext,
SpawnRanks,
_FlagGroupBase,
_validate_parallel,
assert_published,
derive_parallel_widths,
get_context,
@@ -59,6 +60,56 @@ from sglang.srt.server_args import ServerArgs
from sglang.test.test_utils import CustomTestCase
_SRT = _pathlib.Path(next(iter(_sglang.__path__))).resolve() / "srt"
_PACKAGE = _pathlib.Path(next(iter(_sglang.__path__))).resolve()
def _sources():
"""Every Python file this checkout ships, package and siblings alike.
The package alone is the wrong subject set for anything about entries or
public names: `benchmark/`, `examples/` and the top-level `test/` call the
same doors and are not covered by any suite that would notice them break.
An installed package has no siblings, and then this is the package alone."""
roots = [_PACKAGE]
checkout = _PACKAGE.parents[1]
roots += [
checkout / name
for name in ("benchmark", "examples", "scripts", "test")
if (checkout / name).is_dir()
]
for root in roots:
for path in root.rglob("*.py"):
yield path
def _scope_entries_that_say_nothing(paths):
"""Draft-scope entries that do not state `owns_attention`, as `path:line`.
The scope either narrows the draft's attention and expert identity or
leaves the target's in place, and only the worker knows which -- so the
keyword has no default. Omitting it is a `TypeError`, but only on the path
that runs, and those paths want a GPU and a draft model.
"""
import ast
missing = []
for path in paths:
try:
tree = ast.parse(path.read_text(encoding="utf-8"))
except (SyntaxError, UnicodeDecodeError):
continue
for node in ast.walk(tree):
if not isinstance(node, ast.Call):
continue
func = node.func
name = getattr(func, "attr", None) or getattr(func, "id", None)
if name not in ("draft_tp_context", "patch_tensor_parallel_group"):
continue
if not any(kw.arg == "owns_attention" for kw in node.keywords):
missing.append(f"{path}:{node.lineno}")
return missing
_PS = "sglang.srt.distributed.parallel_state"
_DP = "sglang.srt.layers.dp_attention"
@@ -347,9 +398,8 @@ class TestStampedRanks(_IsolatedOverrides):
"""`attn_dp_rank` comes from the stamp, and says so when there is none.
It is the one rank no group answers with: `initialize_dp_attention`
computes it from this process's `tp_rank`, and an elastic scale-up
replaces it with a rank in the expanded WORLD. Falling back to anything
would be inventing a placement for this process.
computes it from this process's `tp_rank`. Falling back to anything would
be inventing a placement for this process.
"""
def setUp(self):
@@ -407,21 +457,71 @@ class TestStampedRanks(_IsolatedOverrides):
)
self.assertIs(mode, DpPaddingMode.SUM_LEN)
def test_a_scale_up_stamps_the_width_and_the_rank_together(self):
"""The two describe one topology; a reader that saw only one moved
would place this process in a group it is not in."""
def test_the_gather_slot_follows_the_list_that_was_gathered(self):
"""The DP sync gathers over the attention-DP replicas, or over the
expanded WORLD once a scale-up has moved the gather there. The index
into that list is a property of the gather, so it is read beside the
flag that says which one happened rather than kept on the topology."""
from sglang.srt.layers.dp_attention import dp_gather_slot
self.addCleanup(reset_context)
dp_flags = get_flags().dp
saved = (
dp_flags.use_world_group_for_gather,
dp_flags.joiner_skip_all_gather,
)
def restore():
(
dp_flags.use_world_group_for_gather,
dp_flags.joiner_skip_all_gather,
) = saved
self.addCleanup(restore)
publish(
ServerArgs(
model_path="dummy", tp_size=8, dp_size=8, enable_dp_attention=True
),
role="test",
ranks=SpawnRanks(world_rank=3),
)
parallel = get_parallel()
dp_flags.use_world_group_for_gather = False
self.assertEqual(dp_gather_slot(), parallel.attn_dp_rank)
# After a scale-up the gather spans the expanded WORLD, and the joining
# cohort is numbered from its offset.
dp_flags.use_world_group_for_gather = True
dp_flags.joiner_skip_all_gather = False
parallel.override_permanently(ep_join_rank_offset=8)
self.assertEqual(dp_gather_slot(), 8 + parallel.tp_rank)
# and the topology it was read off is untouched
self.assertEqual(parallel.attn_dp_size, 8)
self.assertEqual(parallel.tp_size, 8)
def test_a_scale_up_writes_no_width(self):
"""The identities stay unconditional because nothing overrides them:
the scale-up only points the gather at the expanded WORLD."""
from sglang.srt.layers.dp_attention import update_dp_attention_post_scale
# It also flips a process-wide gather flag; put it back, or every
# later test in this process runs as if a scale-up had happened.
dp_flags = get_flags().dp
saved_gather = dp_flags.use_world_group_for_gather
self.addCleanup(setattr, dp_flags, "use_world_group_for_gather", saved_gather)
self.addCleanup(reset_context)
publish(
ServerArgs(
model_path="dummy", tp_size=8, dp_size=8, enable_dp_attention=True
),
role="test",
ranks=SpawnRanks(world_rank=3),
)
parallel = get_parallel()
before = (parallel.attn_dp_size, parallel.attn_dp_rank, parallel.tp_size)
update_dp_attention_post_scale(new_dp_size=16, new_dp_rank=11)
self.assertEqual(parallel.attn_dp_size, 16)
self.assertEqual(parallel.attn_dp_rank, 11)
self.assertTrue(dp_flags.use_world_group_for_gather)
self.assertEqual(
(parallel.attn_dp_size, parallel.attn_dp_rank, parallel.tp_size), before
)
class TestEveryDeclaredParallelNameIsStatable(_IsolatedOverrides):
@@ -1898,14 +1998,26 @@ class TestDerivedWidths(_IsolatedOverrides):
def test_a_topology_is_stated_by_naming_the_width(self):
"""Overriding a leaf does not move the quotient -- the quotient is not
recomputed on read. Naming it is how a test states one."""
recomputed on read. Naming it is how a caller states one, and naming
only some of them is refused: the caller owns the arithmetic, the
context only checks it."""
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy", tp_size=8), role="test")
self.assertEqual(get_parallel().attn_tp_size, 8)
with get_parallel().override(tp_size=2):
self.assertEqual(get_parallel().attn_tp_size, 8)
with get_parallel().override(attn_tp_size=4):
with self.assertRaises(ValueError) as caught:
with get_parallel().override(tp_size=2):
pass
self.assertIn(
"tp_size == attn_tp_size * attn_dp_size * attn_cp_size",
str(caught.exception),
)
with get_parallel().override(tp_size=2, attn_tp_size=2, moe_tp_size=2):
self.assertEqual(get_parallel().tp_size, 2)
self.assertEqual(get_parallel().attn_tp_size, 2)
with get_parallel().override(attn_tp_size=4, tp_size=4, moe_tp_size=4):
self.assertEqual(get_parallel().attn_tp_size, 4)
def test_an_unstated_topology_still_fails(self):
@@ -2127,17 +2239,11 @@ class TestDerivedWidths(_IsolatedOverrides):
)
self.assertEqual(published, recomputed)
def test_initialize_model_parallel_no_longer_touches_the_bag(self):
"""`initialize_model_parallel` used to recompute and
permanently override the six derived widths on `get_parallel()`
after building its groups; that call is gone. Publish a placeholder
config (tp_size defaults to 1), then build real groups at a
different width -- the published leaf must now stay exactly what it
was, because nothing corrects it. This is the behavior a caller
relies on being told about, loudly, the first time it publishes and
builds inconsistently -- see
`test_recomputing_from_published_leaves_matches_the_publish_bag`
for why every real caller must not do that.
def test_initialize_model_parallel_builds_at_the_published_widths(self):
"""The build takes every width from the context rather than from an
argument, so "published one width, built another" is no longer a state
a caller can reach -- there is nothing left to translate, and nothing
to correct afterwards either.
"""
from unittest.mock import Mock
@@ -2145,11 +2251,11 @@ class TestDerivedWidths(_IsolatedOverrides):
reset_context()
self.addCleanup(reset_context)
publish(ServerArgs(model_path="dummy"), role="test")
self.assertEqual(get_parallel().attn_tp_size, 1)
self.assertEqual(get_parallel().moe_ep_size, 1)
world_size = 8
publish(ServerArgs(model_path="dummy", tp_size=world_size), role="test")
self.assertEqual(get_parallel().attn_tp_size, world_size)
built_at = []
with (
patch.object(parallel_state, "_WORLD", None),
patch.object(parallel_state, "_TP", None),
@@ -2168,25 +2274,20 @@ class TestDerivedWidths(_IsolatedOverrides):
patch.object(
parallel_state,
"init_model_parallel_group",
return_value=Mock(device_group=Mock()),
side_effect=lambda group_ranks, *a, **k: (
built_at.append(group_ranks),
Mock(device_group=Mock()),
)[1],
),
patch.object(parallel_state, "get_world_group") as mock_world_group,
):
mock_world_group.return_value = Mock(device_group=Mock(), local_rank=0)
parallel_state.initialize_model_parallel(
tensor_model_parallel_size=world_size,
expert_model_parallel_size=world_size,
)
parallel_state.initialize_model_parallel()
self.addCleanup(parallel_state.destroy_model_parallel)
self.assertEqual(
get_parallel().attn_tp_size,
1,
"initialize_model_parallel must not touch the published leaf -- "
"a caller that needs it corrected must publish a config that "
"already matches the width it is about to build",
)
self.assertEqual(get_parallel().moe_ep_size, 1)
# The first group built is TP, one group spanning the published width.
self.assertEqual(built_at[0], [list(range(world_size))])
self.assertEqual(get_parallel().attn_tp_size, world_size)
class TestTheDerivedHalfIsDeclared(CustomTestCase):
@@ -2294,6 +2395,225 @@ class TestAnEntryThatBuildsARunnerHandsOverItsPlacement(CustomTestCase):
)
class TestTheAccessorsHaveNoCallersOutsideTheirPackage(CustomTestCase):
"""`parallel_state`'s getters are the definition, not a second spelling.
Business code asks `get_parallel()`; a call that goes straight to the getter
is a read the context cannot redirect, which is what a scope needs it to be
able to do. The package that defines them is exempt -- a read there would
go through the context back into itself -- and so is `multimodal_gen`, which
has its own parallel state.
"""
#: Not topology. `get_self_pp_group` builds the single-rank group a draft
#: pipeline scope installs, so there is nothing for the context to answer
#: with until the scope has installed it.
ALLOWED = {
"get_self_pp_group",
"get_default_distributed_backend",
"get_mooncake_transfer_engine",
}
def _accessors(self):
"""Derived from the source, not listed here: a guard whose subject set
is written by hand stops watching whatever gets added next."""
from sglang.srt.distributed import parallel_state as parallel_state_module
source = _pathlib.Path(parallel_state_module.__file__).read_text().splitlines()
return {
line[len("def ") : line.index("(")]
for line in source
if line.startswith("def get_") or line.startswith("def is_")
}
def _callers(self, name):
import re
from sglang.srt.distributed import parallel_state as parallel_state_module
root = _pathlib.Path(parallel_state_module.__file__).parents[2]
pattern = re.compile(rf"(?<![.\w]){re.escape(name)}\(")
hits = []
for path in root.rglob("*.py"):
rel = path.relative_to(root).as_posix()
if rel.startswith(("srt/distributed/", "multimodal_gen/", "test/")):
continue
for number, line in enumerate(path.read_text().splitlines(), 1):
if line.lstrip().startswith(("def ", "#")):
continue
if pattern.search(line):
hits.append(f"{rel}:{number}")
return hits
def test_no_business_code_calls_them(self):
offenders = {}
for name in sorted(self._accessors() - self.ALLOWED):
callers = self._callers(name)
if callers:
offenders[name] = callers
self.assertEqual(
offenders,
{},
"read these through get_parallel() instead, or say here why the "
"context cannot answer them",
)
def test_the_guard_would_notice_a_caller(self):
"""The subject set is derived, so this checks the search finds a real
call rather than that the list happens to be empty: `get_self_pp_group`
is exempt and does have one caller."""
self.assertTrue(self._callers("get_self_pp_group"))
class TestTheTopologyIdentities(CustomTestCase):
"""One set of identities, checked wherever the layout is written.
Each one is injected in the direction that breaks it and in the direction
that keeps it: a guard that only ever fires is as uninformative as one that
never does. They hold unconditionally -- a caller that states one leaf owes
the quotients that follow from it, because a namespace describing no real
layout is what the guard exists to refuse.
"""
def _publish_square(self):
"""tp=4 over two attention-DP replicas of two: every identity holds."""
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(
model_path="dummy", tp_size=4, dp_size=2, enable_dp_attention=True
),
role="scheduler",
ranks=SpawnRanks(world_rank=3, dp_rank=1),
)
def test_a_published_topology_is_consistent(self):
"""The quiet direction, and the reason publish can check everything:
it is the one write that establishes the whole layout."""
self._publish_square()
parallel = get_parallel()
self.assertEqual(parallel.tp_size, 4)
self.assertEqual(parallel.attn_dp_size, 2)
self.assertEqual(parallel.attn_tp_size, 2)
self.assertEqual(parallel.tp_rank, 3)
def test_a_width_that_does_not_factor_is_refused(self):
self._publish_square()
with self.assertRaises(ValueError) as caught:
with get_parallel().override(
tp_size=4, attn_tp_size=3, attn_dp_size=1, attn_cp_size=1
):
pass
message = str(caught.exception)
self.assertIn("set by override", message)
self.assertIn("attn_tp_size * attn_dp_size * attn_cp_size", message)
self.assertIn("4 != 3 * 1 * 1", message)
def test_a_rank_at_its_width_is_refused(self):
self._publish_square()
with self.assertRaises(ValueError) as caught:
with get_parallel().override(tp_rank=4, tp_size=4):
pass
self.assertIn("0 <= tp_rank < tp_size", str(caught.exception))
def test_a_moe_width_that_does_not_factor_is_refused(self):
self._publish_square()
with self.assertRaises(ValueError) as caught:
with get_parallel().override(
tp_size=4, moe_ep_size=1, moe_dp_size=1, moe_tp_size=3
):
pass
message = str(caught.exception)
self.assertIn("moe_ep_size * moe_dp_size * moe_tp_size", message)
self.assertIn("4 != 1 * 1 * 3", message)
def test_a_rank_the_attention_layout_cannot_produce_is_refused(self):
"""`tp_rank` is not free of the attention ranks: the layout derives one
from the other, so a set that does not satisfy it places this process
in two different seats at once."""
self._publish_square()
with self.assertRaises(ValueError) as caught:
with get_parallel().override(
tp_rank=0,
attn_dp_rank=1,
attn_cp_rank=0,
attn_tp_rank=0,
attn_cp_size=1,
attn_tp_size=2,
):
pass
self.assertIn("attn_tp_size + attn_tp_rank", str(caught.exception))
def test_the_published_ranks_satisfy_the_layout(self):
"""The quiet direction for the same identity: publish derives the
attention ranks from `tp_rank` through that very equation, so a
published process always sits in one seat."""
self._publish_square()
parallel = get_parallel()
self.assertEqual(
parallel.tp_rank,
(parallel.attn_dp_rank * parallel.attn_cp_size + parallel.attn_cp_rank)
* parallel.attn_tp_size
+ parallel.attn_tp_rank,
)
def test_a_refused_write_leaves_nothing_behind(self):
"""The scope never opened, so the value it tried to state must not be
readable afterwards -- a half-applied override is the state this guard
exists to prevent."""
self._publish_square()
with self.assertRaises(ValueError):
with get_parallel().override(
tp_size=4, attn_tp_size=3, attn_dp_size=1, attn_cp_size=1
):
pass
self.assertEqual(get_parallel().attn_tp_size, 2)
def test_a_group_built_at_another_width_is_refused(self):
"""The other end of the same identity: what the configuration says and
what the coordinators were actually built at, checked where the
disagreement is still attributable to the build."""
from sglang.srt.distributed import parallel_state
from sglang.srt.distributed.parallel_state import GroupCoordinator
self._publish_square()
wrong = GroupCoordinator.__new__(GroupCoordinator)
wrong.world_size = 8
wrong.rank_in_group = 0
with patch.object(parallel_state, "_TP", wrong):
with self.assertRaises(ValueError) as caught:
_validate_parallel(get_parallel(), "group build")
message = str(caught.exception)
self.assertIn("set by group build", message)
self.assertIn("tp_group.world_size == tp_size", message)
self.assertIn("built 8, configured 4", message)
def test_a_group_built_at_the_configured_width_is_quiet(self):
from sglang.srt.distributed import parallel_state
from sglang.srt.distributed.parallel_state import GroupCoordinator
self._publish_square()
right = GroupCoordinator.__new__(GroupCoordinator)
right.world_size = 4
right.rank_in_group = 3
with patch.object(parallel_state, "_TP", right):
_validate_parallel(get_parallel(), "group build")
def test_a_draft_scope_states_a_consistent_topology(self):
"""The scope narrows four names at once, so the identity applies to it
-- and holds, which is what step lets the guard stay on."""
from sglang.srt.distributed import parallel_state
from sglang.srt.distributed.parallel_state import GroupCoordinator
self._publish_square()
group = GroupCoordinator.__new__(GroupCoordinator)
group.world_size = 2
group.rank_in_group = 1
with patch.object(parallel_state, "_TP", group):
with parallel_state.patch_tensor_parallel_group(group, owns_attention=True):
self.assertEqual(get_parallel().attn_tp_size, 2)
class TestWhoAnswersDuringADraftScope(CustomTestCase):
"""A draft worker runs in one process with the target, under a scope.
@@ -2397,29 +2717,36 @@ class TestWhoAnswersDuringADraftScope(CustomTestCase):
worker states it, and a caller that forgets is the bug this catches --
`owns_attention` has no default, but a missing one is a TypeError only
on the path that runs, and these paths need a GPU and a draft model."""
import ast
self.assertEqual(
_scope_entries_that_say_nothing(_sources()),
[],
"these enter the draft scope without saying",
)
package = _pathlib.Path(next(iter(_sglang.__path__))).resolve()
checkout = package.parents[1]
roots = [package] + [
checkout / name for name in ("test",) if (checkout / name).is_dir()
]
missing = []
for path in (q for root in roots for q in root.rglob("*.py")):
try:
tree = ast.parse(path.read_text(encoding="utf-8"))
except (SyntaxError, UnicodeDecodeError):
continue
for node in ast.walk(tree):
if not isinstance(node, ast.Call):
continue
func = node.func
name = getattr(func, "attr", None) or getattr(func, "id", None)
if name not in ("draft_tp_context", "patch_tensor_parallel_group"):
continue
if not any(kw.arg == "owns_attention" for kw in node.keywords):
missing.append(f"{path}:{node.lineno}")
self.assertEqual(missing, [], "these enter the draft scope without saying")
def test_the_census_would_notice_one(self):
"""Two ways for it to report zero and still be wrong: the matcher does
not recognise the call, or the walk never reaches the file. A draft
scope entered from `benchmark/` breaks the same way as one in the
package and no suite covers it, so the roots are part of the check."""
trees = {
part
for path in _sources()
for part in ("benchmark", "examples", "scripts", "test")
if f"/{part}/" in path.as_posix()
}
self.assertEqual(
trees,
{"benchmark", "examples", "scripts", "test"},
"the walk misses a tree that can enter the scope",
)
with tempfile.TemporaryDirectory() as tmp:
probe = _pathlib.Path(tmp) / "probe.py"
probe.write_text(
"with self.draft_tp_context(runner.tp_group):\n pass\n",
encoding="utf-8",
)
self.assertEqual(_scope_entries_that_say_nothing([probe]), [f"{probe}:1"])
def test_a_full_width_swap_leaves_the_attention_layout_alone(self):
"""The other caller. A draft built outside any scope carries the