Skip to content
147 changes: 143 additions & 4 deletions megatron/core/optimizer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,10 @@
from torch.optim import AdamW as Adam, SGD

from .ademamix import AdEMAMix
from .prodigy import Prodigy
from .mars import MARS, exists
from .adopt import ADOPT
from .muon import Muon

from megatron.core import mpu

Expand Down Expand Up @@ -52,6 +56,8 @@ def _get_param_groups(
min_lr: float,
decoupled_lr: Optional[float],
decoupled_min_lr: Optional[float],
muon_matched_adamw_rms: Optional[float],
use_muon: bool = False,
) -> List[Dict]:
"""Create parameter groups for optimizer.

Expand Down Expand Up @@ -82,6 +88,7 @@ def _get_param_groups(

# Map (wd_mult, lr_mult, is_expert_parallel, is_decoupled_lr) to params.
params_map = {}
muon_params_map = {}
for model_chunk in model_chunks:
for name, param in model_chunk.named_parameters():
if not param.requires_grad:
Expand Down Expand Up @@ -117,10 +124,25 @@ def _get_param_groups(
):
is_decoupled_lr = True

key = (wd_mult, _lr_mult, is_expert_parallel, is_decoupled_lr)
if key not in params_map:
params_map[key] = []
params_map[key].append(param)
# key = (wd_mult, _lr_mult, is_expert_parallel, is_decoupled_lr)
# if key not in params_map:
# params_map[key] = []
# params_map[key].append(param)
bias_flag = name.endswith(".bias")
shape_flag = param.dim() == 2
embedding_flag = "embedding" in name or "output_layer" in name
muon_flag = use_muon and shape_flag \
and (not bias_flag) and (not embedding_flag)
if muon_flag:
key = (wd_mult, _lr_mult, is_expert_parallel)
if key not in muon_params_map:
muon_params_map[key] = []
muon_params_map[key].append(param)
else:
key = (wd_mult, _lr_mult, is_expert_parallel, is_decoupled_lr)
if key not in params_map:
params_map[key] = []
params_map[key].append(param)

param_groups = []
for (wd_mult, _lr_mult, is_expert_parallel, is_decoupled_lr), params in params_map.items():
Expand All @@ -142,6 +164,20 @@ def _get_param_groups(
decoupled_min_lr=decoupled_min_lr,
)

for (wd_mult, _lr_mult, is_expert_parallel), params in muon_params_map.items():
if len(params) == 0:
continue
param_groups.append(
{
'params': params,
'wd_mult': wd_mult,
'lr_mult': _lr_mult,
'is_expert_parallel': is_expert_parallel,
'use_muon': True,
'is_decoupled_lr': False,
}
)

return param_groups


Expand Down Expand Up @@ -224,6 +260,8 @@ def _get_param_groups_and_buffers(
min_lr=config.min_lr,
decoupled_lr=config.decoupled_lr,
decoupled_min_lr=config.decoupled_min_lr,
muon_matched_adamw_rms=config.muon_matched_adamw_rms,
use_muon = config.optimizer == 'muon',
)
param_groups = list(filter(filter_fn, param_groups))
buffers = {}
Expand Down Expand Up @@ -326,7 +364,108 @@ def init_state_fn(opt, config=None):
else:
opt.initialize_state(p)

elif config.optimizer == 'prodigy':
kwargs = {
"params": param_groups,
"lr": config.lr,
"weight_decay": config.weight_decay,
"betas": (config.adam_beta1, config.adam_beta2),
"beta3": config.prodigy_beta3,
"decouple": config.prodigy_decouple,
"use_bias_correction": config.prodigy_use_bias_correction,
"safeguard_warmup": config.prodigy_safeguard_warmup,
"fsdp_in_use": config.prodigy_fsdp_in_use,
}

optimizer = Prodigy(**kwargs)

def init_state_fn(opt, config=None):
for group in opt.param_groups:
for p in group['params']:
if 'step' not in opt.state[p]:
opt.state[p]['step'] = 0
opt.state[p]['s'] = torch.zeros_like(p.data).detach()
opt.state[p]['p0'] = p.detach().clone()
opt.state[p]['exp_avg'] = torch.zeros_like(p.data).detach()
opt.state[p]['exp_avg_sq'] = torch.zeros_like(p.data).detach()
else:
opt.initialize_state(p)

elif config.optimizer == 'mars':
kwargs = {
"params": param_groups,
"lr": config.mars_lr,
"betas": (config.mars_beta1, config.mars_beta2),
"weight_decay": config.weight_decay,
"amsgrad": config.mars_amsgrad,
"gamma": config.mars_vr_gamma,
"is_approx": config.mars_is_approx,
"mars_type": config.mars_type,
"optimize_1d": config.mars_optimize_1d,
"lr_1d": config.lr,
"betas_1d": (config.adam_beta1, config.adam_beta2),
"weight_decay_1d": config.mars_weight_decay_1d,
}

optimizer = MARS(**kwargs)

def init_state_fn(opt, config=None):
for group in opt.param_groups:
for p in filter(lambda p: exists(p.grad), group["params"]):
amsgrad = group["amsgrad"]
if len(opt.state[p]) <= 1:
opt.state[p]["step"] = 0
opt.state[p]["exp_avg"] = torch.zeros_like(p.data)
opt.state[p]["last_grad"] = torch.zeros_like(p.data)
opt.state[p]["exp_avg_sq"] = torch.zeros_like(p.data)
if amsgrad:
opt.state[p]["max_exp_avg_sq"] = torch.zeros_like(p.data)
if amsgrad and "max_exp_avg_sq" not in opt.state[p]:
opt.state[p]["max_exp_avg_sq"] = torch.zeros_like(p.data)
else:
opt.initialize_state(p)

elif config.optimizer == 'adopt':
kwargs = {
"params": param_groups,
"lr": config.lr,
"weight_decay": config.weight_decay,
"betas": (config.adam_beta1, config.adam_beta2),
"eps": config.adopt_eps,
"decouple": config.adopt_decouple,
}

optimizer = ADOPT(**kwargs)

def init_state_fn(opt, config=None):
for group in opt.param_groups:
for p in group['params']:
if len(opt.state[p]) == 0:
opt.state[p]['exp_avg'] = torch.zeros_like(p.data)
opt.state[p]['exp_avg_sq'] = torch.zeros_like(p.data)
else:
opt.initialize_state(p)

elif config.optimizer == 'muon':
optimizer = Muon(param_groups,
lr=config.lr, weight_decay=config.weight_decay,
matched_adamw_rms=config.muon_matched_adamw_rms,
momentum=config.muon_momentum,
nesterov=config.muon_nesterov,
ns_steps=config.muon_ns_steps,
adamw_betas=(config.adam_beta1, config.adam_beta2),
adamw_eps=config.adam_eps)

def init_state_fn(opt, config=None):
for group in opt.param_groups:
for p in group['params']:
if len(opt.state[p]) == 0:
if config is None or not config.use_precision_aware_optimizer:
opt.state[p]['exp_avg'] = torch.zeros_like(p.data)
opt.state[p]['exp_avg_sq'] = torch.zeros_like(p.data)
else:
opt.initialize_state(p)

elif config.optimizer == 'sgd':
optimizer = SGD(
param_groups,
Expand Down
106 changes: 106 additions & 0 deletions megatron/core/optimizer/adopt.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
"""
Here is an original implementation of ADOPT.
Source: https://github.com/iShohei220/adopt
"""

import torch


def exists(val):
return val is not None


from typing import Callable, Optional, Tuple

import torch


def adopt_clip_fn(step: int) -> float:
return step ** 0.25


class ADOPT(torch.optim.Optimizer):
def __init__(
self,
params,
lr: float = 1e-3,
betas: Tuple[float, float] = (0.9, 0.9999),
eps: float = 1e-6,
clip_lambda: Optional[Callable[[int], float]] = adopt_clip_fn,
weight_decay: float = 0.0,
decouple: bool = True,
):
if not 0.0 <= lr:
raise ValueError(f"Invalid learning rate: {lr}")
if not 0.0 <= eps:
raise ValueError(f"Invalid epsilon value: {eps}")
if not 0.0 <= betas[0] < 1.0:
raise ValueError(f"Invalid beta parameter at index 0: {betas[0]}")
if not 0.0 <= betas[1] < 1.0:
raise ValueError(f"Invalid beta parameter at index 1: {betas[1]}")
if not 0.0 <= weight_decay:
raise ValueError(f"Invalid weight_decay value: {weight_decay}")

defaults = dict(
lr=lr,
betas=betas,
eps=eps,
weight_decay=weight_decay,
decouple=decouple,
clip_lambda=clip_lambda,
step=0,
)
super().__init__(params, defaults)

@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()

for group in self.param_groups:
group['step'] += 1
step = group['step']
beta1, beta2 = group["betas"]
lr = group["lr"]

for p in group["params"]:
if p.grad is None:
continue
grad = p.grad

if grad.is_sparse:
raise RuntimeError("ADOPT does not support sparse gradients")

state = self.state[p]

if len(state) ==0:
state["exp_avg"] = torch.zeros_like(p, memory_format=torch.preserve_format)
state["exp_avg_sq"] = torch.zeros_like(p, memory_format=torch.preserve_format)

exp_avg = state["exp_avg"]
exp_avg_sq = state["exp_avg_sq"]

if step == 1:
exp_avg_sq.addcmul_(grad, grad)
continue

if group["weight_decay"] != 0:
if group["decouple"]:
p.data.mul_(1 - lr * group["weight_decay"])
else:
grad = grad.add(p, alpha=group["weight_decay"])

denom = torch.clamp(exp_avg_sq.sqrt(), group["eps"])
normed_grad = grad.div(denom)

if group["clip_lambda"] is not None:
clip = group["clip_lambda"](step)
normed_grad.clamp_(-clip, clip)

exp_avg.lerp_(normed_grad, 1 - beta1)
p.data.add_(exp_avg, alpha=-lr)
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)

return loss
Loading