Skip to content

[PyTorch] Add newton_schulz_tp optimizer step function - #2920

Merged
vcherepanov-nv merged 17 commits into
NVIDIA:mainfrom
vcherepanov-nv:muon
Aug 25, 2026
Merged

vcherepanov-nv merged 17 commits into
NVIDIA:mainfrom
vcherepanov-nv:muon

Conversation

@vcherepanov-nv

@vcherepanov-nv vcherepanov-nv commented Apr 23, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Add a newton_schulz_tp utility function on top of the existing newton_schulz. To be used in Muon optimizer step.

Fixes # (issue)

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

Please list the changes introduced in this PR:

  • Add a newton_schulz_tp function
  • refactor test_newton_schulz.py to do multiple check during a single torchrun

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

vcherepanov-nv and others added 2 commits April 23, 2026 18:50
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
@greptile-apps

greptile-apps Bot commented Apr 23, 2026 •

Copy link
Copy Markdown
Contributor

Greptile Summary

The PR adds tensor-parallel Newton–Schulz orthogonalization and restores compatibility for the former module path.

  • Supports replicated, row-partitioned, and column-partitioned tensors.
  • Moves the implementation under the optimizer package while retaining a legacy import shim.
  • Consolidates distributed numerical coverage into single- and multi-GPU worker runs.

Confidence Score: 5/5

The PR appears safe to merge.

No blocking failure remains.

Important Files Changed

Filename Overview
transformer_engine/pytorch/optimizers/newton_schulz.py Adds the tensor-parallel wrapper, process-group accessor, replicated gather path, and row/column partition handling without an eligible blocking follow-up issue.
transformer_engine/pytorch/newton_schulz.py Restores the legacy module path through a public-symbol re-export shim.
transformer_engine/pytorch/init.py Exposes the new tensor-parallel function from the optimizer implementation.
tests/pytorch/distributed/run_newton_schulz.py Consolidates distributed numerical cases and adds coverage for tensor-parallel modes and replicated inputs.
tests/pytorch/distributed/test_newton_schulz.py Launches consolidated one- and multi-GPU worker runs through the active Python interpreter.

Flowchart

%%{init: {'theme': 'neutral'}}%%
flowchart TD
  A["newton_schulz_tp(x, partition_dim, tp_mode)"] --> B{"Replicated input?"}
  B -- Yes --> C["Shard larger dimension"]
  B -- No --> D{"Duplicated mode?"}
  D -- Yes --> E["Gather full tensor"]
  D -- No --> F{"Row partition?"}
  F -- Yes --> G["Transpose local shard"]
  F -- No --> H["Use column shard directly"]
  C --> I["Distributed Newton–Schulz kernel"]
  E --> I
  G --> I
  H --> I
  I --> J["Gather/transpose/slice as needed"]
  J --> K["Copy result into x"]
Loading

Reviews (9): Last reviewed commit: "[pre-commit.ci] auto fixes from pre-comm..." | Re-trigger Greptile

Comment on lines +186 to +191
def step(self, closure=None):
"""Perform a single optimization step."""
loss = None
if closure is not None:
loss = closure()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Closure called inside @torch.no_grad(), preventing gradient computation

closure() is invoked while torch.no_grad() is active. Any loss.backward() call inside the closure will silently produce zero/no gradients. The standard PyTorch pattern (used in SGD, Adam, etc.) is to wrap the closure in with torch.enable_grad():.

Suggested change
def step(self, closure=None):
"""Perform a single optimization step."""
loss = None
if closure is not None:
loss = closure()
@torch.no_grad()
def step(self, closure=None):
"""Perform a single optimization step."""
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()

Comment on lines +28 to +33
scale_mode: str,
extra_scale_factor: float,
eps: float,
) -> torch.Tensor:
global_shape = [grad.size(0), grad.size(1)]
global_shape[partition_dim] *= world_size

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Reference global_shape incorrectly scales an already-full tensor

_reference_orthogonalize receives the full matrix (shape full_shape) but then multiplies global_shape[partition_dim] by world_size a second time. For partition_dim=1 with world_size=2 and full_shape=(96, 128) this gives global_shape=[96, 256], so get_muon_scale_factor returns max(96,256)^0.5 = 16. The optimizer, operating on the shard (96, 64), correctly reconstructs global_shape=[96, 128] and computes max(96,128)^0.5 ≈ 11.3. This √2 discrepancy means the reference cannot correctly validate the optimizer's output.

