Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions deepspeed/runtime/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -2123,6 +2123,7 @@ def _configure_optimizer(self, client_optimizer, model_parameters):
log_dist(f"DeepSpeed Basic Optimizer = {basic_optimizer.__class__.__name__}", ranks=[0])

optimizer_wrapper = self._do_optimizer_sanity_check(basic_optimizer)
self._check_muon_can_reach_its_parameters(basic_optimizer, optimizer_wrapper)

if optimizer_wrapper == ZERO_OPTIMIZATION:
self.optimizer = self._configure_zero_optimizer(basic_optimizer)
Expand All @@ -2147,6 +2148,27 @@ def _configure_optimizer(self, client_optimizer, model_parameters):
self.compression_scheduler = self._configure_compression_scheduler()
self.quantizer = self._configure_quantization()

def _check_muon_can_reach_its_parameters(self, basic_optimizer, optimizer_wrapper):
"""Refuse the one wrapper that hands Muon flat partitions and does not orthogonalize them.

`MuonWithAuxAdam.step` tells the two cases apart by shape: a matrix is the weight itself
and is orthogonalized there, a 1-D tensor is a ZeRO partition whose update the ZeRO
optimizer already applied. `BF16_Optimizer` breaks that reading - it replaces the param
groups with flat fp32 partitions (`param_group['params'] = [self.fp32_groups_flat_partition[i]]`)
and knows nothing about `use_muon`, so the update is never applied and the step is SGD.

The original shapes are not recoverable from `step`, so this is a refusal rather than a
fix; implementing Muon inside BF16_Optimizer is its own change. Reached by bf16 with
`grad_accum_dtype: fp32` at ZeRO stage 1.
"""
if not isinstance(basic_optimizer, MuonWithAuxAdam) or optimizer_wrapper != BFLOAT16:
return
raise ZeRORuntimeException(
"Muon cannot be used with the BF16_Optimizer, which this configuration selects: bf16 "
"with grad_accum_dtype fp32 at ZeRO stage 1. That optimizer hands Muon flat fp32 "
"partitions and never applies the Newton-Schulz update, so training would silently "
"proceed as SGD. Drop grad_accum_dtype, or use ZeRO stage 2 or 3.")

def _configure_autoep_folding_optimizer_gradient_reduction(self):
configure = getattr(self.optimizer, "configure_autoep_folding_tp_gradient_reduction", None)
if configure is None:
Expand Down
23 changes: 20 additions & 3 deletions deepspeed/runtime/zero/muon/muon_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
try:
from deepspeed.runtime.zero.muon.original_muon import MuonWithAuxAdam as BaseMuonWithAuxAdam
from deepspeed.runtime.zero.muon.original_muon import adam_update
from deepspeed.runtime.zero.muon.original_muon import muon_update
except ImportError:
pass

Expand Down Expand Up @@ -60,11 +61,27 @@ def step(self, closure=None, step_id=None):
loss = closure()
for group in self.param_groups:
if group["use_muon"]:
# we move the muon update part to the deepspeed's optimizer since the parameter here is a flat version
# thus not suitable for muon update
for p in group["params"]:
if p.grad is None:
p.grad = torch.zeros_like(p) # force synchronization
if p.dim() < 2:
# A flat ZeRO partition. ZeRO 1/2 orthogonalizes in get_flat_partition and
# ZeRO-3 in its sub-group loop, so the gradient already holds the update
# and only the weight decay and step size are left to apply.
update = p.grad
else:
# The weight itself, so nothing has orthogonalized it: no ZeRO optimizer is
# in play. Muon has to run here or the step degenerates to SGD.
state = self.state[p]
if len(state) == 0:
state["momentum_buffer"] = torch.zeros_like(p)
update = muon_update(p.grad,
state["momentum_buffer"],
beta=group["momentum"],
ns_method=group.get("ns_method", "gram"),
is_expert_group=getattr(p, "is_expert_group", False))
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(p.grad.reshape(p.shape), alpha=-group["lr"])
p.add_(update.reshape(p.shape), alpha=-group["lr"])

aux_param_groups = [group for group in self.param_groups if not group["use_muon"]]
if self.aux_optimizer is not None:
Expand Down
183 changes: 183 additions & 0 deletions tests/unit/runtime/zero/test_muon_without_zero_optimizer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,183 @@
# SPDX-License-Identifier: Apache-2.0

# DeepSpeed Team
"""Newton-Schulz has to run whether or not a ZeRO optimizer is there to do it.

`MuonWithAuxAdam.step` applied an update it assumed had been orthogonalized already, which holds
under ZeRO - the parameters it sees there are flat partitions and `get_flat_partition` or the
ZeRO-3 sub-group loop did the work. At stage 0, the default, no ZeRO optimizer exists to have
done it, and applying a raw gradient is SGD. Training runs and the loss falls, which is why
counting the Newton-Schulz calls is the assertion that means something here.
"""

import contextlib

import pytest
import torch

import deepspeed
from deepspeed.accelerator import get_accelerator
import deepspeed.runtime.zero.muon.original_muon as original_muon
from deepspeed.runtime.zero.utils import ZeRORuntimeException
from unit.common import DistributedTest

NS_KERNELS = ("zeropower_via_gram_newtonschulz", "zeropower_via_newtonschulz5")


@contextlib.contextmanager
def counting_newton_schulz():
"""Counts every Newton-Schulz call, whichever kernel the config selects.

Patched inside the test body rather than in a fixture: `DistributedTest` runs the body in a
worker process that a fixture in the parent would not reach.
"""
calls = []
originals = {name: getattr(original_muon, name) for name in NS_KERNELS}

