diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index f7b8825477..6b70032a13 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -1668,7 +1668,7 @@ def train_llama3(): def train_gptoss(): from examples.mlperf.models.gpt_oss import GPTOSS, GPT_OSS_20B, apply_grad, FP8_DTYPE from examples.mlperf.lr_schedulers import CosineAnnealingLRWithWarmup - from examples.mlperf.optim import GradAccClipAdamW, clip_grads + from examples.mlperf.optim import GradAccClipAdamW, GradAccClipAdamWGroup, clip_grads BENCHMARK = getenv("BENCHMARK") @@ -1734,7 +1734,12 @@ def train_gptoss(): is_offload_optim = bool(getenv("OFFLOAD_OPTIM")) is_fake_offload = Device.DEFAULT == "NULL" optim_device = ("CPU" if not is_fake_offload else "NULL:99") if is_offload_optim else None - optim = GradAccClipAdamW(params, lr=0.0, b1=opt_adamw_beta_1, b2=opt_adamw_beta_2, eps=opt_adamw_epsilon, weight_decay=opt_adamw_weight_decay, grad_acc=grad_acc, device=optim_device) + params_wd = [p for p in params if p.ndim >= 3] + params_no_wd = [p for p in params if p.ndim < 3] + optim = GradAccClipAdamWGroup( + GradAccClipAdamW(params_wd, lr=0.0, b1=opt_adamw_beta_1, b2=opt_adamw_beta_2, eps=opt_adamw_epsilon, weight_decay=opt_adamw_weight_decay, grad_acc=grad_acc, device=optim_device), + GradAccClipAdamW(params_no_wd, lr=0.0, b1=opt_adamw_beta_1, b2=opt_adamw_beta_2, eps=opt_adamw_epsilon, weight_decay=0.0, grad_acc=grad_acc, device=optim_device), + ) for p in optim.params: grad_dtype = dtypes.bfloat16 if p.dtype == FP8_DTYPE else p.dtype diff --git a/examples/mlperf/optim.py b/examples/mlperf/optim.py index 5c6d4c53e8..d9130d7dfa 100644 --- a/examples/mlperf/optim.py +++ b/examples/mlperf/optim.py @@ -1,6 +1,6 @@ from tinygrad.tensor import Tensor from tinygrad.dtype import dtypes -from tinygrad.nn.optim import Optimizer +from tinygrad.nn.optim import Optimizer, OptimizerGroup from tinygrad.helpers import FUSE_OPTIM, getenv from tinygrad.uop.ops import UOp, Ops @@ -121,3 +121,21 @@ class GradAccClipAdamW(Optimizer): return ret.shard_like(t) if offloaded else ret out = new_w.cast(t.dtype) return out.shard_like(t) if offloaded else out + +class GradAccClipAdamWGroup(OptimizerGroup): + def fstep(self, grads:list[Tensor], grad_norm:Tensor|None=None): + offset = 0 + to_realize = [] + for o in self.optimizers: + n = len(o.params) + to_realize += o.fschedule_step(grads[offset:offset+n]) + offset += n + Tensor.realize(*to_realize, *([grad_norm] if grad_norm is not None else [])) + @property + def lr(self): return self.optimizers[0].lr + @property + def device(self): return self.optimizers[0].device + @property + def master_params(self): + mp = [mp for o in self.optimizers for mp in (o.master_params or [])] + return mp if mp else None