The global_shape[partition_dim] *= world_size line should be removed since the input is already the full matrix.

Comment on lines +33 to +34
if mode == "unit_rms_norm":
return (size_out / size_in) ** 0.5

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 unit_rms_norm mode can divide by zero when size_in == 0

(size_out / size_in) ** 0.5 raises ZeroDivisionError when size_in is 0. While the optimizer validates that the partition dimension is non-empty, it doesn't ensure the other dimension is non-zero. Consider adding a guard or documenting that both dimensions must be strictly positive.

Comment thread transformer_engine/pytorch/optimizers/muon.py Outdated
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
@vcherepanov-nv vcherepanov-nv changed the title [Draft] [PyTorch] Add distributed Muon optimizer [PyTorch] Add distributed Muon optimizer Apr 27, 2026
@vcherepanov-nv
vcherepanov-nv requested a review from cyanguwa April 27, 2026 18:12

@skyw skyw left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd advice NOT to expose it in public API. Keeping it in test only if that is the purpose.

Having an optimizer with most code copied invites fragmentation.

Before this, all optimizer TE provides are more optimized fused version. I'd say a highly optimized Fused Muon with similar concept can be justified, but would need more consideration because it has more dependencies on other part of the training pipeline than elementwise optimizers.

on tensor-parallel parameter shards. The local parameter shard must represent a
partition of a logical 2D matrix across the provided NCCL process group.

Args:

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Q: Does TE use numpy style docstring instead of Google style?


