core.optimizer.layer_sharded_muon#

Layer-sharded Muon: layer sharding for Newton-Schulz over the GTP_remat x TP domain.

Instead of per-weight all-gather + redundant full-matrix NS on every rank, each weight is assigned one NS home rank in the (GTP_remat x TP) domain. Two all_to_all stages route the momentum shards so the home holds the complete (P, Q) matrix, Newton-Schulz runs there with zero communication and zero redundancy — the exact same full-matrix NS as duplicated mode — and two reverse all_to_all stages scatter the result back to the original shards. All collectives use the existing gtp_remat / tp process groups.

Module Contents#

Classes#

ParamSharding

Which axes of a param group’s (gtp_remat, tp) domain shard a 2-D weight.

ParamShardSpec

How the model sharded one parameter over a (gtp_remat, tp) domain.

_GroupExchangePlan

Everything step() needs for one param group that is constant across steps.

LayerShardedMuon

Muon with layer sharding over the GTP_remat x TP domain.

Functions#

_validate_ns_config

Reject a Newton-Schulz configuration the installed stack cannot run.

_phase

Phase-level NVTX range (active only under --profile with --nvtx-ranges).

tp_partition_dim

The dim TP splits p along, or None when p is replicated across TP.

Data#

API#

core.optimizer.layer_sharded_muon.__all__#

[‘LayerShardedMuon’, ‘ParamShardSpec’, ‘ParamSharding’, ‘tp_partition_dim’]

core.optimizer.layer_sharded_muon.logger#

‘getLogger(…)’

core.optimizer.layer_sharded_muon._BATCHED_NS_MIN_EO_VERSION#

‘0.3.0’

core.optimizer.layer_sharded_muon._BATCHED_SYRK_MIN_EO_VERSION#

‘0.5.0a0’

core.optimizer.layer_sharded_muon._SYRK_VALIDATED_SMS#

((8, 0), (9, 0), (10, 0), (10, 3))

core.optimizer.layer_sharded_muon._validate_ns_config(use_syrk: bool, ns_batch_size: int) → None#

Reject a Newton-Schulz configuration the installed stack cannot run.

One place, one exception type (ValueError, like the parent’s own gates) for every “this install cannot do that” condition: batched Newton-Schulz needs an emerging-optimizers that accepts 3-D input, and SYRK needs Triton >= 3.4.0, an SM emerging-optimizers validated the kernel on, and for batched chunks the batched SYRK kernel. The Triton / SM conditions mirror the guard in emerging-optimizers’ Muon.__init__ (which this class does not inherit from) but raise instead of downgrading: a run must not silently switch kernels, and with them numerics, with the hardware or the installed version. The parent’s emerging-optimizers version gate for use_syrk itself still applies. Follow-up: generalize the SYRK conditions to every Muon mode in TensorParallelMuon.

core.optimizer.layer_sharded_muon._phase(name: str)#

Phase-level NVTX range (active only under --profile with --nvtx-ranges).

Kernel-name classification cannot separate the forward from the reverse all_to_all, nor the momentum update from the weight update, so the step is annotated explicitly; a handful of ranges per step, not per param.

class core.optimizer.layer_sharded_muon.ParamSharding(*args, **kwds)#

Bases: enum.Enum

Which axes of a param group’s (gtp_remat, tp) domain shard a 2-D weight.

Derived from the model’s sharding attributes and the domain sizes (:meth:ParamShardSpec.from_param); the exchanges follow from it: stage 1 runs for a gtp_remat axis, stage 2 for a TP axis.

Initialization

REPLICATED#

‘replicated’

Whole on every rank of the domain (MoE router, latent projections, or any param in a single-rank domain): no exchange, every rank runs the same local Newton-Schulz.

GTP_REMAT#

‘gtp_remat’

dim 0 split over gtp_remat: stage-1 exchange, then every TP peer of the home column holds the full matrix and runs the same Newton-Schulz.

TP#

‘tp’

Split over TP only (the domain has no gtp_remat axis): stage-2 exchange only.

GTP_REMAT_AND_TP#

‘gtp_remat_and_tp’

gtp_remat shards of a TP shard: stage 1 assembles the TP-local matrix on the home column, stage 2 assembles the full matrix on the (g_home, t_home) rank.

core.optimizer.layer_sharded_muon.tp_partition_dim(p: torch.Tensor) → int | None#

