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#
Which axes of a param group’s (gtp_remat, tp) domain shard a 2-D weight. |
|
How the model sharded one parameter over a (gtp_remat, tp) domain. |
|
Everything |
|
Muon with layer sharding over the GTP_remat x TP domain. |
Functions#
Reject a Newton-Schulz configuration the installed stack cannot run. |
|
Phase-level NVTX range (active only under |
|
The dim TP splits |
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 foruse_syrkitself 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
--profilewith--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.EnumWhich 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
palong, or None whenpis replicated across TP.tensor_model_parallelis the sharded/replicated flag andpartition_dimis only meaningful when it is set, the same conventionparam_is_not_tensor_parallel_duplicateuses. Megatron marks duplicated-mode TE weightstensor_model_parallel=Falsewhile TE still stampspartition_dim=0on them, sopartition_dimalone 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 instep(). 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,
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 == 1every param is TP-replicated, and in a single-rank domain everything is REPLICATED (plain local Newton-Schulz).- Raises:
ValueError –
pis 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_planand invalidated by both setters;step()only consumes it. Index spaces:iindexes the group’s grad-bearing params,nthe routed sub-list,kthe 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.TensorParallelMuonMuon 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) whenparam.tensor_model_parallelis 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_lengthtrailing 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, andstep()only consumes it.step()runs, per param group:Momentum update on the local shard (elementwise, identical to base Muon).
Stage-1 all_to_all over
gtp_remat_group(dim 0): each param’s GTP_remat extent is assembled on its assignedg_homecolumn.Stage-2 all_to_all over
tp_group(alongpartition_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.Full-matrix Newton-Schulz on the home — bit-identical to duplicated mode.
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ᵀandB = 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); withns_batch_size > 1also an emerging-optimizers with the batched SYRK kernel (>= 0.5.0a0, PR #276). Unmet requirements raise at construction; nothing downgrades silently. Only takes effect withfp32_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 usebaddbmmand 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_sizeor 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 particularsplit_qkv/is_qkv_fn/qkv_split_shapes,tp_modeandpg_collectiononly take effect on the paths that delegate to the parent (the empty-param_ns_homesfallback 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-trivialpg_collection.tpand partition_dim-tagged params (direct API only), the parent path still issues TP collectives.
.. note::
Nonefor 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 ingtp_remat_groupand intp_groupthat runs NS for it.t_homeis ignored for params that are not TP-sharded and whentp_groupis 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],
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_grouppassed to the constructor.- Parameters:
group_process_groups – Maps the index of a
self.param_groupsentry 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,
Apply one weight update through the base-class hook points.
OrthogonalizedOptimizer.step()brackets everyp.add_withpre_weight_update_fn_inplace/post_weight_update_fn_inplace; this helper keeps layer sharding’s overriddenstep()honouring them too, and keeps the two update sites (replicated, routed) from diverging. No dtype cast on purpose: the base class’sp.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 theweight_update_hookconstructor 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_sizeto 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,
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,
Read every param’s sharding once and assemble the routing tables.
- _step_groups(
- streams: list | None,
- ready: torch.cuda.Event | None,