From 04274ee7d7cfaaed21a82671d671511d782b4704 Mon Sep 17 00:00:00 2001 From: Less Wright Date: Tue, 6 Apr 2021 20:35:33 -0700 Subject: [PATCH] warmup implemented --- src/Ranger21.py | 74 +++++++++++++++++++----- src/__pycache__/Ranger21.cpython-37.pyc | Bin 3436 -> 4547 bytes 2 files changed, 61 insertions(+), 13 deletions(-) diff --git a/src/Ranger21.py b/src/Ranger21.py index 038ad67..ac9cae7 100644 --- a/src/Ranger21.py +++ b/src/Ranger21.py @@ -20,6 +20,19 @@ import collections import copy from torch import linalg as LA +def centralize_gradient(x, gc_conv_only=False): + """credit - https://github.com/Yonghongwei/Gradient-Centralization """ + + size = len(list(x.size())) + #print(f"size = {size}") + + if gc_conv_only: + if size > 3: + x.add_(-x.mean(dim=tuple(range(1, size)), keepdim=True)) + else: + if size > 1: + x.add_(-x.mean(dim=tuple(range(1, size)), keepdim=True)) + return x class Ranger21(TO.Optimizer): def __init__( @@ -31,11 +44,13 @@ class Ranger21(TO.Optimizer): eps=1e-8, num_batches_per_epoch=None, num_epochs=None, - num_warmup_iterations=1000, - weight_decay=0, + use_warmup = True, + num_warmup_iterations=None, + weight_decay=1e-4, decay_type="stable", warmup_type="linear", - use_GC=True, + use_gradient_centralization=True, + gc_conv_only=False ): # todo - checks on incoming params @@ -46,9 +61,10 @@ class Ranger21(TO.Optimizer): 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.use_gc = use_gradient_centralization, + self.gc_conv_only=gc_conv_only, self.starting_lr = lr # decay @@ -56,26 +72,46 @@ class Ranger21(TO.Optimizer): self.decay_type = decay_type self.param_size = 0 + # warmup - we'll use default recommended in Ma/Yarats unless user specifies num iterations + self.use_warmup = use_warmup + if num_warmup_iterations is None: + self.num_warmup_iters = math.ceil((2 / (1-betas[1]))) # default untuned linear warmup + else: + self.num_warmup_iters = num_warmup_iterations + # logging self.variance_sum_tracking = [] + # print out initial settings to make usage easier + print(f"Ranger21 optimizer ready with following settings:\n") + print(f"Learning rate of {self.starting_lr}") + if self.use_warmup: + print(f"{self.warmup_type} warmup, over {self.num_warmup_iters} iterations") + + print(f"Stable weight decay of {self.decay}") + if self.use_gc: + print(f"Gradient Centralization = On") + else: + print("Gradient Centralization = Off") + + def __setstate__(self, state): super().__setstate__(state) - def warmup_dampening(self, step): + + def warmup_dampening(self, lr, step): # not usable yet style = self.warmup_type - step += 1 warmup = self.num_warmup_iters if style is None: return 1.0 if style == "linear": - return min(1.0, (step / warmup)) + return lr * min(1.0, (step / warmup)) elif style == "exponential": - return 1.0 - math.exp(-step / warmup) + return lr * (1.0 - math.exp(-step / warmup)) else: raise ValueError(f"warmup type {style} not implemented.") @@ -99,13 +135,12 @@ class Ranger21(TO.Optimizer): with torch.grad(): loss = closure() - # if closure is not None: - # with torch.enable_grad(): - # loss = closure() param_size = 0 variance_ma_sum = 0.0 + #phase 1 - accumulate all of the variance_ma_sum to use in stable weight decay + for i, group in enumerate(self.param_groups): for j, p in enumerate(group["params"]): if p.grad is None: @@ -114,7 +149,7 @@ class Ranger21(TO.Optimizer): if not self.param_size: param_size += p.numel() - # Perform optimization step + grad = p.grad if grad.is_sparse: @@ -135,6 +170,15 @@ class Ranger21(TO.Optimizer): p, memory_format=torch.preserve_format ) + # centralize gradients + if self.use_gc: + grad = centralize_gradient( + grad, + gc_conv_only=self.gc_conv_only, + ) + #else: + # grad = uncentralized_grad + state["step"] += 1 beta1, beta2 = group["betas"] @@ -186,6 +230,10 @@ class Ranger21(TO.Optimizer): decay = group["weight_decay"] eps = group["eps"] lr = group["lr"] + if self.use_warmup: + lr = self.warmup_dampening(lr,step) + #if step < 10: + # print(f"warmup dampening at step {step} = {lr} vs {group['lr']}") if decay: p.data.mul_(1 - decay * lr / variance_normalized) diff --git a/src/__pycache__/Ranger21.cpython-37.pyc b/src/__pycache__/Ranger21.cpython-37.pyc index 788bebd4016b423fe466cdf7969cf89143dc191b..f3141831ed5d6b8e177e568fde7862de1e5f0c04 100644 GIT binary patch literal 4547 zcma)9&2JmW6`$EJE|(NVeN&Q?Y>YNAn^=x@XxbpLlQ?mLHjoj=s8b=qcEwpyD=$Ah zyOb?rd1z!bKvB282MMHOE;;uP=%I(U#{!*W8Z>ZQpg9C+(NlkKmQ)Bg1xjh&o1HhG zZ{EjmUYVKk7=CBZb$+twGxlftnEWhMUP4LV0TE2_gl!j`cZ_Y57jS1;O4}~_xWW?lJ-%Ijoe4*j?lDn{%zf|3+^z`zHfxq2Vs^IGpLJGCOmGksC+G5f`Sg?i&?$ z!a3`nLr>%X>6tTZ%O1Yel97lr|GdAKW&QN|3m4jPwl~<>X!SZ5e%$M}_fYOc@rBoA zDB`G_oxcJi!z3PtS={UT4{84oImnHsp-d5XROwa}^-*b-l$k_bWhHT%DJzYKk+MP| z0%digu&eBB&`%;|%dp#ynx^7+Rkhs;TD|T;(Ca3LvW!_jVX_YP? zZ9^=7z;?OtMCD$2+tl@{n7QY|8fs$p9)mU1#T>XJ=EVYdDRPg?Crqq}MX_`b*SGC` zz$~_1fwCLQo$U4#qy78e(eLk<+fRS6{HH(s^Bs(Y{rjZPC3tK zPYlj-Cuf0~mq4w&3To#wyD+%G$!nUIFrEc&?tzx`3aAHDV?iZ%L92OLuv@NtKbtwO zk?G|qtE|D+VStd&cAdOTyv;!C*_^PCV3^=oci;VxeFT%mF4%rm{wA9Pw?BlpfX{-v zA6etxm}9>@)!UJuV-Am`T^F?q)dU%Fb*pJBBazDO^g6JBK}Xp;Q5L4kjQXjn!g||# zSsvDxPiahOOoE5xZC#AD9gap=e^r2j`|$NX^r*)sKIK2+ zQ~rzoweIkLiXq+IZO$u8#H~!(>7XCs@IerF<17eNr90>ZJ7Lz^i&Ev$N7V1N_EJ?T zHW_4xeHhXpg_pD>*^#Qix+I-1fEsPRL*;2!boTn83_F4LV2=hCGr^C;Y)@IOC{9$J zRvNEBr%%<=K`7&}+lqn|A_R@rEi9qzzKpwBvyO-mCA-SebEk5RF69i7bs~?^zLutr z65y1@N#wVv!_%i7rWP$!E+3~3O(Nf>UhDd-q1%ua(n&EH<#l@dvqOPgS*9?4_F0Z4dVG#9<;oQW~HYWFU zngGZe$$Pa2lRIY&{>pPs0bN@--)b82Nk~*R2w;i0tt<)xG+v}eTYopXd=6hzd;T6i zcTiGVf$`ze)E>T`rA6z~&qi@juArp5An-soqIAPY{J4X>L9QN&ji#gTSZ>frvlDmaGrER3l=q`B8APv1 z*^^CM>q=W03@?lNq_#tfsKo+>-bnrW!c~22RU+*4Bl33ANJ?pg>3es%g-kM!(lCar zr<#3qREq)*{fJZ1+9frUc5!;#EtePY4Mw;5=sv3LC<`Wzbbt>3eVfxgPH&TP0);re z-Z@n0PFC5w#>m8N0wk=mkqvHwJIEFm@`bZ&VMo@Mc0+_yBLsA;Kxw{WI2zR)9dk5o6>Cd4+_t-sQ&ymAG zIx6So{j&Tx^8^yw`$&xYNS{!(msN8lF`zj^eh%mHg#7}`%>9|%1+59^zJa#VH?+Ag zkjc!-e`R$_u={hOeBZ!2n89YNOnA>Bf}oEhD*Bo$BdjD!SaSxP)~*Q;2_OAe*#Z*w zXf~fYp3f^Me18E+-ODQtCaU)dBGk}cMf)OFB1sjJnfn04c`ct^XQMjIrGC7m&uiV@ zxk)=KJ>coDp$8U}vSm@zX?$7C=C<%w*d&odQqN1%E#r~D%~r(RSW}JT)!aH>JJE$=#D}&Mn0

6t3D`ApoU~FeS*#Z`!BdKx(;6MN@QdPYpyay1cQu!TBs%#P_svM_5p*K~%K0vnV z6sbZ-GfDANdDu)Z4U+g)q-OgvN+Wp?jkPDAq770)C?JH@86<(Eho*Az$Z651GJ$;L zv()EBA{6J96ZVl>CB=Iz_g!(%UlAh(^9lT`Uo1p)xT|F_(xBg64V=a!akGmkuPJ zQsCng%}(7Pfi1;j`5H~c_Z9EU>$En-XnB>$>qLG)$KwAG`n~p$$t?deMNbNUv3GR!G1)c|1D(MKymqQ86=FAV6og*HLB4%o=At zk~Z{0KgBxf&p?=E@kQ=&deT`|gRh}(AT3u-4af}7tQxLiX?fjpxPd>+tZI+tr#^k#;mO@~b$RUbCWoY(qbllwC06#x-rj@*4W$DbVX5O2V7Uc=KiL9jbh7^0d7&-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=