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#

params_by_home

Group param indices by their NS home rank in the group.

route_to_ns_home

Forward all_to_all for layer sharding: redistribute momentum shards.

route_from_ns_home

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,
) → list[list[int]]#

Group param indices by their NS home rank in the group.

result[r] lists the params homed on rank r in increasing index order, the order every exchange uses. home_of must 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,
) → tuple[list[torch.Tensor], list[int]]#

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,
) → list[torch.Tensor | 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).