diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 7821094702f..e6ab74223f8 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -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 @@ -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. @@ -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: @@ -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(): @@ -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 @@ -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 = {} @@ -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, diff --git a/megatron/core/optimizer/adopt.py b/megatron/core/optimizer/adopt.py new file mode 100644 index 00000000000..7b6c8ced179 --- /dev/null +++ b/megatron/core/optimizer/adopt.py @@ -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 diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 252ee9646cc..653340a0433 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -22,6 +22,10 @@ HAVE_APEX_OR_TE = False from .ademamix import AdEMAMix +from .prodigy import Prodigy +from .mars import MARS +from .adopt import ADOPT +from .muon import Muon, MuonDistMeta from .. import tensor_parallel from ..config_logger import has_config_logger_enabled, log_config_to_disk @@ -44,6 +48,7 @@ _zero_grad_group_helper, ) from .optimizer_config import OptimizerConfig +from megatron.core.parallel_state import get_tensor_model_parallel_group try: # This will be used when "--fp8-param-gather" is enabled. @@ -148,6 +153,7 @@ def _build_model_gbuf_param_range_map( sub_param_start = max(0, gbuf_world_range.start - param_world_start) sub_param_range = param_local_range.normalize(sub_param_start) param_range_map[param] = { + "world_indexes": (param_world_start, param_world_end), "gbuf_world": param_world_range, "gbuf_world_in_bucket": param_world_range_in_bucket, "gbuf_local": param_local_range, @@ -335,13 +341,22 @@ def _build_model_and_main_param_groups( shard_fp32_groups.append(shard_fp32_params_this_group) shard_fp32_from_float16_groups.append(shard_fp32_from_float16_params_this_group) + dist_metas = {} + for model_param in group_range["params"]: assert model_param.requires_grad gbuf_index, dtype, bucket_index = param_gbuf_map[model_param] gbuf_range = gbuf_ranges[gbuf_index][dtype][bucket_index] - param_range = gbuf_range["param_map"][model_param]["param"] + param_gbuf_ranges = gbuf_range["param_map"][model_param] + param_range = param_gbuf_ranges["param"] + + # gen dist meta + param_world_indexes = param_gbuf_ranges["world_indexes"] + tp_split_dim = -1 if getattr(model_param, 'tensor_model_parallel', False) else \ + getattr(model_param, 'partition_dim') + dist_meta = MuonDistMeta(gbuf_index, bucket_index, model_param.shape, param_world_indexes, tp_split_dim) # fp16, bf16 params. if model_param.type() in ['torch.cuda.HalfTensor', 'torch.cuda.BFloat16Tensor']: @@ -395,6 +410,9 @@ def _build_model_and_main_param_groups( shard_float16_params_this_group.append(shard_model_param) shard_fp32_from_float16_params_this_group.append(shard_main_param) + # add to dist metas + dist_metas[shard_main_param] = dist_meta + # fp32 params. elif model_param.type() == 'torch.cuda.FloatTensor': shard_model_param = model_param.view(-1)[param_range.start : param_range.end] @@ -433,7 +451,7 @@ def _build_model_and_main_param_groups( shard_float16_groups, shard_fp32_groups, shard_fp32_from_float16_groups, - ) + ), dist_metas def __init__( self, @@ -501,8 +519,22 @@ def __init__( self.optimizer_keys = ("param", "exp_avg_slow", "exp_avg_fast", "exp_avg_sq") else: self.optimizer_keys = ("param", "exp_avg_slow", "exp_avg_sq") + elif isinstance(optimizer, Prodigy): + self.optimizer_name = 'prodigy' + self.optimizer_keys = ("param", "exp_avg", "exp_avg_sq", "s", "p0") + elif isinstance(optimizer, MARS): + self.optimizer_name = 'mars' + if config.mars_amsgrad: + self.optimizer_keys = ("param", "exp_avg", "exp_avg_sq", "last_grad", "max_exp_avg_sq") + else: + self.optimizer_keys = ("param", "exp_avg", "exp_avg_sq", "last_grad") + elif isinstance(optimizer, ADOPT): + self.optimizer_name = 'adopt' + self.optimizer_keys = ("param", "exp_avg", "exp_avg_sq") + elif isinstance(optimizer, Muon): + self.optimizer_name = 'muon' else: - raise Exception(f"Unrecognized optimizer {type(optimizer)}, only Adam and AdEMAMix are supported for now.") + raise Exception(f"Unrecognized optimizer {type(optimizer)}.") # when freezing sub-models we have no real optimizer # but still need a stub DistributedOptimizer class @@ -577,7 +609,7 @@ def __init__( self.shard_float16_groups, self.shard_fp32_groups, self.shard_fp32_from_float16_groups, - ) = self._build_model_and_main_param_groups( + ), dist_metas = self._build_model_and_main_param_groups( self.gbuf_ranges, self.model_param_gbuf_map, self.opt_group_ranges, config ) @@ -587,6 +619,17 @@ def __init__( self.optimizer.param_groups = [g["orig_group"] for g in self.opt_group_ranges] self.optimizer.load_state_dict(self.optimizer.state_dict()) + if isinstance(self.optimizer, Muon): + assert all(grad_buffer.grad_dtype == torch.float32 for grad_buffer in self.buffers), \ + "all grad buffer should only contains float32 type for muon optimizer" + gbuf_sizes = [ [(bucket.grad_data.numel(), bucket.offset) for bucket in buffer.buckets ] + for buffer in self.buffers ] + self.optimizer.enable_distributed_mode( + gbuf_sizes, self.data_parallel_group, + get_tensor_model_parallel_group(), + dist_metas, + ) + self.is_stub_optimizer = False def _get_model_param_range_map(self, param: torch.nn.Parameter): @@ -714,6 +757,20 @@ def load_state_dict(self, state_dict): tensors = {"exp_avg_slow": init_shard(), "exp_avg_fast": init_shard(), "exp_avg_sq": init_shard()} else: # beta1 == 0 tensors = {"exp_avg_slow": init_shard(), "exp_avg_sq": init_shard()} + elif self.optimizer_name == 'prodigy': + tensors = {"exp_avg": init_shard(), "exp_avg_sq": init_shard(), "s": init_shard(), "p0": init_shard()} + elif self.optimizer_name == 'mars': + if len(self.optimizer_keys) == 5: + tensors = {"exp_avg": init_shard(), "exp_avg_sq": init_shard(), "last_grad": init_shard(), "max_exp_avg_sq": init_shard()} + else: + tensors = {"exp_avg": init_shard(), "exp_avg_sq": init_shard(), "last_grad": init_shard()} + elif self.optimizer_name == 'adopt': + tensors = {"exp_avg": init_shard(), "exp_avg_sq": init_shard()} + elif self.optimizer_name == 'muon': + tensors = {"exp_avg": init_shard(), "exp_avg_sq": init_shard()} + tensors["muon_buffer"] = tensors["exp_avg"] + tensors["adamw_exp_avg"] = tensors["exp_avg"] + tensors["adamw_exp_avg_sq"] = tensors["exp_avg_sq"] if self.config.use_precision_aware_optimizer: tensors["master_param"] = init_shard() state_dict_state.append((state_order, tensors)) @@ -793,6 +850,16 @@ def _get_main_param_and_optimizer_states(self, model_param): main_param = self.optimizer.param_groups[group_index]["params"][group_order] optim_state = self.optimizer.state[main_param] tensors = {"param": main_param, **optim_state} + + # process muon to be compatiable with adam ( always save to exp_avg / exp_avg_sq ) + if isinstance(self.optimizer, Muon): + use_muon = self.optimizer.param_groups[group_index].get("use_muon", False) + if use_muon: + tensors["exp_avg"] = tensors["muon_buffer"] + tensors["exp_avg_sq"] = torch.zeros_like(tensors["param"]) + else: + tensors["exp_avg"] = tensors["adamw_exp_avg"] + tensors["exp_avg_sq"] = tensors["adamw_exp_avg_sq"] return tensors def _set_main_param_and_optimizer_states(self, model_param, tensors): @@ -818,6 +885,8 @@ def _set_main_param_and_optimizer_states(self, model_param, tensors): optim_state = self.optimizer.state[main_param] dst_tensors = {"param": main_param, **optim_state} for key in dst_tensors: + if not key in tensors: + continue dst_tensors[key].copy_(tensors[key]) def get_parameter_state_fs_bucket_space(self): diff --git a/megatron/core/optimizer/mars.py b/megatron/core/optimizer/mars.py new file mode 100644 index 00000000000..8adb528524e --- /dev/null +++ b/megatron/core/optimizer/mars.py @@ -0,0 +1,316 @@ +""" +Here is an original implementation of MARS. +Source: https://github.com/AGI-Arena/MARS +""" + +# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates +# SPDX-License-Identifier: Apache-2.0 +import math + +import torch + +@torch.compile +def zeropower_via_newtonschulz5(G, steps=10, eps=1e-7): + """ + Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a + quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose + of minimizing steps, it turns out to be empirically effective to keep increasing the slope at + zero even beyond the point where the iteration no longer converges all the way to one everywhere + on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T + where S' is diagonal with S_{ii}' \sim Uniform(0.5, 1.5), which turns out not to hurt model + performance at all relative to UV^T, where USV^T = G is the SVD. + """ + assert len(G.shape) == 2 + a, b, c = (3.4445, -4.7750, 2.0315) + X = G.bfloat16() + X /= X.norm() + eps # ensure top singular value <= 1 + if G.size(0) > G.size(1): + X = X.T + for _ in range(steps): + A = X @ X.T + B = b * A + c * A @ A + X = a * X + B @ X + if G.size(0) > G.size(1): + X = X.T + return X + + +def exists(val): + return val is not None + + +def update_fn( + p, + grad, + exp_avg, + exp_avg_sq, + lr, + wd, + beta1, + beta2, + last_grad, + eps, + amsgrad, + max_exp_avg_sq, + step, + gamma, + mars_type, + is_grad_2d, + optimize_1d, + lr_1d_factor, + betas_1d, + weight_decay_1d, +): + # optimize_1d: use MARS for 1d para, not: use AdamW for 1d para + if optimize_1d or is_grad_2d: + c_t = (grad - last_grad).mul(gamma * (beta1 / (1.0 - beta1))).add(grad) + c_t_norm = torch.norm(c_t) + if c_t_norm > 1.0: + c_t = c_t / c_t_norm + exp_avg.mul_(beta1).add_(c_t, alpha=1.0 - beta1) + if (mars_type == "mars-adamw") or ( + mars_type == "mars-shampoo" and not is_grad_2d + ): + exp_avg_sq.mul_(beta2).addcmul_(c_t, c_t, value=1.0 - beta2) + bias_correction1 = 1.0 - beta1**step + bias_correction2 = 1.0 - beta2**step + if amsgrad: + torch.max(max_exp_avg_sq, exp_avg_sq, out=max_exp_avg_sq) + denom = ( + max_exp_avg_sq.sqrt() + .mul(1 / math.sqrt(bias_correction2)) + .add(eps) + .mul(bias_correction1) + ) + else: + denom = ( + exp_avg_sq.sqrt() + .mul(1 / math.sqrt(bias_correction2)) + .add(eps) + .mul(bias_correction1) + ) + real_update_tmp = -lr * torch.mul(p.data, wd).add(exp_avg.div(denom)) + elif mars_type == "mars-lion": + real_update_tmp = -lr * torch.mul(p.data, wd).add(exp_avg.sign()) + elif mars_type == "mars-shampoo" and is_grad_2d: + factor = max(1, grad.size(0) / grad.size(1)) ** 0.5 + real_update_tmp = ( + zeropower_via_newtonschulz5(exp_avg.mul(1.0 / (1.0 - beta1)), eps=eps) + .mul(factor) + .add(wd, p.data) + .mul(-lr) + ) + p.data.add_(real_update_tmp) + else: + beta1_1d, beta2_1d = betas_1d + exp_avg.mul_(beta1_1d).add_(grad, alpha=1 - beta1_1d) + exp_avg_sq.mul_(beta2_1d).addcmul_(grad, grad, value=1 - beta2_1d) + bias_correction1 = 1.0 - beta1_1d**step + bias_correction2 = 1.0 - beta2_1d**step + if amsgrad: + torch.max(max_exp_avg_sq, exp_avg_sq, out=max_exp_avg_sq) + denom = ( + max_exp_avg_sq.sqrt() + .mul(1 / math.sqrt(bias_correction2)) + .add(eps) + .mul(bias_correction1) + ) + else: + denom = ( + exp_avg_sq.sqrt() + .mul(1 / math.sqrt(bias_correction2)) + .add(eps) + .mul(bias_correction1) + ) + real_update_tmp = ( + -lr + * lr_1d_factor + * torch.mul(p.data, weight_decay_1d).add(exp_avg.div(denom)) + ) + p.data.add_(real_update_tmp) + return exp_avg, exp_avg_sq + + +class MARS(torch.optim.Optimizer): + def __init__( + self, + params, + lr=3e-3, + betas=(0.95, 0.99), + eps=1e-8, + weight_decay=0.0, + amsgrad=False, + gamma=0.025, + is_approx=True, + mars_type="mars-adamw", + optimize_1d=False, + lr_1d=3e-3, + betas_1d=(0.9, 0.95), + weight_decay_1d=0.1, + ): + if not 0.0 <= lr: + raise ValueError("Invalid learning rate: {}".format(lr)) + if not 0.0 <= eps: + raise ValueError("Invalid epsilon value: {}".format(eps)) + if not 0.0 <= betas[0] < 1.0: + raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) + if not 0.0 <= betas[1] < 1.0: + raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) + assert mars_type in [ + "mars-adamw", + "mars-lion", + "mars-shampoo", + ], "MARS type not supported" + defaults = dict( + lr=lr, + betas=betas, + eps=eps, + weight_decay=weight_decay, + amsgrad=amsgrad, + mars_type=mars_type, + gamma=gamma, + optimize_1d=optimize_1d, + weight_decay_1d=weight_decay_1d, + ) + super(MARS, self).__init__(params, defaults) + self.eps = eps + self.update_fn = update_fn + self.lr = lr + self.weight_decay = weight_decay + self.amsgrad = amsgrad + self.step_num = 0 + self.is_approx = is_approx + self.gamma = gamma + self.mars_type = mars_type + self.optimize_1d = optimize_1d + self.lr_1d_factor = lr_1d / lr + self.weight_decay_1d = weight_decay_1d + self.betas_1d = betas_1d + + @torch.no_grad() + def update_last_grad(self): + if not self.is_approx: + for group in self.param_groups: + for p in group["params"]: + state = self.state[p] + if "last_grad" not in state: + state["last_grad"] = torch.zeros_like(p) + state["last_grad"].zero_().add_(state["previous_grad"], alpha=1.0) + + @torch.no_grad() + def update_previous_grad(self): + if not self.is_approx: + for group in self.param_groups: + for p in group["params"]: + if p.grad is None: + print(p, "grad is none") + continue + state = self.state[p] + if "previous_grad" not in state: + state["previous_grad"] = torch.zeros_like(p) + state["previous_grad"].zero_().add_(p.grad, alpha=1.0) + + def __setstate__(self, state): + super(MARS, self).__setstate__(state) + for group in self.param_groups: + group.setdefault("amsgrad", False) + + @torch.no_grad() + def step( + self, + closure=None, + grads=None, + output_params=None, + scale=None, + grad_norms=None, + grad_scaler=None, + ): + """Performs a single optimization step. + + Arguments: + closure (callable, optional): A closure that reevaluates the model + and returns the loss. + """ + if any(p is not None for p in [grads, output_params, scale, grad_norms]): + raise RuntimeError( + "FusedAdam has been updated. Simply initialize it identically to torch.optim.Adam, and call step() with no arguments." + ) + + loss = None + if exists(closure): + with torch.enable_grad(): + loss = closure() + gamma = self.gamma + for group in self.param_groups: + for p in filter(lambda p: exists(p.grad), group["params"]): + if p.grad is None: + continue + grad = p.grad.data + if grad.is_sparse: + raise RuntimeError( + "Adam does not support sparse gradients, please consider SparseAdam instead" + ) + amsgrad = group["amsgrad"] + + state = self.state[p] + # ('----- starting a parameter state', state.keys(), 'Length of state', len(state)) + # State initialization + if len(state) <= 1: + state["step"] = 0 + # Exponential moving average of gradient values + state["exp_avg"] = torch.zeros_like(p.data) + # Last Gradient + state["last_grad"] = torch.zeros_like(p) + # state['previous_grad'] = torch.zeros_like(p) + # Exponential moving average of squared gradient values + state["exp_avg_sq"] = torch.zeros_like(p.data) + if amsgrad: + # Maintains max of all exp. moving avg. of sq. grad. values + state["max_exp_avg_sq"] = torch.zeros_like(p.data) + exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"] + last_grad = state["last_grad"] + lr, wd, beta1, beta2 = ( + group["lr"], + group["weight_decay"], + *group["betas"], + ) + if amsgrad: + max_exp_avg_sq = state["max_exp_avg_sq"] + else: + max_exp_avg_sq = 0 + + if "step" in state: + state["step"] += 1 + else: + state["step"] = 1 + step = state["step"] + is_grad_2d = len(grad.shape) == 2 + exp_avg, exp_avg_sq = self.update_fn( + p, + grad, + exp_avg, + exp_avg_sq, + lr, + wd, + beta1, + beta2, + last_grad, + self.eps, + amsgrad, + max_exp_avg_sq, + step, + gamma, + mars_type=self.mars_type, + is_grad_2d=is_grad_2d, + optimize_1d=self.optimize_1d, + lr_1d_factor=self.lr_1d_factor, + betas_1d=self.betas_1d, + weight_decay_1d=( + self.weight_decay if self.optimize_1d else self.weight_decay_1d + ), + ) + if self.is_approx: + state["last_grad"] = grad + self.step_num = step + + return loss \ No newline at end of file diff --git a/megatron/core/optimizer/muon.py b/megatron/core/optimizer/muon.py new file mode 100644 index 00000000000..9024fce46ba --- /dev/null +++ b/megatron/core/optimizer/muon.py @@ -0,0 +1,316 @@ + +from typing import Tuple, Dict + +import torch +import math +import torch.distributed as dist + + +# copy from https://github.com/KellerJordan/Muon/tree/master +# @torch.compile +def zeropower_via_newtonschulz5(G, steps): + """ + Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a + quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose + of minimizing steps, it turns out to be empirically effective to keep increasing the slope at + zero even beyond the point where the iteration no longer converges all the way to one everywhere + on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T + where S' is diagonal with S_{ii}' ~ Uniform(0.5, 1.5), which turns out not to hurt model + performance at all relative to UV^T, where USV^T = G is the SVD. + """ + assert len(G.shape) == 2 + a, b, c = (3.4445, -4.7750, 2.0315) + X = G + if G.size(0) > G.size(1): + X = X.T + + # Ensure spectral norm is at most 1 + X = X / (X.norm() + 1e-7) + # Perform the NS iterations + for _ in range(steps): + A = X @ X.T + B = b * A + c * A @ A # adapted from suggestion by @jxbz, @leloykun, and @YouJiacheng + X = a * X + B @ X + + if G.size(0) > G.size(1): + X = X.T + return X + +def normalize_range(range: Tuple[int, int], start): + return (range[0] - start, range[1] - start) + +class MuonDistMeta: + + # which buffer and bucket param belongs to + buffer_idx: int = 0 + bucket_idx: int = 0 + # param shape after tp + shape: torch.Size = None + # param location in global buffer + global_range: Tuple[int, int] = None + tp_split_dim: int = -1 + # param location in global buffer (current dp slice) + local_range: Tuple[int, int] = None + + def __init__(self, buffer_idx: int, bucket_idx: int, shape: torch.Size, global_range: Tuple[int, int], tp_split_dim: int): + self.buffer_idx = buffer_idx + self.bucket_idx = bucket_idx + self.shape = shape + self.global_range = global_range + self.tp_split_dim = tp_split_dim + + def set_local_buffer_range(self, local_buffer_range: Tuple[int, int]): + start = max(self.global_range[0], local_buffer_range[0]) + end = min(self.global_range[1], local_buffer_range[1]) + self.local_range = (start, end) if start < end else (local_buffer_range[0], local_buffer_range[0]) + +# adjust LR based on: https://github.com/MoonshotAI/Moonlight +def adjust_lr_wd_for_muon(lr, matched_adamw_rms, param_shape): + A, B = param_shape[:2] + adjusted_ratio = math.sqrt(max(A, B)) * matched_adamw_rms + adjusted_lr = lr * adjusted_ratio + return adjusted_lr + +# copy from https://github.com/KellerJordan/Muon/tree/master and support distributed solution +class Muon(torch.optim.Optimizer): + """ + Muon - MomentUm Orthogonalized by Newton-schulz + + Muon internally runs standard SGD-momentum, and then performs an orthogonalization post- + processing step, in which each 2D parameter's update is replaced with the nearest orthogonal + matrix. To efficiently orthogonalize each update, we use a Newton-Schulz iteration, which has + the advantage that it can be stably run in bfloat16 on the GPU. + + Some warnings: + - We believe this optimizer is unlikely to work well for training with small batch size. + - We believe it may not work well for finetuning pretrained models, but we haven't tested this. + + Arguments: + param_groups: The parameters to be optimized. + lr: The learning rate. The updates will have spectral norm of `lr`. (0.02 is a good default) + momentum: The momentum used by the internal SGD. (0.95 is a good default) + matched_adamw_rms: The AdamW Update RMS that Muon is designed to match. (0.2~0.4 recommended) + nesterov: Whether to use Nesterov-style momentum in the internal SGD. (recommended) + ns_steps: The number of Newton-Schulz iterations to run. (5 is probably always enough) + {0, 1}-D or are detected as being the embed or lm_head will be optimized by AdamW as well. + adamw_betas: The betas for the internal AdamW. + adamw_eps: The epsilon for the internal AdamW. + adamw_wd: The weight decay for the internal AdamW. + """ + def __init__(self, param_groups, lr=2e-2, weight_decay=0.1, + matched_adamw_rms=0.2, momentum=0.95, nesterov=True, ns_steps=5, + adamw_betas=(0.95, 0.95), adamw_eps=1e-8): + + defaults = dict(lr=lr, weight_decay=weight_decay, + matched_adamw_rms=matched_adamw_rms, + momentum=momentum, nesterov=nesterov, ns_steps=ns_steps, + adamw_betas=adamw_betas, adamw_eps=adamw_eps,) + + super().__init__(param_groups, defaults) + self.distributed_mode = False + + + def enable_distributed_mode(self, global_buffer_sizes, dist_group, tp_group, + dist_metas: Dict[torch.nn.Parameter, MuonDistMeta]): + """ + enable distributed mode + Args: + global_buffer_size: global buffer size + dist group: optimizer sharding group + tp group: param tp group + dist metas: dist metas for all param + """ + + self.global_buffer_sizes = global_buffer_sizes + self.dist_group = dist_group + self.tp_group = tp_group + self.dist_metas = dist_metas + + world_size = dist.get_world_size(dist_group) + rank = dist.get_rank(dist_group) + + # calc local buffer range + self.local_buffer_sizes = [] + self.local_buffer_ranges = [] + for bucket_sizes in global_buffer_sizes: + local_bucket_sizes = [] + local_bucket_ranges = [] + for (global_bucket_size, bucket_offset) in bucket_sizes: + assert global_bucket_size % world_size == 0 + local_buffer_size = global_bucket_size // world_size + local_buffer_start = local_buffer_size * rank + bucket_offset + local_buffer_range = (local_buffer_start, local_buffer_start + local_buffer_size) + local_bucket_sizes.append(local_buffer_size) + local_bucket_ranges.append(local_buffer_range) + + self.local_buffer_sizes.append(local_bucket_sizes) + self.local_buffer_ranges.append(local_bucket_ranges) + + # calc local range for params + for dist_meta in dist_metas.values(): + local_buffer_range = self.local_buffer_ranges[dist_meta.buffer_idx][dist_meta.bucket_idx] + dist_meta.set_local_buffer_range(local_buffer_range) + + self.distributed_mode = True + + def step(self): + + dtype = torch.bfloat16 + device = torch.cuda.current_device() + + ns_inputs = {} + + # update muon momentum first + for group in self.param_groups: + + if not group.get("use_muon", False): + continue + + momentum = group['momentum'] + params = group["params"] + + for p in params: + + g = p.grad + assert g is not None + # 1-dim grad for distributed mode + assert self.distributed_mode or g.dim() == 2 + + # prepare muon buffer in state + state = self.state[p] + if not "muon_buffer" in state: + state["muon_buffer"] = torch.zeros_like(g) + buf = state["muon_buffer"] + buf.mul_(momentum).add_(g) + + # save to ns input + g = g.add(buf, alpha=momentum) if group['nesterov'] else buf + ns_inputs[p] = g.bfloat16() + + # rewrite ns_inputs if distributed + if self.distributed_mode: + + # initialize buffers + ns_input_local_buffers = [ + [ torch.empty((local_buffer_size), device=device, dtype=dtype) + for local_buffer_size in local_bucket_sizes ] + for local_bucket_sizes in self.local_buffer_sizes + ] + ns_input_global_buffers = [ + [ torch.empty((global_buffer_size), device=device, dtype=dtype) + for (global_buffer_size, bucket_offset) in global_bucket_sizes ] + for global_bucket_sizes in self.global_buffer_sizes + ] + + # fill ns input data to local buffer + for param, ns_input in ns_inputs.items(): + dist_meta = self.dist_metas[param] + ns_input_local_buffer = ns_input_local_buffers[dist_meta.buffer_idx][dist_meta.bucket_idx] + local_buffer_range = self.local_buffer_ranges[dist_meta.buffer_idx][dist_meta.bucket_idx] + local_range = normalize_range(dist_meta.local_range, local_buffer_range[0]) + ns_input_local_buffer[local_range[0]:local_range[1]].copy_(ns_input.view(-1)) + + # all gather buffers + for ns_input_global_buffer, ns_input_local_buffer in zip(ns_input_global_buffers, ns_input_local_buffers): + for ns_input_global_bucket, ns_input_local_bucket in zip(ns_input_global_buffer, ns_input_local_buffer): + dist.all_gather_into_tensor(ns_input_global_bucket, ns_input_local_bucket, group=self.dist_group) + + # overwrite ns input + for p in ns_inputs.keys(): + dist_meta = self.dist_metas[p] + ns_input_global_buffer = ns_input_global_buffers[dist_meta.buffer_idx][dist_meta.bucket_idx] + global_range = dist_meta.global_range + offset = self.global_buffer_sizes[dist_meta.buffer_idx][dist_meta.bucket_idx][1] + ns_inputs[p] = ns_input_global_buffer[global_range[0] - offset : global_range[1] - offset].view(dist_meta.shape) + + # set tp info + tp_world_size = dist.get_world_size(self.tp_group) + tp_rank = dist.get_rank(self.tp_group) + + # update muon momentum first + for group in self.param_groups: + + if not group.get('use_muon', False): + continue + + lr = group["lr"] + ns_steps = group["ns_steps"] + weight_decay = group["weight_decay"] + matched_adamw_rms = group["matched_adamw_rms"] + params = group["params"] + + for p in params: + + ns_input = ns_inputs[p] + tp_split_dim = -1 + + if self.distributed_mode: + dist_meta = self.dist_metas[p] + tp_split_dim = dist_meta.tp_split_dim + + # gather tensor parallel ( if tp ) + if tp_split_dim != -1: + ns_input_shards = [ torch.empty_like(ns_input) for _ in range(tp_world_size) ] + dist.all_gather(ns_input_shards, ns_input, self.tp_group) + ns_input = torch.cat(ns_input_shards, dim=tp_split_dim) + + # calc update + update = zeropower_via_newtonschulz5(ns_input, steps=ns_steps) + + # only local tp part + if tp_split_dim != -1: + update = update.chunk(tp_world_size, dim=tp_split_dim)[tp_rank] + + # only local buffer part + if self.distributed_mode: + local_range_in_global_range = normalize_range(dist_meta.local_range, dist_meta.global_range[0]) + update = update.reshape(-1)[local_range_in_global_range[0]:local_range_in_global_range[1]] + + # apply weight decay + p.data.mul_(1 - lr*weight_decay) + + # adjust lr and apply update + adjusted_lr = adjust_lr_wd_for_muon(lr, matched_adamw_rms, ns_input.shape) + p.data.add_(update, alpha=-adjusted_lr) + + # use adam for other params + for group in self.param_groups: + + if group.get('use_muon', False): + continue + + # init step + if 'step' in group: + group['step'] += 1 + else: + group['step'] = 1 + + step = group['step'] + params = group["params"] + lr = group['lr'] + weight_decay = group['weight_decay'] + beta1, beta2 = group['adamw_betas'] + eps = group['adamw_eps'] + + for p in params: + + g = p.grad + assert g is not None + state = self.state[p] + + if len(state) == 0: + state['adamw_exp_avg'] = torch.zeros_like(g) + state['adamw_exp_avg_sq'] = torch.zeros_like(g) + + buf1 = state['adamw_exp_avg'] + buf2 = state['adamw_exp_avg_sq'] + buf1.lerp_(g, 1-beta1) + buf2.lerp_(g.square(), 1-beta2) + + g = buf1 / (eps + buf2.sqrt()) + + bias_correction1 = 1 - beta1**step + bias_correction2 = 1 - beta2**step + scale = bias_correction1 / bias_correction2**0.5 + p.data.mul_(1 - lr * weight_decay) + p.data.add_(g, alpha=-lr/scale) \ No newline at end of file diff --git a/megatron/core/optimizer/optimizer_config.py b/megatron/core/optimizer/optimizer_config.py index 9f09b372266..197396b24dc 100644 --- a/megatron/core/optimizer/optimizer_config.py +++ b/megatron/core/optimizer/optimizer_config.py @@ -101,7 +101,7 @@ class OptimizerConfig: adam_eps: float = 1e-08 """Term added to the denominator to improve numerical stability in Adam optimizer.""" - # SGD. + # SGD sgd_momentum: float = 0.9 """Momentum factor for SGD optimizer.""" @@ -118,6 +118,77 @@ class OptimizerConfig: ademamix_alpha_warmup: Optional[int] = None """Number of warmup steps used to increase alpha.""" + # Prodigy + prodigy_beta3: Optional[float] = None + """coefficients for computing the Prodidy stepsize using running averages. + If set to None, uses the value of square root of beta2.""" + + prodigy_decouple: bool = True + """Use AdamW style decoupled weight decay.""" + + prodigy_use_bias_correction: bool = False + """Turn on Adam's bias correction. Off by default.""" + + prodigy_safeguard_warmup: bool = False + """Remove lr from the denominator of D estimate to avoid issues during warm-up stage. Off by default.""" + + prodigy_fsdp_in_use: bool = False + """If you're using sharded parameters, this should be set to True. The optimizer + will attempt to auto-detect this, but if you're using an implementation other + than PyTorch's builtin version, the auto-detection won't work. + """ + + # MARS + mars_beta1: float = 0.95 + """First coefficient for computing running averages of gradient and its square in Adam + optimizer. + """ + + mars_beta2: float = 0.99 + """Second coefficient for computing running averages of gradient and its square in Adam + optimizer. + """ + + mars_type: str = 'mars-adamw' + """Which version of the MARS framework to use.""" + + mars_vr_gamma: float = 0.025 + """The gamma parameter for the variance reduction term in MARS.""" + + mars_is_approx: bool = True + """Whether to use the approximate version of the MARS optimizer.""" + + mars_lr: float = 0.003 + """The learning rate for the MARS optimizer.""" + + mars_amsgrad: bool = False + """Whether to use the AMSGrad variant of the MARS optimizer.""" + + mars_optimize_1d: bool = False + """If set to False, we optimize 1D parameters with AdamW.""" + + mars_weight_decay_1d: float = 0.1 + """The weight decay for 1D parameters in MARS.""" + + # ADOPT + adopt_eps: float = 1e-6 + """Term added to the denominator to improve numerical stability in ADOPT optimizer.""" + + adopt_decouple: bool = True + """Use AdamW style decoupled weight decay.""" + + # Muon + muon_momentum: float = 0.95 + """Momentum factor for Muon optimizer.""" + + muon_nesterov: bool = True + """Whether or not to use Nesterov momentum for Muon.""" + + muon_ns_steps: int = 5 + """The number of Newton-Schulz iterations.""" + + muon_matched_adamw_rms: float = 0.2 + """The adamw update rms that muon is designed to matched, typicially 0.2 ~ 0.4""" ####################### # Distributed optimizer diff --git a/megatron/core/optimizer/prodigy.py b/megatron/core/optimizer/prodigy.py new file mode 100644 index 00000000000..c6bcb38d42d --- /dev/null +++ b/megatron/core/optimizer/prodigy.py @@ -0,0 +1,274 @@ +""" +Here is an original implementation of Prodigy. +Source: https://github.com/konstmish/prodigy +""" + +import math + +import torch +import torch.distributed as dist + + +class Prodigy(torch.optim.Optimizer): + r""" + Implements Adam with Prodigy step-sizes. + Leave LR set to 1 unless you encounter instability. + + Arguments: + params (iterable): + Iterable of parameters to optimize or dicts defining parameter groups. + lr (float): + Learning rate adjustment parameter. Increases or decreases the Prodigy learning rate. + betas (Tuple[float, float], optional): coefficients used for computing + running averages of gradient and its square (default: (0.9, 0.999)) + beta3 (float): + coefficients for computing the Prodidy stepsize using running averages. + If set to None, uses the value of square root of beta2 (default: None). + eps (float): + Term added to the denominator outside of the root operation to improve numerical stability. (default: 1e-8). + weight_decay (float): + Weight decay, i.e. a L2 penalty (default: 0). + decouple (boolean): + Use AdamW style decoupled weight decay + use_bias_correction (boolean): + Turn on Adam's bias correction. Off by default. + safeguard_warmup (boolean): + Remove lr from the denominator of D estimate to avoid issues during warm-up stage. Off by default. + d0 (float): + Initial D estimate for D-adaptation (default 1e-6). Rarely needs changing. + d_coef (float): + Coefficient in the expression for the estimate of d (default 1.0). + Values such as 0.5 and 2.0 typically work as well. + Changing this parameter is the preferred way to tune the method. + growth_rate (float): + prevent the D estimate from growing faster than this multiplicative rate. + Default is inf, for unrestricted. Values like 1.02 give a kind of learning + rate warmup effect. + fsdp_in_use (bool): + If you're using sharded parameters, this should be set to True. The optimizer + will attempt to auto-detect this, but if you're using an implementation other + than PyTorch's builtin version, the auto-detection won't work. + """ + + def __init__( + self, + params, + lr=1.0, + betas=(0.9, 0.999), + beta3=None, + eps=1e-8, + weight_decay=0, + decouple=True, + use_bias_correction=False, + safeguard_warmup=False, + d0=1e-6, + d_coef=1.0, + growth_rate=float("inf"), + fsdp_in_use=False, + ): + if not 0.0 < d0: + raise ValueError("Invalid d0 value: {}".format(d0)) + if not 0.0 < lr: + raise ValueError("Invalid learning rate: {}".format(lr)) + if not 0.0 < eps: + raise ValueError("Invalid epsilon value: {}".format(eps)) + if not 0.0 <= betas[0] < 1.0: + raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) + if not 0.0 <= betas[1] < 1.0: + raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) + + if decouple and weight_decay > 0: + print(f"Using decoupled weight decay") + + defaults = dict( + lr=lr, + betas=betas, + beta3=beta3, + eps=eps, + weight_decay=weight_decay, + d=d0, + d0=d0, + d_max=d0, + d_numerator=0.0, + d_coef=d_coef, + k=0, + growth_rate=growth_rate, + use_bias_correction=use_bias_correction, + decouple=decouple, + safeguard_warmup=safeguard_warmup, + fsdp_in_use=fsdp_in_use, + ) + self.d0 = d0 + super().__init__(params, defaults) + + @property + def supports_memory_efficient_fp16(self): + return False + + @property + def supports_flat_params(self): + return True + + def step(self, closure=None): + """Performs a single optimization step. + + Arguments: + closure (callable, optional): A closure that reevaluates the model + and returns the loss. + """ + loss = None + if closure is not None: + loss = closure() + + d_denom = 0.0 + + group = self.param_groups[0] + use_bias_correction = group["use_bias_correction"] + beta1, beta2 = group["betas"] + beta3 = group["beta3"] + if beta3 is None: + beta3 = math.sqrt(beta2) + k = group["k"] + + d = group["d"] + d_max = group["d_max"] + d_coef = group["d_coef"] + lr = max(group["lr"] for group in self.param_groups) + + if use_bias_correction: + bias_correction = ((1 - beta2 ** (k + 1)) ** 0.5) / (1 - beta1 ** (k + 1)) + else: + bias_correction = 1 + + dlr = d * lr * bias_correction + + growth_rate = group["growth_rate"] + decouple = group["decouple"] + fsdp_in_use = group["fsdp_in_use"] + + d_numerator = group["d_numerator"] + d_numerator *= beta3 + + for group in self.param_groups: + decay = group["weight_decay"] + k = group["k"] + eps = group["eps"] + group_lr = group["lr"] + d0 = group["d0"] + safeguard_warmup = group["safeguard_warmup"] + + if group_lr not in [lr, 0.0]: + raise RuntimeError( + f"Setting different lr values in different parameter groups is only supported for values of 0" + ) + + for p in group["params"]: + if p.grad is None: + continue + if hasattr(p, "_fsdp_flattened"): + fsdp_in_use = True + + grad = p.grad.data + + # Apply weight decay (coupled variant) + if decay != 0 and not decouple: + grad.add_(p.data, alpha=decay) + + state = self.state[p] + + # State initialization + if "step" not in state: + state["step"] = 0 + state["s"] = torch.zeros_like(p.data).detach() + state["p0"] = p.detach().clone() + # Exponential moving average of gradient values + state["exp_avg"] = torch.zeros_like(p.data).detach() + # Exponential moving average of squared gradient values + state["exp_avg_sq"] = torch.zeros_like(p.data).detach() + + exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"] + + s = state["s"] + p0 = state["p0"] + + if group_lr > 0.0: + # we use d / d0 instead of just d to avoid getting values that are too small + d_numerator += ( + (d / d0) + * dlr + * torch.dot(grad.flatten(), (p0.data - p.data).flatten()).item() + ) + + # Adam EMA updates + exp_avg.mul_(beta1).add_(grad, alpha=d * (1 - beta1)) + exp_avg_sq.mul_(beta2).addcmul_( + grad, grad, value=d * d * (1 - beta2) + ) + + if safeguard_warmup: + s.mul_(beta3).add_(grad, alpha=((d / d0) * d)) + else: + s.mul_(beta3).add_(grad, alpha=((d / d0) * dlr)) + d_denom += s.abs().sum().item() + + ###### + + d_hat = d + + # if we have not done any progres, return + # if we have any gradients available, will have d_denom > 0 (unless \|g\|=0) + if d_denom == 0: + return loss + + if lr > 0.0: + if fsdp_in_use: + dist_tensor = torch.zeros(2).cuda() + dist_tensor[0] = d_numerator + dist_tensor[1] = d_denom + dist.all_reduce(dist_tensor, op=dist.ReduceOp.SUM) + global_d_numerator = dist_tensor[0] + global_d_denom = dist_tensor[1] + else: + global_d_numerator = d_numerator + global_d_denom = d_denom + + d_hat = d_coef * global_d_numerator / global_d_denom + if d == group["d0"]: + d = max(d, d_hat) + d_max = max(d_max, d_hat) + d = min(d_max, d * growth_rate) + + for group in self.param_groups: + group["d_numerator"] = global_d_numerator + group["d_denom"] = global_d_denom + group["d"] = d + group["d_max"] = d_max + group["d_hat"] = d_hat + + decay = group["weight_decay"] + k = group["k"] + eps = group["eps"] + + for p in group["params"]: + if p.grad is None: + continue + grad = p.grad.data + + state = self.state[p] + + exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"] + + state["step"] += 1 + + denom = exp_avg_sq.sqrt().add_(d * eps) + + # Apply weight decay (decoupled variant) + if decay != 0 and decouple: + p.data.add_(p.data, alpha=-decay * dlr) + + ### Take step + p.data.addcdiv_(exp_avg, denom, value=-dlr) + + group["k"] = k + 1 + + return loss \ No newline at end of file diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index c6c0928a458..a0c022ebf12 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -1220,6 +1220,8 @@ def _add_logging_args(parser): help='Enable world size logging to tensorboard.') group.add_argument('--wandb-project', type=str, default='', help='The wandb project name. Ignore wandb by default.') + group.add_argument('--wandb-entity', type=str, default='', + help='The wandb entity name. Ignore wandb by default.') group.add_argument('--wandb-exp-name', type=str, default='', help='The wandb experiment name.') group.add_argument('--wandb-save-dir', type=str, default='', @@ -1261,9 +1263,53 @@ def _add_regularization_args(parser): help='AdEMAMix warmup period for beta_3') group.add_argument('--ademamix-alpha-warmup', type=int, default=-1, help='AdEMAMix warmup period for aplha') + group.add_argument('--prodigy-beta3', type=float, default=None, + help='If set to None, uses the value of square root of beta2') + group.add_argument('--prodigy-decouple', type=bool, default=True, + help='Decoupled weight decay') + group.add_argument('--prodigy-use-bias-correction', type=bool, default=False, + help='Use bias correction') + group.add_argument('--prodigy-safeguard-warmup', type=bool, default=False, + help='Remove lr from the denominator of D estimate to avoid issues during warm-up stage') + group.add_argument('--prodigy-fsdp-in-use', type=bool, default=False, + help='If set, use FSDP') + group.add_argument('--mars-beta1', type=float, default=0.95, + help='First coefficient for computing running averages ' + 'of gradient and its square') + group.add_argument('--mars-beta2', type=float, default=0.99, + help='Second coefficient for computing running averages ' + 'of gradient and its square') + group.add_argument('--mars-type', type=str, default='mars-adamw', choices=['mars-adamw', 'mars-lion', 'mars-shampoo'], + help='Which version of the MARS framework to use') + group.add_argument('--mars-vr-gamma', type=float, default=0.025, + help='Variance Reduction scaling factor') + group.add_argument('--mars-is-approx', type=bool, default=True, + help='If set, use the approximate version of MARS') + group.add_argument('--mars-lr', type=float, default=0.003, + help='Learning rate for MARS') + group.add_argument('--mars-amsgrad', type=bool, default=False, + help='If set, use AMSGrad for MARS') + group.add_argument('--mars-optimize-1d', type=bool, default=False, + help='If set to False, we optimize 1D parameters with AdamW') + group.add_argument('--mars-weight-decay-1d', type=float, default=0.1, + help='Weight decay for 1D parameters') + group.add_argument('--adopt-eps', type=float, default=1e-6, + help='Term added to the denominator to improve' + 'numerical stability') + group.add_argument('--adopt-decouple', type=bool, default=True, + help='Decoupled weight decay') group.add_argument('--adam-eps', type=float, default=1e-08, help='Term added to the denominator to improve' 'numerical stability') + group.add_argument('--muon-matched-adamw-rms', type=float, default=0.2, + help="The RMS of the matched AdamW's, typically 0.2 ~ 0.4") + group.add_argument('--muon-momentum', type=float, default=0.95, + help='Momentum beta for muon') + group.add_argument('--muon-ns-steps', type=int, default=5, + help='Number of Newton-Schultz iteartion steps for muon') + group.add_argument('--no-muon-nesterov', action='store_false', + dest='muon_nesterov', default=True, + help='If set, disable Nesterov momentum for muon') group.add_argument('--sgd-momentum', type=float, default=0.9, help='Momentum factor for sgd') return parser @@ -1463,7 +1509,7 @@ def _add_training_args(parser): help='Enable bias only in the QKV linear layers', dest='add_qkv_bias') group.add_argument('--optimizer', type=str, default='adam', - choices=['adam', 'sgd', 'ademamix'], + choices=['adam', 'sgd', 'ademamix', 'prodigy', 'mars', 'adopt', 'muon'], help='Optimizer function') group.add_argument('--dataloader-type', type=str, default=None, choices=['single', 'cyclic', 'external'], diff --git a/megatron/training/global_vars.py b/megatron/training/global_vars.py index 70701341ec4..5fd4b4689f9 100644 --- a/megatron/training/global_vars.py +++ b/megatron/training/global_vars.py @@ -187,6 +187,7 @@ def _set_wandb_writer(args): 'dir': save_dir, 'name': args.wandb_exp_name, 'project': args.wandb_project, + 'entity': args.wandb_entity, 'config': vars(args)} os.makedirs(wandb_kwargs['dir'], exist_ok=True) wandb.init(**wandb_kwargs) diff --git a/megatron/training/training.py b/megatron/training/training.py index caa574494de..c9f183c6340 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -784,7 +784,11 @@ def train_step(forward_step_func, data_iterator, # Set grad to zero. for model_chunk in model: model_chunk.zero_grad_buffer() - optimizer.zero_grad() + if optimizer.__class__.__name__ == "MARS": + optimizer.zero_grad(set_to_none=True) + optimizer.update_last_grad() + else: + optimizer.zero_grad() # Forward pass. forward_backward_func = get_forward_backward_func() diff --git a/tests/unit_tests/test_optimizer_muon.py b/tests/unit_tests/test_optimizer_muon.py new file mode 100644 index 00000000000..3dcfeb59a21 --- /dev/null +++ b/tests/unit_tests/test_optimizer_muon.py @@ -0,0 +1,186 @@ + +import os + +import torch +import torch.distributed as dist + +from megatron.core.optimizer.muon import Muon, MuonDistMeta, normalize_range + +def is_rank_0(): + return torch.distributed.get_rank() == 0 + +def print_rank_0(*args): + if is_rank_0(): + print(*args) + +def cdiv(x: int, y: int): + return (x + y - 1) // y + +def gen_param_and_grads(): + + # reset manual seed + torch.manual_seed(0) + torch.cuda.manual_seed(0) + device = 'cuda' + dtype = torch.float32 + + # gen params + params = [ torch.randn(shape, device=device, dtype=dtype) for shape in [ + (100, 100), (124, 324), (456, 124), (676, 876), (128, 128), ] ] + + # gen grads [ [ grad-list ] * step ] + grads = [ [ torch.randn_like(param) for param in params ] for _ in range(10) ] + + return params, grads + +def distribute_params(params, grads, tp_dims, dist_group, tp_group): + """ 将 param 进行 dist & tp shard, 仅保留自己的一部分 """ + + params = params.copy() + grads = [ step_grads.copy() for step_grads in grads ] + + # tp dist + tp_size = dist.get_world_size(tp_group) + tp_rank = dist.get_rank(tp_group) + for i, param in enumerate(params): + tp_dim = tp_dims[i] + if tp_dim == -1: + continue + assert param.shape[tp_dim] % tp_size == 0 + local_range_start = param.shape[tp_dim] // tp_size * tp_rank + local_range_end = param.shape[tp_dim] // tp_size * (tp_rank + 1) + params[i] = param[local_range_start:local_range_end, :] if tp_dim == 0 else \ + param[:, local_range_start:local_range_end].contiguous() + + for step_grads in grads: + step_grads[i] = step_grads[i][local_range_start:local_range_end, :] if tp_dim == 0 else \ + step_grads[i][:, local_range_start:local_range_end].contiguous() + + # distributed + world_size = dist.get_world_size(dist_group) + rank = dist.get_rank(dist_group) + + global_buffer_size = sum(param.numel() for param in params) + local_buffer_size = cdiv(global_buffer_size, world_size) + local_buffer_range = (local_buffer_size * rank, local_buffer_size * (rank + 1)) + global_buffer_size = local_buffer_size * world_size # fix global buffer size + + numel_acc = 0 + dist_params = [] + dist_grads = [[] for _ in grads] + dist_metas = {} + for i, param in enumerate(params): + + # gen meta + numel = param.numel() + dist_meta = MuonDistMeta(0, 0, param.shape, (numel_acc, numel_acc + numel), tp_dims[i]) + dist_meta.set_local_buffer_range(local_buffer_range) + numel_acc += numel + + # skip if no element in this shard + if dist_meta.local_range[0] == dist_meta.local_range[1]: + continue + + # gen param + local_range = normalize_range(dist_meta.local_range, dist_meta.global_range[0]) + dist_param = param.view(-1)[local_range[0]:local_range[1]] + dist_params.append(dist_param) + dist_metas[dist_param] = dist_meta + + # gen grad + for step, step_grads in enumerate(grads): + dist_grad = step_grads[i].view(-1)[local_range[0]:local_range[1]] + dist_grads[step].append(dist_grad) + + return dist_params, dist_grads, global_buffer_size, dist_metas + + +def test_muon_dist(dp_size, tp_size): + + world_size = dist.get_world_size() + rank = dist.get_rank() + assert dp_size * tp_size == world_size + + # init dist group + for i in range(tp_size): + ranks = range(i, world_size, tp_size) + group = dist.new_group(ranks) + if rank in ranks: + dist_group = group + # init tp group + for i in range(dp_size): + ranks = range(i * tp_size, (i + 1) * tp_size) + group = dist.new_group(ranks) + if rank in ranks: + tp_group = group + + print_rank_0("process group initialized") + + params_ref, grads_ref = gen_param_and_grads() + params_test, grads_test = gen_param_and_grads() + tp_dims = [0, 1, -1, 1, 0] + + params_test, grads_test, global_buffer_size, dist_metas \ + = distribute_params(params_test, grads_test, tp_dims, dist_group, tp_group) + + muon_args = { + "use_muon": True, + "lr": 0.1, + "momentum": 0.9, + "nesterov": True, + "ns_steps": 5, + "weight_decay": 0.1, + } + + # gen params + ref_param_groups = [{ + "params": params_ref, + **muon_args + }] + test_param_groups = [{ + "params": params_test, + **muon_args + }] + + ref_muon = Muon(ref_param_groups) + test_muon = Muon(test_param_groups) + test_muon.enable_distributed_mode([[(global_buffer_size, 0)]], dist_group, tp_group, dist_metas) + + for step in range(10): + + # add grad + for i, grad in enumerate(grads_ref[step]): + params_ref[i].grad = grad.clone() + for i, grad in enumerate(grads_test[step]): + params_test[i].grad = grad.clone() + # step + ref_muon.step() + test_muon.step() + # distribute ref params + dist_ref_params, _, _, _ = distribute_params(params_ref, [], tp_dims, dist_group, tp_group) + # verify + for i, params_x2 in enumerate(zip(dist_ref_params, params_test)): + assert (params_x2[0] == params_x2[1]).all(), f"rank {rank} param {i} verify failed" + print_rank_0(f" - step {step} verify passed") + + print_rank_0(f"dist dp = {dp_size} tp = {tp_size} test passed") + +def run_process(rank, world_size): + + # init dist + torch.cuda.set_device(rank) + dist.init_process_group("nccl", rank=rank, world_size=world_size) + + test_muon_dist(dp_size=4, tp_size=2) + test_muon_dist(dp_size=2, tp_size=4) + + dist.destroy_process_group() + +if __name__ == "__main__": + + world_size = 8 + os.environ['MASTER_ADDR'] = 'localhost' + os.environ['MASTER_PORT'] = '12345' + os.environ['CUDA_DEVICE_MAX_CONNECTIONS'] = '1' + + torch.multiprocessing.spawn(run_process, args=(world_size,), nprocs=world_size, join=True) \ No newline at end of file