def counted(kernel):

def wrapper(*args, **kwargs):
calls.append(1)
return kernel(*args, **kwargs)

return wrapper

for name, kernel in originals.items():
setattr(original_muon, name, counted(kernel))
try:
yield calls
finally:
for name, kernel in originals.items():
setattr(original_muon, name, kernel)


def _model():
return torch.nn.Sequential(torch.nn.Linear(32, 32, bias=False), torch.nn.Linear(32, 32, bias=False))


def _config(stage, dtype="fp32"):
config = {
"train_micro_batch_size_per_gpu": 2,
"gradient_accumulation_steps": 1,
"gradient_clipping": 0.0,
"optimizer": {
"type": "Muon",
"params": {
"lr": 0.02
}
},
}
if stage is not None:
config["zero_optimization"] = {"stage": stage, "reduce_scatter": stage != 3}
if dtype != "fp32":
config[dtype] = {"enabled": True}
if dtype == "fp16":
config[dtype]["initial_scale_power"] = 4
return config


def _skip_if_unsupported(dtype):
"""Mirror the check the engine itself makes.

`_do_sanity_check` raises `Type fp16 is not supported on your device.` on
`not get_accelerator().is_fp16_supported()`, which is a different predicate from
`supported_dtypes()` -- the cpu-torch-latest runner reports fp16 in the latter and
False from the former, so guarding on the wrong one still fails there.
"""
supported = {
"fp16": get_accelerator().is_fp16_supported,
"bf16": get_accelerator().is_bf16_supported,
}.get(dtype)
if supported is not None and not supported():
pytest.skip(f"{dtype} not supported on this accelerator")


class TestMuonRunsWithoutAZeroOptimizer(DistributedTest):
world_size = 1

@pytest.mark.parametrize("dtype", ["fp32", "bf16", "fp16"])
def test_newton_schulz_runs_at_stage_zero(self, dtype):
"""Every stage-0 wrapper: unwrapped for fp32, FP16_UnfusedOptimizer for bf16 and fp16.

Each hands `step` the weight itself rather than a flat partition, so nothing upstream has
orthogonalized it. On master all three do zero orthogonalizations and train as SGD.
"""
_skip_if_unsupported(dtype)
model = _model()
engine, _, _, _ = deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(0, dtype))

# Counting starts after initialize: FP16_UnfusedOptimizer steps once at construction to
# allocate state, and that call must not be what the assertion below is satisfied by.
with counting_newton_schulz() as calls:
x = torch.ones(2, 32, device=engine.device, dtype=next(engine.module.parameters()).dtype)
engine.backward(engine(x).square().sum())
engine.step()

assert len(calls) == 2, f"Newton-Schulz ran {len(calls)} times for two Muon matrices; expected one each"

def test_the_default_config_runs_muon(self):
"""`zero_optimization.stage` defaults to 0, so this is the plainest Muon config there is."""
model = _model()
engine, _, _, _ = deepspeed.initialize(model=model, model_parameters=model.parameters(), config=_config(None))

with counting_newton_schulz() as calls:
x = torch.ones(2, 32, device=engine.device)
engine.backward(engine(x).square().sum())
engine.step()

assert len(calls) == 2, f"Newton-Schulz ran {len(calls)} times; the default config trained as SGD"

@pytest.mark.parametrize("stage", [1, 2, 3])
def test_newton_schulz_runs_on_the_supported_stages(self, stage):
"""The positive control, and the assertion the existing tests were missing.

They check that training progresses, which SGD does too. Counting the orthogonalizations
is what distinguishes Muon from the update it degenerates to.
"""
model = _model()
engine, _, _, _ = deepspeed.initialize(model=model, model_parameters=model.parameters(), config=_config(stage))

with counting_newton_schulz() as calls:
x = torch.ones(2, 32, device=engine.device, dtype=next(engine.module.parameters()).dtype)
engine.backward(engine(x).square().sum())
engine.step()

assert len(calls) == 2, \
f"Newton-Schulz ran {len(calls)} times for two Muon matrices; the step was not Muon"


class TestMuonRefusesBF16Optimizer(DistributedTest):
"""The one wrapper that hands Muon flat partitions without orthogonalizing them.

`BF16_Optimizer` replaces the param groups with flat fp32 partitions and knows nothing about
`use_muon`, so the shape test in `step` reads them as "ZeRO already did the update" and the
step is SGD. The original shapes are not recoverable there, so this is refused rather than
fixed. Selected by bf16 with `grad_accum_dtype: fp32` at ZeRO stage 1.
"""
world_size = 1

def test_bf16_optimizer_with_muon_is_refused(self):
_skip_if_unsupported("bf16")
model = _model()
config = _config(1, "bf16")
config["data_types"] = {"grad_accum_dtype": "fp32"}

with pytest.raises(ZeRORuntimeException, match="BF16_Optimizer"):
deepspeed.initialize(model=model, model_parameters=model.parameters(), config=config)

def test_the_same_config_without_grad_accum_dtype_still_runs_muon(self):
"""The neighbouring config, so the refusal is shown to be narrow."""
_skip_if_unsupported("bf16")
model = _model()
engine, _, _, _ = deepspeed.initialize(model=model,
model_parameters=model.parameters(),
config=_config(1, "bf16"))

with counting_newton_schulz() as calls:
x = torch.ones(2, 32, device=engine.device, dtype=next(engine.module.parameters()).dtype)
engine.backward(engine(x).square().sum())
engine.step()

assert len(calls) == 2
Loading