The dim TP splits p along, or None when p is replicated across TP.

tensor_model_parallel is the sharded/replicated flag and partition_dim is only meaningful when it is set, the same convention param_is_not_tensor_parallel_duplicate uses. Megatron marks duplicated-mode TE weights tensor_model_parallel=False while TE still stamps partition_dim=0 on them, so partition_dim alone misclassifies them.

class core.optimizer.layer_sharded_muon.ParamShardSpec#

How the model sharded one parameter over a (gtp_remat, tp) domain.

Holds only model-imposed facts, fixed by the forward/backward parallelism: the optimizer reads them once (:meth:from_param) and never recomputes them in step(). Everything the optimizer derives from them is a property.

  • tp_dim: the dim TP splits (0 column-parallel, 1 row-parallel) or None (:func:tp_partition_dim, ignored when the domain has no TP axis).

  • gtp_sharded: dim 0 of the TP-local shard is split over gtp_remat (is_gtp_param) and the domain has a gtp_remat axis.

  • pad_length: GTP alignment padding, trailing zero rows on the gtp-gathered TP-local dim 0 (0 when not GTP-sharded).

  • full_shape: shape of the matrix Newton-Schulz runs on (pad stripped).

tp_dim: int | None#

None

gtp_sharded: bool#

None

pad_length: int#

None

full_shape: tuple[int, ...]#

None

classmethod from_param(
p: torch.Tensor,
gtp_remat_size: int,
tp_size: int,
) → core.optimizer.layer_sharded_muon.ParamShardSpec#

Read p’s sharding for a (gtp_remat, tp) domain of the given sizes.

An axis of size 1 is absent from the domain, so its tag is ignored: with tp_size == 1 every param is TP-replicated, and in a single-rank domain everything is REPLICATED (plain local Newton-Schulz).

Raises:

ValueError – p is TP-sharded but not GTP-sharded while the domain has a gtp_remat axis. Such a param is replicated across gtp_remat, and the stage-1 exchange would concatenate its copies as dim-0 shards and silently corrupt the update.

property sharding: core.optimizer.layer_sharded_muon.ParamSharding#

Which domain axes shard the param; decides the exchanges it joins.

property ns_cost: int#

Newton-Schulz cost estimate on full_shape, ~ max(M, N) * min(M, N)^2, the weight NS-home balancing uses; non-2-D shapes count their elements.

class core.optimizer.layer_sharded_muon._GroupExchangePlan#

Everything step() needs for one param group that is constant across steps.

Built once per param identity tuple by :meth:LayerShardedMuon._build_plan and invalidated by both setters; step() only consumes it. Index spaces: i indexes the group’s grad-bearing params, n the routed sub-list, k the subset of the routed list that stage 1 delivers to this rank (stage1_routed_indices).

param_ids: tuple[int, ...]#

None

specs: list[core.optimizer.layer_sharded_muon.ParamShardSpec]#

None

replicated: list[int]#

None

routed: list[int]#

None

ns_homes: list[tuple[int, int]]#

None

g_home: dict[int, int]#

None

pad_lengths: list[int]#

None

stage1_routed_indices: list[int]#

None

tp_exchanges: dict[int, tuple[list[int], dict[int, int]]]#

None

tp_complete: list[int]#

None

route_plans: dict#

‘field(…)’

class core.optimizer.layer_sharded_muon.LayerShardedMuon(
params: torch.optim.optimizer.ParamsT,
lr: float = 0.0003,
momentum: float = 0.95,
weight_decay: float = 0.01,
*,
nesterov: bool = True,
fp32_matmul_prec: emerging_optimizers.utils.FP32MatmulPrecT = 'medium',
coefficient_type: emerging_optimizers.orthogonalized_optimizers.muon_utils.NSCoeffT = 'quintic',
num_ns_steps: int = 5,
scale_mode: emerging_optimizers.orthogonalized_optimizers.muon.MuonScaleT = 'spectral',
extra_scale_factor: float = 1.0,
gtp_remat_group: torch.distributed.ProcessGroup | None,
tp_group: torch.distributed.ProcessGroup | None = None,
ns_batch_size: int = 1,
use_syrk: bool = False,
concurrent_groups: bool = True,
use_decoupled_weight_decay: bool = True,
split_qkv: bool = False,
is_qkv_fn: Callable[[torch.Tensor], bool] | None = None,
qkv_split_shapes: list[int] | None = None,
pg_collection: megatron.core.process_groups_config.ProcessGroupCollection | None = None,
tp_mode: Literal[blockwise, duplicated, distributed, auto] = 'duplicated',
)#