def __init__(
self,
params: Iterable[torch.nn.Parameter | dict],

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nit: The type here doesn't match PyTorch internal. Should be fine for the purpose of this class.

scale_mode: MuonScaleT = "spectral",
extra_scale_factor: float = 1.0,
process_group: Optional[dist.ProcessGroup] = None,
partition_dim: int = 1,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fix: partition_dim is per parameter.

raise ValueError(f"Invalid weight_decay value: {weight_decay}")
if num_ns_steps < 1:
raise ValueError(f"num_ns_steps must be at least 1, got {num_ns_steps}")
if partition_dim not in (0, 1):

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Q: Does this class intend to support non-distributed case? partition_dim would be -1 in TE in such case.


if process_group is None:
if not dist.is_initialized():
raise RuntimeError("MuonOptimizer requires torch.distributed to be initialized.")

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same question above regarding single GPU support.

if process_group is None:
if not dist.is_initialized():
raise RuntimeError("MuonOptimizer requires torch.distributed to be initialized.")
process_group = dist.group.WORLD

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggestion: This silent behavior is dangerous. If user forgot to pass the correct TP group, wrong group will be used.

Comment thread transformer_engine/pytorch/optimizers/muon.py Outdated
global_shape[partition_dim] *= world_size

orth_grad = grad.clone()
transposed = partition_dim == 0

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Attn: This is from common Row and Column wise tensor parallelism in most LLM. It would be sub optimal for anything other than that. Add comment if the assumption is made.

@vcherepanov-nv

Copy link
Copy Markdown
Collaborator Author

Having an optimizer with most code copied invites fragmentation.

The idea was to give something to users, who use TE, but not Megatron-LM. By fragmentation you mean that we want to encourage everyone to use Megatron-LM? Or that the optimizer being relatively thin thing on top of newton_schulz call, and the users should have no trouble creating it themselves?

I don't think we gain anything by putting it into tests, since we already have tests for newton_schulz call. So we need to decide whether we want this PR, or should abandon it altogether. @cyanguwa

@skyw

skyw commented Apr 28, 2026

Copy link
Copy Markdown

Having an optimizer with most code copied invites fragmentation.

The idea was to give something to users, who use TE, but not Megatron-LM. By fragmentation you mean that we want to encourage everyone to use Megatron-LM? Or that the optimizer being relatively thin thing on top of newton_schulz call, and the users should have no trouble creating it themselves?

I don't think we gain anything by putting it into tests, since we already have tests for newton_schulz call. So we need to decide whether we want this PR, or should abandon it altogether. @cyanguwa

Fragmentation means there will be different flavor of muon in emerging optimizer and TE, also a lot of copied code. TE can have stalled feature when emerging optimizer updates. Megatron-LM will always have its own version because there are implementation specific things need to be hooked together. For example, how QKV is implemetned, or fused swighlu.
For TE, I think an example of how to build a version of emerging optimizer use TE NS backend would be good to have. But providing optimizer (not fusion optimized version) confuses customers.
Having said that, I would love for TE to have a more optimized version. similar idea as fusedAdam, etc.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we move newton_schulz.py to this directory? Also, how do we expect Megatron to call us for this functionality? Thanks.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we move newton_schulz.py to this directory?

No, don't think so.

Should we move newton_schulz.py to this directory?

Megatron will call newton_shulz directly from their optimizers. This one is for other users.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd prefer moving them into something like transformer_engine/pytorch/cusolver. But I suppose that is orthogonal to this PR.

vcherepanov-nv and others added 3 commits May 1, 2026 07:27
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
@cyanguwa

cyanguwa commented May 4, 2026

Copy link
Copy Markdown
Collaborator
  1. run CI here "/te-ci torch L1"; add tests to qa/Lx_pytorch_unittest/test.sh?
  2. please create newton_schulz_tp API to include partition_dim/mode params, and for Megatron integration
  3. please test non-distributed cases (per comment above)
  4. please move newton_schulz.py to te/pytorch/optimizers; if we have more solvers to integrate to TE or more use cases of Newton-Schulz, we can definitely restructure the code, but we don't see that in the near future

@cyanguwa

cyanguwa commented May 4, 2026

Copy link
Copy Markdown
Collaborator

@skyw Just following up on the discussion above - our purpose for this PR was two-fold. One was to provide an equivalent newton_schulz_tp API for Megatron; the other one was to provide a dialed-down version of Muon Optimizer class so direct TE users can access the Newton-Schulz solver. I understand this may cause divergence in TE and Megatron's Muon support, but we do want to expose this feature to direct users of TE as well. Hope that helps. Thanks.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd prefer moving them into something like transformer_engine/pytorch/cusolver. But I suppose that is orthogonal to this PR.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LAUNCH_CMD = ["torchrun", f"--nproc_per_node={NUM_PROCS}"]


def _run_test(dtype: str, partition_dim: int, weight_decay_mode: str) -> None:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Each torchrun launch is somewhat expensive. Instead of launching a separate torchrun for each test case, it's better to launch a single torchrun instance and to perform multiple tests internally. See distributed/test_fusible_ops.py for an example.

vcherepanov-nv and others added 4 commits May 4, 2026 23:21
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
@skyw

skyw commented May 5, 2026

Copy link
Copy Markdown

@skyw Just following up on the discussion above - our purpose for this PR was two-fold. One was to provide an equivalent newton_schulz_tp API for Megatron; the other one was to provide a dialed-down version of Muon Optimizer class so direct TE users can access the Newton-Schulz solver. I understand this may cause divergence in TE and Megatron's Muon support, but we do want to expose this feature to direct users of TE as well. Hope that helps. Thanks.

Megatron is wrapped over emgering-optimizers with megatron specific details, like TP and how QKV are organized. The most optimizer logic is in emerging-optimizers. Could TE do the same? I understand introducing a new dependency may have concern, let me know. The biggest concern is actually large portion of duplicated code.

What I would favor is having newton_schulz_tp in one release, test it out. and have a well optimized version of Muon optimizer class in the next. There are a lot of optimizations (fusion, batch, graph capturability etc.) that can and I believe should go into TE.

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
@vcherepanov-nv
vcherepanov-nv requested a review from ksivaman as a code owner May 18, 2026 19:51
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
@vcherepanov-nv vcherepanov-nv changed the title [PyTorch] Add distributed Muon optimizer [PyTorch] Add newton_schulz_tp optimizer step function May 18, 2026
vcherepanov-nv and others added 2 commits May 18, 2026 22:48
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
partition_dim: Optional[int] = None,
tp_mode: Literal["duplicated", "distributed"] = "duplicated",
) -> None:
"""Compute tensor-parallel Newton-Schulz orthogonalization in-place.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is there some requirement for the operation to be in-place? I see we are doing a few x.copy_()s in the code and was just wondering if it's necessary.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, it's intentional

ctx : CusolverMpCtx
cuSolverMp context created for the tensor-parallel process group.
num_iterations : int, optional
Number of Newton-Schulz iterations. Default: 5.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should this be optional, or mandatory? Shouldn't users always be specifying the exact number of iterations they want to run NS with? Otherwise, they might get successful runs but incorrect results silently?


output_shards = [torch.empty_like(local_work) for _ in range(ctx.nranks)]
dist.all_gather(output_shards, local_work, group=ctx.group)
output = torch.cat(output_shards, dim=1)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I feel the functionality of _orthogonalize_replicated() is a bit odd. In the case of partition_dim is None, we probably don't need any chunking or all-gathering, given that tensor x is local, non-distributed. In the case of tp_mode == "duplicated", we seem to be all-gathering, chunking, and then all-gathering again - is this necessary? Looking at emerging-optimizers, it seems to be all lowering down to a simple newton_schulz() call?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

It's necessary for this cuSolverMp. cuSolverMp interprets each rank’s input as a column shard; the wrapper must shard replicated inputs before the call and gather the distributed result afterward. Emerging-Optimizers can avoid this only because its local PyTorch implementation is different.

tp_mode=tp_mode,
)
else:
newton_schulz(x_local, ctx, num_iterations, coefficients=coefficients)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Non-TP is a special case of TP, right? So we probably can just call newton_schulz_tp() here regardless of whether api == "tp"? We just need to set up partition_dim and tp_mode properly for the non-TP case?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think it makes sense to test both APIs, even if they end up running the same thing internally.

return (5e-2, 5e-2)
if check == "orthogonality" and world_size == 1:
return (2e-2, 2e-2)
return (1e-2, 1e-2)

@cyanguwa cyanguwa Jun 1, 2026 •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Are the tolerances set a bit high? Usually we have 1e-2 for BF16 and even smaller for FP32. Could you do a bit of research and tighten them up if necessary? A couple of reference points are:

https://github.com/NVIDIA-NeMo/Emerging-Optimizers/blob/main/tests/test_muon_utils.py
https://github.com/vcherepanov-nv/TransformerEngine/blob/3c4dfffb788f01009bf6741d25d4b69078e95d13/tests/pytorch/test_permutation.py#L192

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Are we committing to only supporting Newton-Shulz, and to only use it for Muon? It would be more general to put this code in a transformer_engine.pytorch.cusolvermp submodule.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think in principle that's a good idea, but if there's realistically no other use case of it in TE, it might be better to simplify the file structure?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could you check if there are other optimizers in Emerging-Optimizers that could be using cusolvermp?

tp_mode=tp_mode,
)
else:
newton_schulz(x_local, ctx, num_iterations, coefficients=coefficients)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think it makes sense to test both APIs, even if they end up running the same thing internally.

@KshitijLakhani KshitijLakhani added 2.18 and removed 2.17 labels Jun 25, 2026
@vcherepanov-nv vcherepanov-nv added 2.19 and removed 2.18 labels Jul 30, 2026
vcherepanov-nv and others added 3 commits August 21, 2026 20:39
# Conflicts:
#	tests/pytorch/distributed/run_newton_schulz.py
#	transformer_engine/pytorch/optimizers/newton_schulz.py
Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

@cyanguwa cyanguwa left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Looks like the CI is clean as well. Thanks!

@vcherepanov-nv
vcherepanov-nv merged commit 5df5a0d into NVIDIA:main Aug 25, 2026
12 of 16 checks passed
fheinecke pushed a commit that referenced this pull request Aug 25, 2026
* Add distributed Muon optimizer

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix Muon closure and reference test

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* Fix Muon optimizer distributed API handling

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* Fix Muon optimizer docs and params typing

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Add tensor-parallel Newton-Schulz wrapper

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* Move Newton-Schulz wrapper into optimizers

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* Use tensor-parallel Newton-Schulz in Muon

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Limit Muon branch to Newton-Schulz TP

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* Run Newton-Schulz distributed cases in one launch

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* Add single-GPU Newton-Schulz coverage

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix Newton-Schulz compatibility and replicated coverage

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Signed-off-by: Vladimir Cherepanov <vcherepanov@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
(cherry picked from commit 5df5a0d)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants