-
Notifications
You must be signed in to change notification settings - Fork 9k
config: the derived parallel widths are computed from the leaves #36790
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -131,6 +131,57 @@ def _parallel_config_leaves() -> frozenset: | |
| ) | ||
|
|
||
|
|
||
| def derive_attention_widths( | ||
| *, tp_size: int, attn_cp_size: int, dp_size: int, enable_dp_attention: bool | ||
| ) -> tuple: | ||
| """(attn_dp_size, attn_tp_size) from the leaves. | ||
|
|
||
| Split out because the rank computation in | ||
| `dp_attention.compute_dp_attention_world_info` needs the same two numbers | ||
| and must not carry a second copy of the arithmetic. | ||
| """ | ||
| attn_dp_size = dp_size if enable_dp_attention else 1 | ||
| return attn_dp_size, tp_size // attn_dp_size // attn_cp_size | ||
|
|
||
|
|
||
| def derive_parallel_widths( | ||
| *, | ||
| tp_size: int, | ||
| attn_cp_size: int, | ||
| attn_dp_size: int, | ||
| moe_ep_size: int, | ||
| moe_dp_size: int, | ||
| dcp_size: int, | ||
| dcp_enabled: bool, | ||
| ) -> dict: | ||
| """The parallel widths no flag sets, from the leaves that do. | ||
|
|
||
| `tp_size` and its siblings are configured; these are quotients of them, so | ||
| the arithmetic lives here rather than being read back off the group | ||
| coordinators. | ||
|
|
||
| `world_size` is not among them: it is not a quotient, and `get_world_size()` | ||
| answers with the live WORLD group, which stays right through an elastic | ||
| scale-up that a stamp taken at group build would not survive. | ||
| """ | ||
| return { | ||
| "attn_dp_size": attn_dp_size, | ||
| # `attn_dp_size` is already the effective width (1 when DP attention is | ||
| # off), so the flag is spent here; a caller passing the raw `dp_size` | ||
| # leaf with the attention disabled would get tp/dp/cp instead of tp/1/cp. | ||
| "attn_tp_size": derive_attention_widths( | ||
| tp_size=tp_size, | ||
| attn_cp_size=attn_cp_size, | ||
| dp_size=attn_dp_size, | ||
| enable_dp_attention=True, | ||
| )[1], | ||
| "moe_ep_size": moe_ep_size, | ||
| "moe_tp_size": tp_size // moe_ep_size // moe_dp_size, | ||
| "dcp_enabled": dcp_enabled, | ||
| "attn_dcp_size": dcp_size if dcp_enabled else 1, | ||
| } | ||
|
|
||
|
|
||
| class ParallelContext: | ||
| """Parallel-topology namespace: one spelling per name. | ||
|
|
||
|
|
@@ -154,11 +205,12 @@ class ParallelContext: | |
| different names rather than two answers to one name. | ||
| """ | ||
|
|
||
| __slots__ = ("_overrides", "_config") | ||
| __slots__ = ("_overrides", "_config", "_derived") | ||
|
|
||
| def __init__(self): | ||
| self._overrides = {} | ||
| self._config = None # parallel config bag, wired at publish | ||
| self._derived = {} # widths stamped when the groups are built | ||
|
ch-wan marked this conversation as resolved.
|
||
|
|
||
| def __getattr__(self, name): | ||
| if name.startswith("_"): | ||
|
|
@@ -181,6 +233,45 @@ def _v(self, name, getter): | |
| overrides = self._overrides | ||
| return overrides[name] if name in overrides else getter() | ||
|
|
||
| def stamp_derived_widths(self, **widths) -> None: | ||
| """Record the widths derived from the leaves, as the groups are built. | ||
|
|
||
| `initialize_model_parallel` computes the set through | ||
| `derive_parallel_widths` and hands it here; `initialize_dp_attention` | ||
| stamps `attn_dp_size` again once it knows the effective width, and | ||
| elastic EP restamps it where it already updates the live one. A stamped | ||
| width is what the readers answer with. | ||
| """ | ||
| self._derived.update(widths) | ||
|
|
||
| def clear_derived_widths(self) -> None: | ||
| self._derived.clear() | ||
|
|
||
| def _derived_width(self, name, getter): | ||
| """A width the leaves imply: the stamp, else the live group. | ||
|
|
||
| The fallback keeps a process that installed groups without going | ||
| through `initialize_model_parallel` working. When neither is there, | ||
| the failure says which of the two is missing rather than surfacing a | ||
| group getter's bare assertion. | ||
| """ | ||
| overrides = self._overrides | ||
| if name in overrides: | ||
| return overrides[name] | ||
| derived = self._derived | ||
| if name in derived: | ||
| return derived[name] | ||
| try: | ||
| return getter() | ||
| except (AssertionError, AttributeError, RuntimeError) as exc: | ||
| raise RuntimeError( | ||
| f"derived parallel width {name!r} is not available: it is " | ||
| "computed from the configured leaves when the process groups " | ||
| "are built (initialize_model_parallel / " | ||
| "initialize_dp_attention), and neither a stamp nor a live " | ||
| "group is present" | ||
| ) from exc | ||
|
|
||
| @contextmanager | ||
| def override(self, **kwargs): | ||
| """Temporarily force parallel values, restoring on exit. Validates keys and | ||
|
|
@@ -213,7 +304,9 @@ def pp_rank(self) -> int: | |
|
|
||
| @property | ||
| def moe_ep_size(self) -> int: | ||
| return self._v("moe_ep_size", _ps().get_moe_expert_parallel_world_size) | ||
| return self._derived_width( | ||
| "moe_ep_size", _ps().get_moe_expert_parallel_world_size | ||
| ) | ||
|
|
||
| @property | ||
| def moe_ep_rank(self) -> int: | ||
|
|
@@ -225,15 +318,19 @@ def moe_dp_rank(self) -> int: | |
|
|
||
| @property | ||
| def moe_tp_size(self) -> int: | ||
| return self._v("moe_tp_size", _ps().get_moe_tensor_parallel_world_size) | ||
| return self._derived_width( | ||
| "moe_tp_size", _ps().get_moe_tensor_parallel_world_size | ||
| ) | ||
|
|
||
| @property | ||
| def moe_tp_rank(self) -> int: | ||
| return self._v("moe_tp_rank", _ps().get_moe_tensor_parallel_rank) | ||
|
|
||
| @property | ||
| def attn_tp_size(self) -> int: | ||
| return self._v("attn_tp_size", _ps().get_attn_tensor_model_parallel_world_size) | ||
| return self._derived_width( | ||
| "attn_tp_size", _ps().get_attn_tensor_model_parallel_world_size | ||
| ) | ||
|
|
||
| @property | ||
| def attn_tp_rank(self) -> int: | ||
|
|
@@ -254,11 +351,11 @@ def getter(): | |
| return False | ||
| return _ps().get_dcp_world_size() > 1 | ||
|
|
||
| return self._v("dcp_enabled", getter) | ||
| return self._derived_width("dcp_enabled", getter) | ||
|
|
||
| @property | ||
| def attn_dcp_size(self) -> int: | ||
| return self._v( | ||
| return self._derived_width( | ||
| "attn_dcp_size", | ||
| lambda: _ps().get_dcp_world_size() if self.dcp_enabled else 1, | ||
| ) | ||
|
|
@@ -271,7 +368,7 @@ def attn_dcp_rank(self) -> int: | |
|
|
||
| @property | ||
| def attn_dp_size(self) -> int: | ||
| return self._v("attn_dp_size", _dp().get_attention_dp_size) | ||
| return self._derived_width("attn_dp_size", _dp().get_attention_dp_size) | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
When code enters Useful? React with 👍 / 👎. |
||
|
|
||
| @property | ||
| def attn_dp_rank(self) -> int: | ||
|
|
@@ -1512,14 +1609,17 @@ def reset_context() -> None: | |
| """Clear the context-owned store (unit-test teardown): drop the published | ||
| ``server_args`` and install fresh ``Flags`` and ``Resources``. | ||
|
|
||
| Wrapper subsystems (``parallel``) hold no state and are unaffected. | ||
| ``parallel`` holds the stamped derived widths, which go with the lifecycle | ||
| that stamped them: `_derived_width` prefers the stamp over the live group, | ||
| so leaving one behind lets the next test read the previous topology. | ||
| """ | ||
| _CONTEXT._server_args = None | ||
| _CONTEXT._config_bags = None | ||
| _adaptive_draft_token_bound.cache_clear() | ||
| _CONTEXT._overrides_log = [] | ||
| _CONTEXT._publish_role = None | ||
| _CONTEXT.parallel._config = None | ||
| _CONTEXT.parallel.clear_derived_widths() | ||
| _CONTEXT.flags = Flags() | ||
| _CONTEXT.resources = Resources() | ||
| _CONTEXT.forward = ForwardFlags() | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.