From 66ca8c69abfa2f65059df82790db975deb352345 Mon Sep 17 00:00:00 2001 From: Less Wright Date: Mon, 5 Apr 2021 20:23:55 -0700 Subject: [PATCH] stable weight decay --- src/Ranger21.py | 260 ++++++++++++++++++++---- src/__pycache__/Ranger21.cpython-37.pyc | Bin 0 -> 3436 bytes 2 files changed, 220 insertions(+), 40 deletions(-) create mode 100644 src/__pycache__/Ranger21.cpython-37.pyc diff --git a/src/Ranger21.py b/src/Ranger21.py index 5593c40..a30d3c3 100644 --- a/src/Ranger21.py +++ b/src/Ranger21.py @@ -12,7 +12,7 @@ import torch -import torch.optim +import torch.optim as TO import math import collections @@ -20,65 +20,245 @@ import collections import copy from torch import linalg as LA -class Ranger21(torch.optim.Optimizer): - def __init__(self, - params, - lr, - eps=1e-8, - num_batches_per_epoch = None, - num_epochs = None, - num_warmup_iterations = 1000, - weight_decay=0, - decay_type = "stable", - warmup_type = 'linear', - use_GC=True): +class Ranger21(TO.Optimizer): + def __init__( + self, + params, + lr, + betas=(0.9, 0.999), # temp for checking tuned warmups + momentum=0.9, + eps=1e-8, + num_batches_per_epoch=None, + num_epochs=None, + num_warmup_iterations=1000, + weight_decay=0, + decay_type="stable", + warmup_type="linear", + use_GC=True, + ): - # todo - checks on incoming params - defaults = dict(lr=lr, eps=eps, weight_decay = weight_decay) - super().__init__(params, defaults) + # todo - checks on incoming params + defaults = dict( + lr=lr, momentum=momentum, betas=betas, eps=eps, weight_decay=weight_decay + ) + super().__init__(params, defaults) - self.num_batches = num_batches_per_epoch - self.num_epochs = num_epochs - self.num_warmup_iters = num_warmup_iterations - self.warmup_type=warmup_type - self.use_GC = use_GC - self.starting_lr = lr + self.num_batches = num_batches_per_epoch + self.num_epochs = num_epochs + self.num_warmup_iters = num_warmup_iterations + self.warmup_type = warmup_type + self.use_GC = use_GC + self.starting_lr = lr - #decay - self.decay = weight_decay - self.decay_type = decay_type + # decay + self.decay = weight_decay + self.decay_type = decay_type + self.param_size = 0 + + # logging + self.variance_sum_tracking = [] + + def __setstate__(self, state): + super().__setstate__(state) def warmup_dampening(self, step): # not usable yet style = self.warmup_type - step +=1 + step += 1 warmup = self.num_warmup_iters if style is None: return 1.0 - if style=='linear': - return min(1.0, (step/warmup) ) + if style == "linear": + return min(1.0, (step / warmup)) - elif style=='exponential': - return 1.0 - math.exp(-step/warmup) + elif style == "exponential": + return 1.0 - math.exp(-step / warmup) else: raise ValueError(f"warmup type {style} not implemented.") + def get_variance(self): + return self.variance_sum_tracking + def get_state_values(self, group, state): + beta1, beta2 = group["betas"] + mean_avg = state["mean_avg"] + variance_avg = state["variance_avg"] - @torch.no_grad - def step(self, - closure = None, - passed_loss = None): + return beta1, beta2, mean_avg, variance_avg - # let's build in a loss pass through for HyperExplorer - loss = None - if closure is not None and isinstance(closure, collections.Callable): - with torch.grad(): - loss = closure() - + # @staticmethod + @torch.no_grad() + def step(self, closure=None): + loss = None + # if closure is not None and isinstance(closure, collections.Callable): + # with torch.grad(): + # loss = closure() + # if closure is not None: + # with torch.enable_grad(): + # loss = closure() + param_size = 0 + variance_ma_sum = 0.0 + + for i, group in enumerate(self.param_groups): + for j, p in enumerate(group["params"]): + if p.grad is None: + continue + + if not self.param_size: + param_size += p.numel() + + # Perform optimization step + grad = p.grad + + if grad.is_sparse: + raise RuntimeError("sparse matrix not supported atm") + + state = self.state[p] + + # State initialization + if len(state) == 0: + # print("init state") + state["step"] = 0 + # Exponential moving average of gradient values + state["grad_ma"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + # Exponential moving average of squared gradient values + state["variance_ma"] = torch.zeros_like( + p, memory_format=torch.preserve_format + ) + + state["step"] += 1 + + beta1, beta2 = group["betas"] + grad_ma = state["grad_ma"] + variance_ma = state["variance_ma"] + + bias_correction2 = 1 - beta2 ** state["step"] + + # update the exp averages + grad_ma.mul_(beta1).add_(grad, alpha=1 - beta1) + + variance_ma.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) + + variance_ma_debiased = variance_ma / bias_correction2 + + variance_ma_sum += variance_ma_debiased.sum() + + # print(f"variance hat sum = {exp_avg_sq_hat_sum}") + # Calculate the sqrt of the mean of all elements in exp_avg_sq_hat + + # we will run this first epoch only and then memoize + if not self.param_size: + self.param_size = param_size + print(f"params size saved") + print(f"total param groups = {i+1}") + print(f"total params in groups = {j+1}") + + if not self.param_size: + raise ValueError("failed to set param size") + + # debugging + self.variance_sum_tracking.append(variance_ma_sum.item()) + + variance_normalized = math.sqrt(variance_ma_sum / self.param_size) + + # print(f"variance mean sqrt = {variance_normalized}") + + # phase 2 - apply weight decay and step + for group in self.param_groups: + for p in group["params"]: + if p.grad is None: + continue + + state = self.state[p] + + step = state["step"] + + # Perform stable weight decay + decay = group["weight_decay"] + eps = group["eps"] + lr = group["lr"] + + if decay: + p.data.mul_(1 - decay * lr / variance_normalized) + + beta1, beta2 = group["betas"] + grad_exp_avg = state["grad_ma"] + variance_ma = state["variance_ma"] + + bias_correction1 = 1 - beta1 ** step + bias_correction2 = 1 - beta2 ** step + + variance_biased_ma = variance_ma / bias_correction2 + + denom = variance_biased_ma.sqrt().add(eps) + + step_size = lr / bias_correction1 + + # update weights + p.addcdiv_(grad_exp_avg, denom, value=-step_size) + + return loss + + """ param_size = 0 + variance_avg_sum = 0. + + for group in self.param_groups: + for p in group['params']: + if p.grad is None: + continue + param_size += p.numel() + + # first part of optimization + grad = p.grad + + + + state = self.state[p] + + #init if needed + if len(state)==0: + print(f"initing state") + state['step']=0 + + state['mean_avg'] = torch.zeros_like(p, memory_format=torch.preserve_format) + state['variance_avg'] = torch.zeros_like(p,memory_format = torch.preserve_format) + + # get state values + #beta1, beta2, mean_avg, variance_avg = self.get_state_values(group, state) + beta1,beta2 = group['betas'] + mean_avg = state['mean_avg'] + variance_avg = state['variance_avg'] + + #print(f"beta1= {beta1}") + + state['step'] +=1 + + bias_correction2 = 1 - beta2**variance_avg + + # bias avgs + mean_avg.mul_(beta1).add_(grad, alpha=1-beta1) + variance_avg.mul_(beta2).addcmul_(grad,grad,value = 1-beta2) + + variance_avg_hat = variance_avg / bias_correction2 + #print(f"variance-avg-hat = {variance_avg_hat}") + + variance_avg_sum += variance_avg_hat.sum() + + print(f"param size = {param_size}") + if not self.paramsize: + self.paramsize = param_size + else: + if self.paramsize != param_size: + raise ValueError("param size changed") + print(f"variance_avg_sum = {variance_avg_sum}") + variance_avg_normalized = math.sqrt(variance_avg_sum / param_size) + print(f"variance sum normalize = {variance_avg_normalized}") +""" diff --git a/src/__pycache__/Ranger21.cpython-37.pyc b/src/__pycache__/Ranger21.cpython-37.pyc new file mode 100644 index 0000000000000000000000000000000000000000..788bebd4016b423fe466cdf7969cf89143dc191b GIT binary patch literal 3436 zcmZuzO>7&-6`t8YF3BZDNt8^9 zi1xBPewaujvW(Teeq%0h6LkdzGglzy+*u>%$L!Q9c+ohsI4it@ zMRrjEbqXKUE$X0N5r9^(mPLM11+5k}(3&7leZW8i!Eg_IyXR^vQQ9A+BRS5eBkewv zd7NomPBI;QEQf>LJQA`W9~$9%4ef|wKi6(Hok*;UqTzU$M^RWGPe;+iIPdSuOxLI* zCux5-(=BR$9IMfE5)E^yGF_kd^1}&Gm}WA1@13v?R8&434a~Md4dZcNMj7_bRos7sg}tVBGMVgX&rHn}y}t^G)QJ$x;o1c;Bix{rQmrh; zE=U^Z!*rY}dKPtsNSnwSi1vltiKj`PLG`Lp)KB6pi=r>se_r0&x>o+r735l``VY=K zTa!bzj14*{84dr7R|>;HR*`alNK5|gI4f@pK-D#yer@C zS&Bw<5Jj2Hfo(3M2+U1lx~AXb>TBq!mxz3wNEhTSl+2vOkJc{8@M4#R9Aw!p$~&}V z2c%$oYzFNV{7p9F=<%6dSb{zBRaaPdq0Pd6WU24v7BrC;_EmOVavgB*2rIWwor2dH zjhyigZj!sgEv(O&@IJSX`Ojg&xLsxYIT_Lb1>e7JsVXX6myaju7)m^hlcS(~#xAK& z_h`90PV???G)W|BxfEMHuYAl^T4|4lW9^LMd>3l}c%o}RjgzT-SE*F>oDwl7%jJY{ zJ4|Hgm7ABqEzr}gx$1~`G?8N{kk-x;{NDedG)X zWCQRw(I99y*aK^3&m4ehZLk^O5jJWMZo+}1@OB(Pm3^1@XIGLK57Iv}jFJD8ODAdCmNs(Ba!}D)4n(eX;CaQ13B+>Uo%{)?l7SJl}-#|_0G({D@8S^*T zDty3fsc4)AMg5HLHQ?c^MZLpB{TX@oCinsPP55sb2`~~oL$oZK#nM$aYvHk5r%UE; z&i)(s2`T%8X9u`5VnW^$jbnHN)Xl;b)hp~X_PI6ZEEScDocYR6@@3JQqq1_kS~%oQ z^7iQ(PQu<~k+bFf&HL=}U(m;!eE6rC_xQ{5CD8lrix)WI>JKM;|2LR7$G5!mYz4EI zfx{~571V9itEktC)e}o~@J8#W0dBimthHIOd=(gx6qt1x(e(?6){7N`<`si(*O z7v?yM)hn25>=g^#IGOCm)G(HHv|I+WE(N}B79U7)w3erNoOI0qPK5Dh-M6|&?K!(U z9ADsecH&_oMK@2oaN+Ysv{@O~C@7jh+1t=niQp&|LWK^@=25x1c1UVP;YzzSCDGMk z7L`}l!JR35?x^&ar5CXiIo58Ts{XF7;iNQ+lHnt%mnKSPQXR;1IGoWmiL?_75osT^ zZ~8W(x^^dOIL_6pB;N*|?8OtfRiVB?9VgpYIZ}~0kF|}329fWA=+?twoJIXqDcPsQ;6~WK zxSfzRE=9Py;Ew6d6F@DDHmG&lYXNu|7+pZXG5-1T_SeIe#YW}Mq&7&5#_32`DIA$x z=#pnP2ZcB%n}9IesXK@{9lptH+(HxzOnTwlH9N3;tBqgFG7j5r*&S;Wyq1SZ1sA@; zR}iOs%xn7rZ^ErRR?YHl@ZjD(-YM6C2DWb@rr9r9M~y{HBLRCgMbS%9j-Wh{?QKQ= z-q@9*XkjAvJxW*VluC}{5wa9ri=t5~rU|u?zwA%rq#QA+lgU!lw`le2L?~g=m2nDX zg_4rl>Bc-G>&qmGq8AXIXFMj^p&(0HT}IaDJB_*86Falsb}Wr=KF_qkKwR@88zd4`>NlpV{w)Xl088)Ber! S_HR?-nNj}gbG#N%r~1F(k7`f= literal 0 HcmV?d00001