Bases: megatron.core.optimizer.emerging_optimizers.TensorParallelMuon

Muon with layer sharding over the GTP_remat x TP domain.

Sharding model per 2D weight of full shape (P, Q):

  • TP shards along param.partition_dim (0 = column-parallel, 1 = row-parallel) when param.tensor_model_parallel is set; otherwise the param is TP-replicated.

  • GTP_remat shards dim 0 of the TP-local shard, for params tagged param.is_gtp_weight_remat (is_gtp_param; absent means unsharded).

  • A param sharded by neither is whole on every rank of the domain (e.g. the MoE router and latent projections): it skips both exchanges and every rank runs the same deterministic NS on its own copy.

  • GTP alignment padding (param.pad_length trailing zero rows on the gtp-gathered, TP-local dim 0) is stripped before Newton-Schulz — so the scale factor sees the true dims, matching the parent’s duplicated path bitwise — and restored before the reverse gtp_remat exchange.

Each param’s sharding is read once from these attributes (:meth:ParamShardSpec.from_param), the per-group exchange plan (homes, routing tables) is cached, and step() only consumes it.

step() runs, per param group:

  1. Momentum update on the local shard (elementwise, identical to base Muon).

  2. Stage-1 all_to_all over gtp_remat_group (dim 0): each param’s GTP_remat extent is assembled on its assigned g_home column.

  3. Stage-2 all_to_all over tp_group (along partition_dim): the full matrix is assembled on the (g_home, t_home) NS home. Params that are not TP-sharded skip this stage — every TP peer of the column already holds the full matrix and runs the same (deterministic) NS so each column can scatter its own updates.

  4. Full-matrix Newton-Schulz on the home — bit-identical to duplicated mode.

  5. Reverse stage-2 / stage-1 all_to_all scatter the scaled NS result back to every rank’s original shard, which applies p -= lr * update.

