CoolFace
Apppublic

procgne/Plonk

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
optimizers.py112 linesDownload Raw Back to utils
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