Spaces:
Sleeping
Sleeping
File size: 8,451 Bytes
a5f2e46 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 | """Combined AdamW + Muon optimizer.
The two underlying optimizers are kept as separate ``torch.optim`` instances
internally; this class is a thin wrapper that exposes the standard
``torch.optim.Optimizer`` interface (``step``, ``zero_grad``, ``param_groups``,
``state_dict`` / ``load_state_dict``) so that LR schedulers, gradient
accumulation, and HuggingFace ``Accelerator`` treat it as a single optimizer.
Routing rule (when constructed from an ``nn.Module``):
* parameters whose name contains an ``embedding_keyword`` → AdamW (low LR)
* parameters whose name contains a ``head_keyword`` → AdamW (no decay)
* remaining parameters with ``ndim >= 2`` → Muon
* remaining parameters with ``ndim < 2`` (biases/norms) → AdamW (no decay)
"""
from __future__ import annotations
from typing import Iterable
import torch
from torch import nn
from torch.optim import AdamW, Muon, Optimizer
class AdamWMuon(Optimizer):
"""Wrapper that drives an inner ``AdamW`` and ``Muon`` in lock-step.
Parameters
----------
model:
``nn.Module`` whose parameters will be auto-routed by name.
lr:
Base AdamW learning rate (used for the ``other`` 1D parameter group).
lr_embedding, lr_head:
Per-group AdamW LR overrides. Default to ``lr``.
lr_muon:
Muon LR for ≥2D body weights. Default ``= 3 * lr`` (Muon tolerates a
higher LR than AdamW because updates are orthogonalised).
weight_decay:
AdamW weight decay applied to the embedding group. Heads and 1D
params get ``0.0`` by convention.
weight_decay_muon:
Muon weight decay. Default = ``weight_decay``.
embedding_keywords / head_keywords:
Substrings used to classify parameters by ``named_parameters`` key.
"""
def __init__(
self,
model: nn.Module,
*,
lr: float = 3e-4,
lr_embedding: float | None = None,
lr_head: float | None = None,
lr_muon: float | None = None,
weight_decay: float = 0.1,
weight_decay_muon: float | None = None,
betas: tuple[float, float] = (0.9, 0.95),
eps: float = 1e-8,
muon_momentum: float = 0.95,
muon_nesterov: bool = True,
muon_ns_steps: int = 5,
embedding_keywords: Iterable[str] = ("embedding", "cls_token"),
head_keywords: Iterable[str] = ("lm_head", "value_head"),
) -> None:
if not isinstance(model, nn.Module):
raise TypeError(
f"AdamWMuon expects an nn.Module to auto-route parameters by name, "
f"got {type(model).__name__}."
)
lr_embedding = lr if lr_embedding is None else lr_embedding
lr_head = lr if lr_head is None else lr_head
lr_muon = 3.0 * lr if lr_muon is None else lr_muon
weight_decay_muon = weight_decay if weight_decay_muon is None else weight_decay_muon
embedding_keywords = tuple(embedding_keywords)
head_keywords = tuple(head_keywords)
emb_params: list[torch.Tensor] = []
head_params: list[torch.Tensor] = []
other_params: list[torch.Tensor] = []
muon_params: list[torch.Tensor] = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
if any(kw in name for kw in embedding_keywords):
emb_params.append(param)
elif any(kw in name for kw in head_keywords):
head_params.append(param)
elif param.ndim >= 2:
muon_params.append(param)
else:
other_params.append(param)
adamw_groups: list[dict] = []
if emb_params:
adamw_groups.append({
"params": emb_params, "lr": lr_embedding,
"weight_decay": weight_decay, "name": "adamw_embedding",
})
if other_params:
adamw_groups.append({
"params": other_params, "lr": lr,
"weight_decay": 0.0, "name": "adamw_other",
})
if head_params:
adamw_groups.append({
"params": head_params, "lr": lr_head,
"weight_decay": 0.0, "name": "adamw_head",
})
if not adamw_groups:
raise ValueError("No AdamW-eligible parameters found in model.")
self.adamw = AdamW(adamw_groups, lr=lr, betas=betas, eps=eps,
weight_decay=weight_decay)
if muon_params:
self.muon: Muon | None = Muon(
[{"params": muon_params, "lr": lr_muon,
"weight_decay": weight_decay_muon, "name": "muon_body"}],
lr=lr_muon, weight_decay=weight_decay_muon,
momentum=muon_momentum, nesterov=muon_nesterov,
ns_steps=muon_ns_steps,
)
else:
self.muon = None
# Initialise the ``Optimizer`` base with a placeholder group so
# ``isinstance(opt, Optimizer)`` checks (Accelerate, schedulers) pass
# and ``self.state`` / ``self.defaults`` exist. We then expose the
# *inner* groups via ``param_groups`` so schedulers mutate them in
# place and both sub-optimizers see the new LR.
all_params = emb_params + other_params + head_params + muon_params
self._init_done = False
super().__init__([{"params": all_params}], defaults={"lr": lr})
self.param_groups = self._combined_param_groups()
self._init_done = True
# ── helpers ─────────────────────────────────────────────────────────
def _combined_param_groups(self) -> list[dict]:
groups = list(self.adamw.param_groups)
if self.muon is not None:
groups += list(self.muon.param_groups)
return groups
# ── Optimizer API ───────────────────────────────────────────────────
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
self.adamw.step()
if self.muon is not None:
self.muon.step()
return loss
def zero_grad(self, set_to_none: bool = True) -> None:
self.adamw.zero_grad(set_to_none=set_to_none)
if self.muon is not None:
self.muon.zero_grad(set_to_none=set_to_none)
def state_dict(self) -> dict:
return {
"adamw": self.adamw.state_dict(),
"muon": self.muon.state_dict() if self.muon is not None else None,
}
def load_state_dict(self, state_dict: dict) -> None:
self.adamw.load_state_dict(state_dict["adamw"])
if self.muon is not None and state_dict.get("muon") is not None:
self.muon.load_state_dict(state_dict["muon"])
# Re-link in case sub-optimizers rebuilt their group dicts.
self.param_groups = self._combined_param_groups()
def add_param_group(self, param_group) -> None:
if not getattr(self, "_init_done", False):
# Called from ``Optimizer.__init__`` with the placeholder group.
super().add_param_group(param_group)
return
raise NotImplementedError(
"add_param_group is not supported on AdamWMuon after construction; "
"build a new optimizer with the full parameter set instead."
)
def __repr__(self) -> str:
n_emb = sum(p.numel() for g in self.adamw.param_groups
if g.get("name") == "adamw_embedding" for p in g["params"])
n_head = sum(p.numel() for g in self.adamw.param_groups
if g.get("name") == "adamw_head" for p in g["params"])
n_other = sum(p.numel() for g in self.adamw.param_groups
if g.get("name") == "adamw_other" for p in g["params"])
n_muon = (sum(p.numel() for g in self.muon.param_groups for p in g["params"])
if self.muon is not None else 0)
return (
f"AdamWMuon(emb={n_emb:,}, other_1d={n_other:,}, "
f"head={n_head:,}, muon_2d={n_muon:,})"
)
# Backwards-compat alias for the previous skeleton class.
AdamWMuonOptim = AdamWMuon |