One read path for every parallel name (#40069)

This commit is contained in:
Cheng Wan
2026-09-18 17:43:56 -07:00
committed by GitHub
parent afe71f4b9e
commit 81421b91e9
5 changed files with 336 additions and 262 deletions
@@ -150,7 +150,7 @@ def _clear_srt_tp_group() -> None:
if srt_parallel_state._ATTN_TP is _TP: if srt_parallel_state._ATTN_TP is _TP:
srt_parallel_state._ATTN_TP = None srt_parallel_state._ATTN_TP = None
get_parallel().clear_derived_widths() get_parallel().clear_stamp()
if srt_parallel_state._TP is _TP: if srt_parallel_state._TP is _TP:
srt_parallel_state._TP = None srt_parallel_state._TP = None
@@ -3163,7 +3163,7 @@ def get_moe_tensor_parallel_rank():
def destroy_model_parallel(): def destroy_model_parallel():
"""Set the groups to none and destroy them.""" """Set the groups to none and destroy them."""
get_parallel().clear_derived_widths() get_parallel().clear_stamp()
dwdp_mgr = get_global_dwdp_manager() dwdp_mgr = get_global_dwdp_manager()
if dwdp_mgr is not None: if dwdp_mgr is not None:
dwdp_mgr.cleanup() dwdp_mgr.cleanup()
+151 -219
View File
@@ -68,10 +68,23 @@ logger = logging.getLogger(__name__)
# Imported lazily so this module has no import-time dependencies: any module can # Imported lazily so this module has no import-time dependencies: any module can
# import get_parallel at module level without risking an import cycle. # import get_parallel at module level without risking an import cycle.
def _ps(): _PARALLEL_STATE = None
from sglang.srt.distributed import parallel_state
return parallel_state
def _ps():
"""The module every rank and group read ends at.
Cached because the import statement dominated the read: a group read is
two attribute lookups plus this, and it runs per row-linear on an eager
forward. The getter is still resolved by name on the returned module, so
a test that patches `parallel_state.get_tp_group` is still seen.
"""
global _PARALLEL_STATE
if _PARALLEL_STATE is None:
from sglang.srt.distributed import parallel_state
_PARALLEL_STATE = parallel_state
return _PARALLEL_STATE
def _dp(): def _dp():
@@ -97,42 +110,60 @@ def _parallel_config_leaves() -> frozenset:
) )
_PARALLEL_FIELDS = frozenset( # Ranks and group handles: the names no configuration carries, each with the
{ # canonical getter that answers it live. This table is their declaration, the
"world_size", # way `arg_groups/fields/parallel.py` is the leaves' and `Derived` is the
"world_rank", # widths'. `None` marks a name only a stamp can answer: no coordinator knows
"tp_size", # this process's attention-DP rank.
"tp_rank", _MISSING_READ = object()
"pp_size",
"pp_rank", _LIVE_READS: dict = {
"moe_ep_size", "world_size": "get_world_size",
"moe_ep_rank", "world_rank": "get_world_rank",
"moe_dp_size", "tp_rank": "get_tensor_model_parallel_rank",
"moe_dp_rank", "pp_rank": "get_pipeline_model_parallel_rank",
"moe_tp_size", "moe_ep_rank": "get_moe_expert_parallel_rank",
"moe_tp_rank", "moe_dp_rank": "get_moe_data_parallel_rank",
"attn_tp_size", "moe_tp_rank": "get_moe_tensor_parallel_rank",
"attn_tp_rank", "attn_tp_rank": "get_attn_tensor_model_parallel_rank",
"attn_cp_size", "attn_cp_rank": "get_attn_context_model_parallel_rank",
"attn_cp_rank", "dcp_rank": "get_dcp_rank",
"dcp_enabled", "attn_dcp_rank": lambda self: self.dcp_rank if self.dcp_enabled else 0,
"dcp_size", "attn_dp_rank": None,
"dcp_rank", "world_group": "get_world_group",
"attn_dcp_size", "tp_group": "get_tp_group",
"attn_dcp_rank", "pp_group": "get_pp_group",
"attn_dp_size", "moe_ep_group": "get_moe_ep_group",
"attn_dp_rank", "moe_dp_group": "get_moe_dp_group",
"world_group", "moe_tp_group": "get_moe_tp_group",
"tp_group", "attn_tp_group": "get_attn_tp_group",
"pp_group", "attn_cp_group": "get_attn_cp_group",
"moe_ep_group", "dcp_group": "get_dcp_group",
"moe_dp_group", }
"moe_tp_group",
"attn_tp_group",
"attn_cp_group", @functools.lru_cache(maxsize=1)
"dcp_group", def _parallel_fields() -> frozenset:
"""Every name `ParallelContext` answers for, read from the declarations.
Three sources, because the namespace has three kinds of name and each one
declares itself somewhere already:
* configured leaves -- the `parallel` namespace of the record;
* derived widths -- the `Derived` declarations beside those leaves;
* ranks and group handles -- `_LIVE_READS`, which is where they are
declared because no configuration carries them.
The set is the union of those three, so `override()` cannot refuse a name
the class answers for.
"""
from sglang.srt.arg_groups.arg_utils import Derived
from sglang.srt.arg_groups.fields.parallel import Parallel
derived = {
name for name, decl in vars(Parallel).items() if isinstance(decl, Derived)
} }
) return frozenset(_parallel_config_leaves() | derived | set(_LIVE_READS))
def derive_attention_widths( def derive_attention_widths(
@@ -265,12 +296,12 @@ class ParallelContext:
different names rather than two answers to one name. different names rather than two answers to one name.
""" """
__slots__ = ("_overrides", "_config", "_derived") __slots__ = ("_overrides", "_stamp", "_config")
def __init__(self): def __init__(self):
self._overrides = {} self._overrides = {} # scoped, restored when the `with` block exits
self._stamp = {} # permanent for the process, dropped by clear_stamp
self._config = None # parallel config bag, wired at publish self._config = None # parallel config bag, wired at publish
self._derived = {} # widths overridden permanently, as the groups are built
def __getattr__(self, name): def __getattr__(self, name):
if name.startswith("_"): if name.startswith("_"):
@@ -278,48 +309,54 @@ class ParallelContext:
# still unset (pickle/copy protocols probe attributes before # still unset (pickle/copy protocols probe attributes before
# __init__ runs). # __init__ runs).
raise AttributeError(name) raise AttributeError(name)
return self._read(name)
def _read(self, name):
"""The one read path, for every kind of name in the namespace.
Scoped override, then the permanent stamp, then what the name is
answered by when nobody has stated it: the published leaf for a
configured value or a derived width, the canonical getter for a rank
or a group handle.
The two override maps stay separate because they are taken down by
different things -- a `with` block and `clear_stamp()` -- and
merging them would let a teardown of one drop the other, and would
turn "which wins" into whichever was written last.
"""
overrides = self._overrides overrides = self._overrides
if name in overrides: if name in overrides:
return overrides[name] return overrides[name]
stamp = self._stamp
if name in stamp:
return stamp[name]
config = self._config config = self._config
if config is not None: if config is not None and name in config._fields:
if name in config._fields: return getattr(config, name)
return getattr(config, name) live = _LIVE_READS.get(name, _MISSING_READ)
elif name in _parallel_config_leaves(): if live is not _MISSING_READ:
if isinstance(live, str):
return getattr(_ps(), live)()
if live is not None:
return live(self)
raise RuntimeError(
f"parallel rank {name!r} is not available: it is computed from "
"this process's `tp_rank` when the attention topology is "
"initialized, so a process that never ran "
"`initialize_dp_attention` has no answer to give"
)
if config is None and name in _parallel_config_leaves():
raise ValueError("config namespace 'parallel' not published") raise ValueError("config namespace 'parallel' not published")
if name in _derived_widths():
raise RuntimeError(
f"derived parallel width {name!r} is not available: it is computed "
"from the configured leaves at publish, and permanently corrected "
"when the process groups are built. Nothing is published and "
"nothing has been set with override_permanently -- publish a "
f"parallel config, or state the width with get_parallel().override({name}=...)"
)
raise AttributeError(f"ParallelContext has no {name!r}") raise AttributeError(f"ParallelContext has no {name!r}")
def _v(self, name, getter):
"""Scoped override, else the permanent stamp, else the live group.
One priority order for ranks and widths alike (`_derived_width`), so a
stamped value wins over the coordinator for both.
"""
overrides = self._overrides
if name in overrides:
return overrides[name]
derived = self._derived
if name in derived:
return derived[name]
return getter()
def _stamped(self, name, why):
"""A per-process fact no configuration implies: scoped override, else
the permanent stamp, else fail.
Unlike a width, this has nothing to fall back on -- the configuration
does not carry this process's rank, and there is no group to ask --
so an unstamped read is a missing initialization rather than a
missing override, and says so.
"""
overrides = self._overrides
if name in overrides:
return overrides[name]
derived = self._derived
if name in derived:
return derived[name]
raise RuntimeError(f"parallel rank {name!r} is not available: {why}")
def override_permanently(self, **values) -> None: def override_permanently(self, **values) -> None:
"""Permanently record a width or rank the published bag can't answer """Permanently record a width or rank the published bag can't answer
or no longer answers correctly -- not `RuntimeContext.override`, or no longer answers correctly -- not `RuntimeContext.override`,
@@ -331,50 +368,25 @@ class ParallelContext:
answer and this only corrects it; a rank is a per-process fact the answer and this only corrects it; a rank is a per-process fact the
configuration never carries, so for those this is the only source. configuration never carries, so for those this is the only source.
Lives beside, not inside, the `@contextmanager` `override` above -- a Lives beside, not inside, the `@contextmanager` `override` below -- a
name it cannot also have on this class -- because these are permanent name it cannot also have on this class -- because these are permanent
for the process, not scoped to a `with` block: none of the real for the process, not scoped to a `with` block: none of the real
callers ever restore the value they set here. callers ever restore the value they set here.
""" """
self._derived.update(values) unknown = set(values) - _parallel_fields()
if unknown:
raise ValueError(f"unknown parallel field(s): {sorted(unknown)}")
self._stamp.update(values)
def clear_derived_widths(self) -> None: def clear_stamp(self) -> None:
self._derived.clear() """Drop every stamped name, ranks included."""
self._stamp.clear()
def _derived_width(self, name):
"""A width the configuration implies: scoped override, else permanent
override, else the published leaf.
The leaf is computed at publish by `parallel_widths_of`; the permanent
override sits above it because an elastic scale-up corrects
`attn_dp_size` after publish, and a scope that swaps in another TP
group states the quotients through the scoped `override` above that.
Nothing is recomputed on read, so overriding `tp_size` does not move
`attn_tp_size`: name the width, or publish a config.
"""
overrides = self._overrides
if name in overrides:
return overrides[name]
derived = self._derived
if name in derived:
return derived[name]
config = self._config
if config is not None and name in config._fields:
return getattr(config, name)
raise RuntimeError(
f"derived parallel width {name!r} is not available: it is computed "
"from the configured leaves at publish, and permanently corrected "
"when the process groups are built. Nothing is published and "
"nothing has been set with override_permanently -- publish a "
f"parallel config, or state the width with get_parallel().override({name}=...)"
)
@contextmanager @contextmanager
def override(self, **kwargs): def override(self, **kwargs):
"""Temporarily force parallel values, restoring on exit. Validates keys and """Temporarily force parallel values, restoring on exit. Validates keys and
supports nesting.""" supports nesting."""
unknown = set(kwargs) - _PARALLEL_FIELDS unknown = set(kwargs) - _parallel_fields()
if unknown: if unknown:
raise ValueError(f"unknown parallel field(s): {sorted(unknown)}") raise ValueError(f"unknown parallel field(s): {sorted(unknown)}")
saved = dict(self._overrides) saved = dict(self._overrides)
@@ -384,124 +396,44 @@ class ParallelContext:
finally: finally:
self._overrides = saved self._overrides = saved
@property
def world_size(self) -> int:
return self._v("world_size", _ps().get_world_size)
@property def _derived_widths() -> dict:
def world_rank(self) -> int: """The declared quotients, by name -- `{name: Derived}`."""
return self._v("world_rank", _ps().get_world_rank)
@property
def tp_rank(self) -> int:
return self._v("tp_rank", _ps().get_tensor_model_parallel_rank)
@property
def pp_rank(self) -> int:
return self._v("pp_rank", _ps().get_pipeline_model_parallel_rank)
@property
def moe_ep_rank(self) -> int:
return self._v("moe_ep_rank", _ps().get_moe_expert_parallel_rank)
@property
def moe_dp_rank(self) -> int:
return self._v("moe_dp_rank", _ps().get_moe_data_parallel_rank)
@property
def moe_tp_rank(self) -> int:
return self._v("moe_tp_rank", _ps().get_moe_tensor_parallel_rank)
@property
def attn_tp_rank(self) -> int:
return self._v("attn_tp_rank", _ps().get_attn_tensor_model_parallel_rank)
@property
def attn_cp_rank(self) -> int:
return self._v("attn_cp_rank", _ps().get_attn_context_model_parallel_rank)
@property
def dcp_rank(self) -> int:
return self._v("dcp_rank", _ps().get_dcp_rank)
@property
def attn_dcp_rank(self) -> int:
return self._v(
"attn_dcp_rank", lambda: self.dcp_rank if self.dcp_enabled else 0
)
@property
def attn_dp_rank(self) -> int:
return self._stamped(
"attn_dp_rank",
"it is computed from this process's `tp_rank` when the attention "
"topology is initialized, so a process that never ran "
"`initialize_dp_attention` has no answer to give",
)
@property
def world_group(self) -> Any:
return self._v("world_group", _ps().get_world_group)
@property
def tp_group(self) -> Any:
return self._v("tp_group", _ps().get_tp_group)
@property
def pp_group(self) -> Any:
return self._v("pp_group", _ps().get_pp_group)
@property
def moe_ep_group(self) -> Any:
return self._v("moe_ep_group", _ps().get_moe_ep_group)
@property
def moe_dp_group(self) -> Any:
return self._v("moe_dp_group", _ps().get_moe_dp_group)
@property
def moe_tp_group(self) -> Any:
return self._v("moe_tp_group", _ps().get_moe_tp_group)
@property
def attn_tp_group(self) -> Any:
return self._v("attn_tp_group", _ps().get_attn_tp_group)
@property
def attn_cp_group(self) -> Any:
return self._v("attn_cp_group", _ps().get_attn_cp_group)
@property
def dcp_group(self) -> Any:
return self._v("dcp_group", _ps().get_dcp_group)
def _install_derived_widths() -> None:
"""Give `ParallelContext` a property per declared quotient.
They are declared in `arg_groups/fields/parallel.py`, in the same class as
the leaves they are computed from -- unannotated, so `collect_input_fields`
leaves them off the record while they still live where the namespace does. Written here as
properties rather than answered by `__getattr__` because they are read
inside compiled model code, where an attribute load is traceable and a
dynamic lookup is not.
"""
from sglang.srt.arg_groups.arg_utils import Derived from sglang.srt.arg_groups.arg_utils import Derived
from sglang.srt.arg_groups.fields.parallel import Parallel from sglang.srt.arg_groups.fields.parallel import Parallel
for name, decl in vars(Parallel).items(): return {
if not isinstance(decl, Derived): name: decl for name, decl in vars(Parallel).items() if isinstance(decl, Derived)
continue }
def _install_parallel_properties() -> None:
"""Give `ParallelContext` a property per name that is not a config leaf.
The quotients are declared in `arg_groups/fields/parallel.py`, beside the
leaves they are computed from; the ranks and group handles are declared in
`_LIVE_READS`, because no configuration carries them. Properties rather
than names left to `__getattr__` because the class surface is what the
guards introspect -- `hasattr(ParallelContext, "tp_group")` and
`vars(ParallelContext)` are how the tests check the set from the class
side -- and because each one carries its `Derived.doc`.
Every one of them resolves through `_read`, so there is a single priority
chain rather than one per kind of name.
"""
docs = {name: decl.doc for name, decl in _derived_widths().items()}
for name in list(docs) + list(_LIVE_READS):
def getter(self, _name=name): def getter(self, _name=name):
return self._derived_width(_name) return self._read(_name)
getter.__name__ = name getter.__name__ = name
getter.__doc__ = decl.doc getter.__doc__ = docs.get(name)
setattr(ParallelContext, name, property(getter)) setattr(ParallelContext, name, property(getter))
_install_derived_widths() _install_parallel_properties()
class _FlagGroupBase(msgspec.Struct): class _FlagGroupBase(msgspec.Struct):
@@ -1796,7 +1728,7 @@ def reset_context() -> None:
``server_args`` and install fresh ``Flags`` and ``Resources``. ``server_args`` and install fresh ``Flags`` and ``Resources``.
``parallel`` holds the permanently-overridden derived widths, which go ``parallel`` holds the permanently-overridden derived widths, which go
with the lifecycle that set them: `_derived_width` prefers them over the with the lifecycle that set them: `_read` prefers them over the
published leaves, so leaving one behind lets the next test read the published leaves, so leaving one behind lets the next test read the
previous topology. previous topology.
""" """
@@ -1805,7 +1737,7 @@ def reset_context() -> None:
_CONTEXT._overrides_log = [] _CONTEXT._overrides_log = []
_CONTEXT._publish_role = None _CONTEXT._publish_role = None
_CONTEXT.parallel._config = None _CONTEXT.parallel._config = None
_CONTEXT.parallel.clear_derived_widths() _CONTEXT.parallel.clear_stamp()
_CONTEXT.flags = Flags() _CONTEXT.flags = Flags()
_CONTEXT.resources = Resources() _CONTEXT.resources = Resources()
_CONTEXT.forward = ForwardFlags() _CONTEXT.forward = ForwardFlags()
@@ -7,7 +7,6 @@ across representative configurations. A field that moves without a declaration
or is projected into the wrong namespace therefore fails on observed state. or is projected into the wrong namespace therefore fails on observed state.
""" """
import ast
import copy import copy
import json import json
import os import os
@@ -135,31 +134,16 @@ def _stash_overlay(server_args):
def _live_topology_leaves(): def _live_topology_leaves():
"""Names `ParallelContext` serves from the live topology, not the config. """Names `ParallelContext` serves from the live topology, not the config.
Read out of the class: each shadowed name arrives as `self._v("<name>", Read out of `_LIVE_READS`, which is where those names are declared.
<getter>)`. Inferring them from "did the read raise" is wrong -- it only Inferring them from "did the read raise" is wrong -- it only raises while
raises while the process groups are missing, so in a process where an the process groups are missing, so in a process where an earlier test built
earlier test built them the property answers the *live* size and a leaf them the property answers the *live* size and a leaf check reads it as a
check reads it as a config mismatch (`parallel.tp_size: bag=1 config mismatch (`parallel.tp_size: bag=1 resolution=2`). Whether a name is
resolution=2`). Whether they are shadowed is a property of the class, not shadowed is a property of the declaration, not of the process.
of the process.
""" """
tree = ast.parse((_SRT / "runtime_context.py").read_text(encoding="utf-8-sig")) from sglang.srt.runtime_context import _LIVE_READS
parallel = next(
node return frozenset(_LIVE_READS)
for node in ast.walk(tree)
if isinstance(node, ast.ClassDef) and node.name == "ParallelContext"
)
names = set()
for node in ast.walk(parallel):
if (
isinstance(node, ast.Call)
and isinstance(node.func, ast.Attribute)
and node.func.attr == "_v"
and node.args
and isinstance(node.args[0], ast.Constant)
):
names.add(node.args[0].value)
return frozenset(names)
class TestResolutionDeclarations(CustomTestCase): class TestResolutionDeclarations(CustomTestCase):
+174 -16
View File
@@ -163,11 +163,11 @@ class TestStampedRanks(_IsolatedOverrides):
def setUp(self): def setUp(self):
super().setUp() super().setUp()
parallel = get_parallel() parallel = get_parallel()
self._saved_derived = dict(parallel._derived) self._saved_derived = dict(parallel._stamp)
parallel.clear_derived_widths() parallel.clear_stamp()
self.addCleanup( self.addCleanup(
lambda: ( lambda: (
parallel.clear_derived_widths(), parallel.clear_stamp(),
parallel.override_permanently(**self._saved_derived), parallel.override_permanently(**self._saved_derived),
) )
) )
@@ -232,6 +232,164 @@ class TestStampedRanks(_IsolatedOverrides):
self.assertEqual(parallel.attn_dp_rank, 11) self.assertEqual(parallel.attn_dp_rank, 11)
class TestEveryDeclaredParallelNameIsStatable(_IsolatedOverrides):
"""The overridable set is read from the declarations, not maintained by hand.
A hand-kept list can hold a name the class does not answer, or miss one it
does; either way `override()` refuses or accepts the wrong thing with
nothing to say so. The three tests below check the set against the
declarations from both sides.
"""
def test_every_declared_name_can_be_stated_and_reads_back(self):
from sglang.srt.runtime_context import _parallel_fields
names = sorted(_parallel_fields())
# Sizes, ranks, groups and the configured leaves of the namespace.
self.assertGreater(len(names), 30)
parallel = get_parallel()
for name in names:
sentinel = object()
with parallel.override(**{name: sentinel}):
self.assertIs(getattr(parallel, name), sentinel, msg=name)
def test_every_name_the_class_answers_for_is_in_the_set(self):
"""Cross-check from the other side: the class's own surface.
Derived from the class rather than from the same declarations the set
is built from, so a source dropped out of `_parallel_fields` shows up
here instead of agreeing with itself.
"""
from sglang.srt.runtime_context import _parallel_fields
answered = {
name
for name, value in vars(ParallelContext).items()
if isinstance(value, property)
}
self.assertTrue(answered)
self.assertEqual(answered - _parallel_fields(), set())
def test_a_live_name_is_never_also_answered_from_the_bag(self):
"""The two answer differently, so a name in both would make the read
order -- not the declaration -- decide which one a caller gets.
The bag carries the declared quotients as well as the operator's
leaves, and both are ahead of the live getter once a configuration is
published: a name in `_LIVE_READS` and in either of them would answer
from the getter before publish and from the bag after."""
from sglang.srt.runtime_context import (
_LIVE_READS,
_derived_widths,
_parallel_config_leaves,
)
self.assertEqual(set(_LIVE_READS) & _parallel_config_leaves(), set())
self.assertEqual(set(_LIVE_READS) & set(_derived_widths()), set())
def test_an_undeclared_name_is_refused(self):
with self.assertRaises(ValueError):
with get_parallel().override(not_a_parallel_name=1):
pass
class TestReadsWithoutAPublishedConfig(_IsolatedOverrides):
"""The namespace has to answer in a process that publishes nothing.
`multimodal_gen` lends its own TP group to shared `srt` layers from a
process with no `srt` config to publish against, and those layers ask for
`attn_tp_size` anyway -- through code `multimodal_gen` does not own, which
is why grepping that package for `get_parallel()` finds nothing while the
read plainly happens.
"""
def setUp(self):
super().setUp()
parallel = get_parallel()
self._saved_stamp = dict(parallel._stamp)
self.addCleanup(
lambda: (
parallel.clear_stamp(),
parallel.override_permanently(**self._saved_stamp),
)
)
reset_context()
self.addCleanup(reset_context)
def test_a_stamped_width_reads_with_nothing_published(self):
parallel = get_parallel()
self.assertIsNone(parallel._config)
parallel.override_permanently(
**derive_parallel_widths(
tp_size=2,
attn_cp_size=1,
attn_dp_size=1,
moe_ep_size=1,
moe_dp_size=1,
dcp_size=1,
dcp_enabled=False,
)
)
self.assertEqual(parallel.attn_tp_size, 2)
self.assertEqual(parallel.moe_tp_size, 2)
def test_an_unstamped_width_still_names_the_cause(self):
"""Without a stamp there is nothing to answer with, and the failure
has to say so rather than invent a width."""
with self.assertRaisesRegex(RuntimeError, r"not available"):
get_parallel().attn_tp_size
class TestPrivateAttributeProbing(_IsolatedOverrides):
def test_probing_a_private_name_does_not_recurse(self):
"""`copy` and `pickle` probe for hooks before `__init__` has run.
`__getattr__` reaches for `self._config`, so if it did not refuse
underscore names outright, probing one on a half-built instance would
recurse until the stack ran out.
"""
fresh = ParallelContext.__new__(ParallelContext) # slots unset
for probe in ("_config", "_stamp", "_overrides", "__deepcopy__"):
with self.assertRaises(AttributeError, msg=probe):
getattr(fresh, probe)
def test_a_built_context_survives_a_copy(self):
import copy
self.assertIsInstance(copy.copy(get_parallel()), ParallelContext)
class TestAWidthReadStaysTraceable(_IsolatedOverrides):
"""A width read inside compiled model code must stay inside the graph.
Shared layers read widths inside a compiled forward. A graph break there
is a performance regression and nothing else -- every suite stays green
through it -- so `fullgraph=True` is what turns it into a failure. This
pins the read path, whichever form it takes: the sibling leaf test
compiles names served by `__getattr__` and they trace too.
"""
def test_a_width_read_compiles_into_the_graph(self):
import torch
reset_context()
self.addCleanup(reset_context)
publish(
ServerArgs(
model_path="dummy", tp_size=8, dp_size=2, enable_dp_attention=True
),
role="test",
)
def read(x):
return x * get_parallel().attn_tp_size
# backend="eager": this pins tracing, not code generation, and stays
# runnable on a box with no inductor toolchain.
compiled = torch.compile(read, fullgraph=True, backend="eager")
self.assertEqual(compiled(torch.ones(3)).tolist(), [4.0, 4.0, 4.0])
class TestParallelOverride(_IsolatedOverrides): class TestParallelOverride(_IsolatedOverrides):
def test_override_takes_precedence(self): def test_override_takes_precedence(self):
p = get_parallel() p = get_parallel()
@@ -1512,11 +1670,11 @@ class TestDerivedWidths(_IsolatedOverrides):
def setUp(self): def setUp(self):
super().setUp() super().setUp()
parallel = get_parallel() parallel = get_parallel()
self._saved_derived = dict(parallel._derived) self._saved_derived = dict(parallel._stamp)
parallel.clear_derived_widths() parallel.clear_stamp()
self.addCleanup( self.addCleanup(
lambda: ( lambda: (
parallel.clear_derived_widths(), parallel.clear_stamp(),
parallel.override_permanently(**self._saved_derived), parallel.override_permanently(**self._saved_derived),
) )
) )
@@ -1567,14 +1725,15 @@ class TestDerivedWidths(_IsolatedOverrides):
self.assertIn("not available", str(caught.exception)) self.assertIn("not available", str(caught.exception))
def test_a_permanent_override_and_a_live_group_both_win_over_the_leaves(self): def test_a_permanent_override_and_a_live_group_both_win_over_the_leaves(self):
"""Order is permanent override, then live group, then the leaves. """Order is scoped override, then the stamp, then the published leaf.
Where a group exists it is the truth -- elastic scale-up moves the
group without a fresh override -- so the leaf derivation only No group is consulted for a width -- `test_the_group_is_never_consulted`
answers where there is none. in this class asserts that -- so a stamp is what an elastic scale-up
leaves behind, and the leaf answers only where there is none.
""" """
parallel = get_parallel() parallel = get_parallel()
parallel.override_permanently(attn_tp_size=7) parallel.override_permanently(attn_tp_size=7)
self.addCleanup(parallel.clear_derived_widths) self.addCleanup(parallel.clear_stamp)
with parallel.override(tp_size=8, attn_dp_size=2): with parallel.override(tp_size=8, attn_dp_size=2):
self.assertEqual(parallel.attn_tp_size, 7) self.assertEqual(parallel.attn_tp_size, 7)
@@ -1662,14 +1821,14 @@ class TestDerivedWidths(_IsolatedOverrides):
# Elastic scaling overrides again where it updates the live width. # Elastic scaling overrides again where it updates the live width.
parallel.override_permanently(attn_dp_size=4) parallel.override_permanently(attn_dp_size=4)
self.assertEqual(parallel.attn_dp_size, 4) self.assertEqual(parallel.attn_dp_size, 4)
parallel.clear_derived_widths() parallel.clear_stamp()
with parallel.override(tp_size=8, attn_dp_size=1): with parallel.override(tp_size=8, attn_dp_size=1):
self.assertEqual(parallel.attn_dp_size, 1) self.assertEqual(parallel.attn_dp_size, 1)
def test_reset_context_drops_the_permanent_override(self): def test_reset_context_drops_the_permanent_override(self):
"""The permanent override belongs to the lifecycle that made it. """The permanent override belongs to the lifecycle that made it.
`_derived_width` prefers it over the published leaf, so one that `_read` prefers it over the published leaf, so one that
outlived `reset_context()` would let the next test read the previous outlived `reset_context()` would let the next test read the previous
topology. topology.
""" """
@@ -1705,8 +1864,7 @@ class TestDerivedWidths(_IsolatedOverrides):
def test_recomputing_from_published_leaves_matches_the_publish_bag(self): def test_recomputing_from_published_leaves_matches_the_publish_bag(self):
"""`initialize_model_parallel` no longer overrides anything -- see """`initialize_model_parallel` no longer overrides anything -- see
`test_initialize_model_parallel_no_longer_touches_the_bag` below -- `test_initialize_model_parallel_no_longer_touches_the_bag` below --
which makes this the load-bearing half of 16-field-registry-design.md so every real caller must forward leaves that already match its own
§6e: every real caller must forward leaves that already match its own
published config, because nothing corrects a mismatch anymore. published config, because nothing corrects a mismatch anymore.
`scheduler.py`'s `ps.attn_dp_size`/`ps.moe_ep_size`/etc, and the `scheduler.py`'s `ps.attn_dp_size`/`ps.moe_ep_size`/etc, and the
weight-cache daemon's own already-published config, both do -- this weight-cache daemon's own already-published config, both do -- this
@@ -1762,7 +1920,7 @@ class TestDerivedWidths(_IsolatedOverrides):
self.assertEqual(published, recomputed) self.assertEqual(published, recomputed)
def test_initialize_model_parallel_no_longer_touches_the_bag(self): def test_initialize_model_parallel_no_longer_touches_the_bag(self):
"""§6e, landed: `initialize_model_parallel` used to recompute and """`initialize_model_parallel` used to recompute and
permanently override the six derived widths on `get_parallel()` permanently override the six derived widths on `get_parallel()`
after building its groups; that call is gone. Publish a placeholder after building its groups; that call is gone. Publish a placeholder
config (tp_size defaults to 1), then build real groups at a config (tp_size defaults to 1), then build real groups at a