core.optimizer.layer_sharded_a2a#
all_to_all routing utilities for layer-sharded Muon.
Layer sharding assigns every 2D weight one Newton-Schulz “home” rank inside the (GTP_remat x TP) weight-shard domain — i.e. the GTP domain (GTP = TP x GTP_remat). The forward exchanges route each rank’s local momentum shards so the home assembles the complete matrix; the backward exchanges scatter the orthogonalized result back to the original shards.
The exchange runs in two stages (route_to_ns_home / route_from_ns_home):
one all_to_all over the GTP_remat group (dim 0), then one over the TP group
(along partition_dim), reusing the process groups Megatron already has.
These helpers are axis-generic — the caller invokes them once with the GTP_remat
group and once with the TP group — so their arguments are named for the role
(group, shard_dim), not for a specific axis. A trivial group (None or
size 1) performs no communication, so a 1-D domain costs a single all_to_all
per direction.
All functions support heterogeneous parameter shapes and uneven home assignments (ranks may own zero matrices in a given exchange, receiving zero-size all_to_all splits).
Module Contents#
Functions#
Group param indices by their NS home rank in the group. |
|
Forward all_to_all for layer sharding: redistribute momentum shards. |
|
Backward all_to_all for layer sharding: distribute NS results as shards. |
API#
- core.optimizer.layer_sharded_a2a.params_by_home(
- num_params: int,
- home_of: dict,
- size: int,
Group param indices by their NS home rank in the group.
result[r]lists the params homed on rankrin increasing index order, the order every exchange uses.home_ofmust cover every index (LayerShardedMuon’s exchange plan always does). Assignments may be uneven: ranks with no params get empty lists and receive zero-size all_to_all splits.
- core.optimizer.layer_sharded_a2a.route_to_ns_home(
- momentum_list: list[torch.Tensor],
- param_to_home_rank: dict,
- group: torch.distributed.ProcessGroup | None,
- shard_dim: int = 0,
- plan: dict | None = None,
Forward all_to_all for layer sharding: redistribute momentum shards.
Each rank holds a (P/S, Q) momentum shard of every param. This redistributes them so each rank ends up with the complete (P, Q) momentum for its assigned subset.
- Parameters:
momentum_list – List of momentum tensors, one per param. Each has shape (P/S, Q) where S = group size (this rank’s shard along
shard_dim).param_to_home_rank – Dict mapping param index -> NS home rank in
group.group – The process group to communicate within (the GTP_remat group in stage 1, the TP group in stage 2). None means size 1: no exchange.
shard_dim – Dimension the shards split (0 for the GTP_remat stage; the param’s partition_dim for the TP stage).
plan – Optional mutable dict caching the routing metadata (index groupings, split sizes, unpack offsets), which is a pure function of shapes, homes and group size — all static across steps. Pass an empty dict on the first call (it is filled) and the same dict on later calls (the metadata rebuild is skipped; only the data movement runs). The CALLER owns validity: reuse a plan only while the participating params, their shapes and their homes are unchanged. None (default) rebuilds every call.
- Returns:
- complete_momentums: List of complete (P, Q) tensors for params assigned to this rank, in the order they appear in momentum_list. - my_param_indices: Indices into momentum_list for params assigned to this rank.
- Return type:
Tuple of
- core.optimizer.layer_sharded_a2a.route_from_ns_home(
- ns_results: list[torch.Tensor],
- my_param_indices: list[int],
- momentum_list: list[torch.Tensor],
- param_to_home_rank: dict,
- group: torch.distributed.ProcessGroup | None,
- shard_dim: int = 0,
- plan: dict | None = None,
Backward all_to_all for layer sharding: distribute NS results as shards.
Each NS-home rank has complete (P, Q) NS results for its assigned params. This redistributes them so every rank gets its (P/S, Q) shard for every param.
- Parameters:
ns_results – Complete (P, Q) NS result tensors, one per assigned param, in the order of my_param_indices.
my_param_indices – Indices into momentum_list for params assigned to this rank.
momentum_list – List of original momentum tensors (provides shapes).
param_to_home_rank – Dict mapping param index -> NS home rank in
group.group – The process group to communicate within (see the fwd docstring). None means size 1: no exchange.
shard_dim – Dimension the shards split (see the fwd docstring).
plan – Optional routing-metadata cache (see the fwd docstring; same ownership rules). The shape-invariant precondition check also runs only when the plan is built.
- Returns:
List of (P/S, Q) NS update shards, one per param in momentum_list order. None for params that did not participate (should not occur in normal usage).