Parameters:
  • params – Parameters to optimize. Every rank in the domain must pass the same params in the same order (they hold different shards of the same weights).

  • gtp_remat_group – GTP_remat weight-shard process group (dim-0 sharding of the TP-local shard).

  • tp_group – TP process group, or None when TP is not used.

  • use_syrk – Use the Triton SYRK kernel for the two symmetric-output NS GEMMs (A = X Xᵀ and B = bA + cA²), computing one triangle only — roughly a third off total NS FLOPs for near-square matrices. Needs Triton >= 3.4.0 and a validated SM (8.0/9.0/10.0/10.3); with ns_batch_size > 1 also an emerging-optimizers with the batched SYRK kernel (>= 0.5.0a0, PR #276). Unmet requirements raise at construction; nothing downgrades silently. Only takes effect with fp32_matmul_prec="medium" and 8-aligned dims. Same math, different kernel — results differ from the GEMM path by kernel-level rounding.

  • ns_batch_size – Maximum number of same-shape matrices fused into one batched Newton-Schulz on a home (see OptimizerConfig.muon_ns_batch_size). Defaults to 1, the bit-exact per-matrix path; batches of more than one use baddbmm and lose bitwise parity with duplicated mode.

  • concurrent_groups – Run each param group’s pipeline on its own CUDA stream instead of serializing them. Groups own disjoint params and, under MoE, disjoint process groups, so nothing orders them against each other; on a single stream one group’s all_to_all stall blocks the other group’s Newton-Schulz even though the GPU is idle. The ops and their order within a group are unchanged, so results are unaffected, but the transient buffers of all groups are live at once – lower ns_batch_size or set this to False if that pushes peak memory too high. No effect with fewer than two param groups or without CUDA. Requires the groups’ domains to be disjoint; groups sharing a (gtp_remat, tp) domain are automatically serialized (NCCL forbids concurrent collectives on one communicator).

  • args (All other) – same as :class:TensorParallelMuon. In particular split_qkv / is_qkv_fn / qkv_split_shapes, tp_mode and pg_collection only take effect on the paths that delegate to the parent (the empty-param_ns_homes fallback and the degenerate single-rank domain, both of which run the parent’s TP-aware full-matrix Newton-Schulz). “Degenerate” means the LAYER-SHARDING domain (gtp_remat_size * tp_size) is trivial, not that the step is collective-free: with a non-trivial pg_collection.tp and partition_dim-tagged params (direct API only), the parent path still issues TP collectives.

.. note::

None for either process group means “no group / size 1”, not torch’s “the default group” — a missing expert group must not silently become the whole world.

Usage::

optimizer = LayerShardedMuon(params, lr=3e-4, gtp_remat_group=gtp, tp_group=tp)
optimizer.set_param_ns_homes({id(p): (g_home, t_home) for ...})
optimizer.step()

Initialization

set_param_ns_homes(param_ns_homes: dict[int, tuple[int, int]]) → None#

Set the NS home for each param (by id).

Parameters:

param_ns_homes – Maps id(param) -> (g_home, t_home): the rank in gtp_remat_group and in tp_group that runs NS for it. t_home is ignored for params that are not TP-sharded and when tp_group is None. An empty mapping is valid: params in single-rank domains run local Newton-Schulz, any other routed param falls back to round-robin homes.

set_group_process_groups(
group_process_groups: dict[int, tuple],
) → None#

Override the (GTP, TP) process groups per param group.

Different param groups can be sharded over different domains — e.g. under MoE, expert weights are sharded over the expert GTP/TP groups while dense weights use the dense ones. Groups absent from the mapping fall back to the gtp_remat_group / tp_group passed to the constructor.

Parameters:

group_process_groups – Maps the index of a self.param_groups entry to (gtp_remat_group, tp_group). Either entry may be None (treated as size 1 / not available).

_apply_update(
p: torch.Tensor,
update: torch.Tensor,
lr: float,
) → None#

Apply one weight update through the base-class hook points.

OrthogonalizedOptimizer.step() brackets every p.add_ with pre_weight_update_fn_inplace / post_weight_update_fn_inplace; this helper keeps layer sharding’s overridden step() honouring them too, and keeps the two update sites (replicated, routed) from diverging. No dtype cast on purpose: the base class’s p.add_(orth_grad, alpha=-lr) (the third path, taken by the no-homes fallback) computes the fused multiply-add in the promoted precision and downcasts once on store, so casting here first would give bf16 params different rounding on the layer-sharded paths than on the fallback and than TensorParallelMuon’s duplicated mode. TODO: forward the weight_update_hook constructor parameter once the emerging-optimizers pin moves past EO #224.

_run_ns(full_by_k: dict) → dict#

Full-matrix Newton-Schulz per home-owned matrix, batched by shape.

Same-shape matrices are stacked into one batched NS: under MoE a home owns hundreds of identically shaped expert weights, and the per-matrix loop is dominated by kernel-launch overhead. Batches are capped at ns_batch_size to bound the transient stack memory. A batch of one stays 2-D, so the unbatched numerics are preserved exactly whenever nothing is actually batched.

_param_group_streams() → list | None#

Per-group CUDA streams, or None when the groups must stay serialized.

Concurrency requires the groups’ communication domains to be disjoint: NCCL serializes collectives per communicator, so two param groups sharing a (gtp_remat, tp) domain (e.g. two dense groups produced by a per-layer lr override) issuing collectives from different streams can interleave and deadlock. Such configurations fall back to serialized execution (bitwise-neutral, see the concurrent_groups docstring).

step(closure: Any = None) → None#

Run one optimizer step: momentum update, shard exchange to the NS homes, full-matrix Newton-Schulz there, reverse exchange, weight update (see the class docstring for the per-stage breakdown).

_plan_for(
group_index: int,
params: list[torch.Tensor],
gtp_remat_group,
tp_group,
) → core.optimizer.layer_sharded_muon._GroupExchangePlan#

The cached exchange plan for this group’s grad-bearing params.

Rebuilt only when the param set changes (a direct-API caller dropping a grad); the wired path hits the cache every step.

_build_plan(
params: list[torch.Tensor],
gtp_remat_group,
tp_group,
) → core.optimizer.layer_sharded_muon._GroupExchangePlan#

Read every param’s sharding once and assemble the routing tables.

_step_groups(
streams: list | None,
ready: torch.cuda.Event | None,
) → None#