procgne/Plonk
0
1"""Lamb optimizer."""2 3import torch4from torch.optim import Optimizer5import math6 7 8class Lamb(Optimizer):9 r"""Implements Lamb algorithm.10 It has been proposed in `Large Batch Optimization for Deep Learning: Training BERT in 76 minutes`_.11 Arguments:12 params (iterable): iterable of parameters to optimize or dicts defining13 parameter groups14 lr (float, optional): learning rate (default: 1e-3)15 betas (Tuple[float, float], optional): coefficients used for computing16 running averages of gradient and its square (default: (0.9, 0.999))17 eps (float, optional): term added to the denominator to improve18 numerical stability (default: 1e-8)19 weight_decay (float, optional): weight decay (L2 penalty) (default: 0)20 adam (bool, optional): always use trust ratio = 1, which turns this into21 Adam. Useful for comparison purposes.22 .. _Large Batch Optimization for Deep Learning: Training BERT in 76 minutes:23 https://arxiv.org/abs/1904.0096224 """25 26 def __init__(27 self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0, adam=False28 ):29 if not 0.0 <= lr:30 raise ValueError("Invalid learning rate: {}".format(lr))31 if not 0.0 <= eps:32 raise ValueError("Invalid epsilon value: {}".format(eps))33 if not 0.0 <= betas[0] < 1.0:34 raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0]))35 if not 0.0 <= betas[1] < 1.0:36 raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1]))37 defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay)38 self.adam = adam39 super(Lamb, self).__init__(params, defaults)40 41 def step(self, closure=None):42 """Performs a single optimization step.43 Arguments:44 closure (callable, optional): A closure that reevaluates the model45 and returns the loss.46 """47 loss = None48 if closure is not None:49 loss = closure()50 51 for group in self.param_groups:52 for p in group["params"]:53 if p.grad is None:54 continue55 grad = p.grad.data56 if grad.is_sparse:57 raise RuntimeError(58 "Lamb does not support sparse gradients, consider SparseAdam instad."59 )60 61 state = self.state[p]62 63 # State initialization64 if len(state) == 0:65 state["step"] = 066 # Exponential moving average of gradient values67 state["exp_avg"] = torch.zeros_like(p.data)68 # Exponential moving average of squared gradient values69 state["exp_avg_sq"] = torch.zeros_like(p.data)70 71 exp_avg, exp_avg_sq = state["exp_avg"], state["exp_avg_sq"]72 beta1, beta2 = group["betas"]73 74 state["step"] += 175 76 # Decay the first and second moment running average coefficient77 # m_t78 exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)79 # v_t80 exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)81 82 # Paper v3 does not use debiasing.83 bias_correction1 = 1 - beta1 ** state["step"]84 bias_correction2 = 1 - beta2 ** state["step"]85 exp_avg_hat = exp_avg / bias_correction186 exp_avg_sq_hat = exp_avg_sq / bias_correction287 # Apply bias to lr to avoid broadcast.88 step_size = group["lr"]89 90 do_layer_adaptation = (91 group["layer_adaptation"]92 if "layer_adaptation" in group93 else group["weight_decay"] > 094 )95 96 adam_step = exp_avg_hat / exp_avg_sq_hat.sqrt().add(group["eps"])97 if group["weight_decay"] != 0:98 adam_step.add_(p.data, alpha=group["weight_decay"])99 if do_layer_adaptation:100 weight_norm = p.data.norm(p=2)101 adam_norm = adam_step.norm(p=2)102 trust_ratio = torch.where(103 weight_norm.ne(0),104 torch.where(adam_norm.ne(0), weight_norm / adam_norm, 1),105 1,106 )107 if self.adam or not do_layer_adaptation:108 trust_ratio = 1109 110 p.data.add_(adam_step, alpha=-step_size * trust_ratio)111 return loss112 