From df35971812f2ca57f16a1394916b734bec43d29f Mon Sep 17 00:00:00 2001 From: Less Wright Date: Sat, 24 Apr 2021 17:30:08 -0700 Subject: [PATCH] NormLoss, Lookahead, new settings display function --- src/Ranger21.py | 147 +++++++++++++++++++++--- src/__pycache__/Ranger21.cpython-37.pyc | Bin 9863 -> 11947 bytes 2 files changed, 131 insertions(+), 16 deletions(-) diff --git a/src/Ranger21.py b/src/Ranger21.py index 347e615..fd13e19 100644 --- a/src/Ranger21.py +++ b/src/Ranger21.py @@ -1,4 +1,4 @@ -# Ranger21 - @lessw2020 +# Ranger21 - @lessw2020 and @NestorDemeure # core components based on: @@ -16,7 +16,18 @@ # big thanks to @kayuksel for suggestion to include agc, and initial code, # @lucidrains and @rwightman for additional code impl reference -# flat lr + cosine decay: original work 2019 +# lookahead: +# Lookahead Optimizer: https://arxiv.org/abs/1907.08610 + +# norm loss: https://arxiv.org/abs/2103.06583v1 +# big thanks to Theodoros Georgiou for TF code implementation, and he and team for inventing norm loss + +# flat lr + cosine decay: original work 2019 (fastai team) + +# Chebyshev fractal steps: + +# This space for rent - send in your improvements! + import torch import torch.optim as TO @@ -73,9 +84,15 @@ class Ranger21(TO.Optimizer): self, params, lr, + use_lookahead=True, + lookahead_mergetime=5, + lookahead_blending_alpha = .5, use_madgrad=False, + use_adabelief=False, using_gc=True, gc_conv_only=False, + use_normloss=True, + normloss_factor = 1e-4, use_adaptive_gradient_clipping=True, agc_clipping_value=1e-2, agc_eps=1e-3, @@ -103,6 +120,34 @@ class Ranger21(TO.Optimizer): ) super().__init__(params, defaults) + # core + # engine + self.use_madgrad = use_madgrad + + self.num_batches = num_batches_per_epoch + self.num_epochs = num_epochs + + if not self.use_madgrad: + self.core_engine = 'AdamW' + else: + self.core_engine = 'madgrad' + + # ada belief: + self.use_adabelief = use_adabelief + + # eps + self.eps = eps + + # norm loss + self.use_normloss = use_normloss + self.normloss_factor = normloss_factor + + # lookahead + self.lookahead_active = use_lookahead + self.lookahead_mergetime = lookahead_mergetime + self.lookahead_step = 0 + self.lookahead_alpha = lookahead_blending_alpha + # agc self.use_agc = use_adaptive_gradient_clipping self.agc_clip_val = agc_clipping_value @@ -137,10 +182,7 @@ class Ranger21(TO.Optimizer): ) self.warmdown_displayed = False # print when warmdown begins... - # engine - self.use_madgrad = use_madgrad - self.num_batches = num_batches_per_epoch - self.num_epochs = num_epochs + self.current_epoch = 0 self.current_iter = 0 @@ -169,6 +211,7 @@ class Ranger21(TO.Optimizer): # warmup - we'll use default recommended in Ma/Yarats unless user specifies num iterations self.use_warmup = use_warmup + self.warmup_complete = False if num_warmup_iterations is None: self.num_warmup_iters = math.ceil( @@ -184,14 +227,29 @@ class Ranger21(TO.Optimizer): engine = "AdamW" if not self.use_madgrad else "MadGrad" # print out initial settings to make usage easier + + self.show_settings() + + def __setstate__(self, state): + super().__setstate__(state) + +# show settings at init or if called + + def show_settings(self): print(f"Ranger21 optimizer ready with following settings:\n") - print(f"Core optimizer = {engine}") + print(f"Core optimizer = {self.core_engine}") print(f"Learning rate of {self.starting_lr}\n") + if self.use_adabelief: + print(f"using AdaBelief for variance computation") if self.use_warmup: print( f"Warm-up: {self.warmup_type} warmup, over {self.num_warmup_iters} iterations\n" ) + if self.lookahead_active: + print(f"Lookahead active, merging every {self.lookahead_mergetime} with blend factor of {self.lookahead_alpha}") + if self.use_normloss: + print(f"Norm Loss active, factor = {self.normloss_factor}") if self.decay: print(f"Stable weight decay of {self.decay}") @@ -211,8 +269,24 @@ class Ranger21(TO.Optimizer): ) print(f"warm down will decay until {self.min_lr} lr") - def __setstate__(self, state): - super().__setstate__(state) + + # lookahead functions + def clear_and_load_backup(self): + for group in self.param_groups: + for p in group['params']: + param_state = self.state[p] + p.data.copy_(param_state['backup_params']) + del param_state['backup_params'] + + + def backup_and_load_cache(self): + for group in self.param_groups: + for p in group['params']: + param_state = self.state[p] + param_state['backup_params'] = torch.zeros_like(p.data) + param_state['backup_params'].copy_(p.data) + p.data.copy_(param_state['lookahead_params']) + def unit_norm(self, x): """ axis-based Euclidean norm""" @@ -263,6 +337,10 @@ class Ranger21(TO.Optimizer): if style is None: return lr + if step ==warmup: + if not self.warmup_complete: + self.warmup_complete = True + return lr if style == "linear": new_lr = lr * min(1.0, (step / warmup)) @@ -386,6 +464,16 @@ class Ranger21(TO.Optimizer): state["variance_ma"] = torch.zeros_like( p, memory_format=torch.preserve_format ) + + if self.lookahead_active: + state['lookahead_params'] = torch.zeros_like(p.data) + state['lookahead_params'].copy_(p.data) + + if self.use_adabelief: + state["variance_ma_belief"] = torch.zeros_like( + p, memory_format=torch.preserve_format + + ) if self.momentum_pnm: state["neg_grad_ma"] = torch.zeros_like( p, memory_format=torch.preserve_format @@ -424,12 +512,16 @@ class Ranger21(TO.Optimizer): # print(f"bias2 = {bias_correction2}") variance_ma = state["variance_ma"] + if self.use_adabelief: + variance_ma_belief = state["variance_ma_belief"] # print(f"variance_ma, upper loop = {variance_ma}") # update the exp averages - # if not self.use_madgrad: - # grad_ma.mul_(beta1).add_(grad, alpha=1 - beta1) + if self.use_adabelief: + grad_ma.mul_(beta1).add_(grad, alpha=1 - beta1) + grad_residual = grad - grad_ma + variance_ma_belief.mul_(beta2).addcmul(grad_residual, grad_residual, value = 1- beta2) # print(f"upper loop grad = {grad.shape}") variance_ma.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) # print(f"variance_ma, grad adjusted") @@ -490,13 +582,13 @@ class Ranger21(TO.Optimizer): # warmup # ====================== - if self.use_warmup: + if self.use_warmup and not self.warmup_complete: lr = self.warmup_dampening(lr, step) # print(f"lr = {lr}") # chebyshev # =================== - if self.use_cheb: + if self.use_cheb and self.warmup_complete: lr = self.get_cheb_lr(lr, step) # warmdown @@ -507,12 +599,21 @@ class Ranger21(TO.Optimizer): # madgrad outer ck = 1 - momentum lamb = lr * math.pow(step, 0.5) - + + # stable decay and / or norm loss + # ================================== if decay: if not self.use_madgrad: + # stable weight decay p.data.mul_(1 - decay * lr / variance_normalized) else: p.data.mul_(1 - decay * lamb / variance_normalized) + + if self.use_normloss: + # apply norm loss + unorm = self.unit_norm(p.data) + correction = 2 * self.normloss_factor*(1 - torch.div(1,unorm+self.eps)) + p.mul_(1 - lr*correction) # innner loop, params for p in group["params"]: @@ -581,6 +682,8 @@ class Ranger21(TO.Optimizer): grad_ma = state["grad_ma"] variance_ma = state["variance_ma"] + if self.use_adabelief: + variance_ma_belief = state["variance_ma_belief"] if self.momentum_pnm: @@ -614,8 +717,8 @@ class Ranger21(TO.Optimizer): inner_grad, gc_conv_only=self.gc_conv_only, ) - - grad_ma.mul_(beta1 ** 2).add_(grad, alpha=1 - beta1 ** 2) + if not self.use_adabelief: + grad_ma.mul_(beta1 ** 2).add_(grad, alpha=1 - beta1 ** 2) noise_norm = math.sqrt((1 + beta2) ** 2 + beta2 ** 2) @@ -637,6 +740,18 @@ class Ranger21(TO.Optimizer): # p.data.add_(weight_mod, alpha=-step_size) # p.addcdiv_(grad_ma, denom, value=-step_size) # print(f"\n End optimizer step\n") + if self.lookahead_active: + + self.lookahead_step += 1 + + if self.lookahead_step >= self.lookahead_mergetime: + self.lookahead_step = 0 + # merge lookahead cached params and save current ones + for group in self.param_groups: + for p in group['params']: + param_state = self.state[p] + p.data.mul_(self.lookahead_alpha).add_(param_state['lookahead_params'], alpha=1.0 - self.lookahead_alpha) # crucial line + param_state['lookahead_params'].copy_(p.data) self.track_epochs(step) return loss diff --git a/src/__pycache__/Ranger21.cpython-37.pyc b/src/__pycache__/Ranger21.cpython-37.pyc index fcda430508428c065efee610c791e42b2458d50b..21608df9b62f4ef688a920b853ecca6398c20f16 100644 GIT binary patch delta 6552 zcmZu#Yiu0Xb)GvryBscemrHU zXDCt3-EJYNQR+ki(Mf>-u8pOkx^0TIjqNl|kQQjs-zi|U{V@%KAWacCMT`7tksof+ z@7!5Zlw%Tm?z!ij`DJlsc14#x?wX_dB;QFJ1Xo>TSjj``={$%pUS@ z>Az%8ye}hH7(46#n{l8uMuPZrzXsrF7a7+s>iozZ?F!@leBgHX#fZW)eDJn;(cnXT z__lU2%CkI&HpWNzDB3PQ#>df`e1cD+?dEwtg*MKo`3%|~KFjCOCip(SA8nE!;03g) zXZbuo2&C6ZxBIp=ev}{Lhi|i6>|*~N&Co6mysQ*su4gY-oFG;y)f`(CCqqGP3T+pD5&YZ_n5KD})MdZy zl>>5&v%Aseq{nV)ZCx@sxveuz#-wINWCqZXgMd+~l_Ty`Tn8Sj#dw5T*#+2zIW)_Z zaoN(WZrOuQToSuS8h{B-Fex#Cm68d-UYP_;%M@Ur90Kg;8X%@A_lpnYAlKGR@eMD- zqqo>i<_*fhTa5TCBsGNXWHN*eLw?x(5OOh+lcOyVvoZ&WUpXiKPmTZ|@kY@bk%Jh7 zfoiHTPCFl0TbodDQci%9mnPtp>;{~MhBRwN!CAHRoRZup)7)Gm-hSB!ctDN;7UVd6 z|GW}8D6u!|kXq!h%+mL4G4U1eNc$-7=5g7}<2fy>jcVKUUB_xMnWa`OavYnDY%%v9 z*K0cOY38NYUf9-{#x=vj9i zt20f|Q(LcE%eGfsbzG6733`fk;aE;>1-2F>N!I4}vQsHJSHzPfkRV#EE~=Hf>jtm( z?x5BcyXe(LkgnA0Yxb&Ra|?vhH76LjPpLX$#qmm2CrI5Vx}MVrdhhq`N@Labe;B{c z0)MIJwEtcYVe^T{)4kViQC)9X#d@_-aXiOgPCU{IlJ08#y5%~aSE{YJg{i=BoywIU z*06BjU%XLua`uW>}ja?j49z!TyfP!Sc&i9OV$cn{Ph zjjWl#QiP%Hg4E3UoUy9dF4>~(6saQEBVdQMT86tcir8Q z9qy^T|D*1IZ*7i}spiRQF-4`4=T6bSkzcQQrAmG?lCOxukjUc5#W@166Cf8=-XLBh zY>mKY36Ot_&k@-5Siz}#iNFg476~X9{Y}DJuMm8d0J(v35*Q}sOe`m`~hu93^6k&7LT zyao$}&P_iJPR1J(pTX zX)$pNj8Iik+#B{5KCvrwlRN&={^#|_ng50U z6D|5mIwBY*B68D!nkfOyO7=2ksm7Wvj)9}Qhd8B6$WcMZF@e4$nTJCmsZs5gp3$07 zm1r6=s$^)aha-=8(RNJgJsQbvv#gfE4d!)m9O;__2PPa_*E)>FjS6$E9wUFaR1)GvFr;Up6ZwjHrkuGD&1|$aKP+;!Qu3U13l7GCMP)Vy3cp z6-ZPndpyi&7sOxC?0@P1Lw1(E?SGVgxc49k-T0|s_oXiXQ0~-0eIx5pZG!vj8Es3y zWWJ_1qs`bB+xQrC)G0Y+{NS)lKli_uvsuOepWKO7-ix$T(k%DE;p03GhX;hi_lGof z-w)u{Jn2QskOS=HU3Yaz^*|OkCU})!;U!LLzd}nY`9a9v^oFReC=c^g>#pAH#$>Um zmcy-n0(F_%LJZM4fla5;CeilgG<;5~*$Y_{v@|jR!XY<+`k~neyc=}7*lD{zLca%n zs-yJm3{T`TwTWH@!SQ{&G?q?Zm9xr+?(JTl~+t4&xrEgVPqbZ zZy`SeYe;t8Cn=TF(6%?phgChNmGkW>X|$)e)h*@Cv}dtogdDgzgx;JS+GhC554|3q zr4d`k2I{DC499I;N;G4krj$l=(8@40I)XMUTlZBQ1)=8wRcqyaTM@B2 zC|DEQI*vpu7rNnC1wU zc!!LxorUnQOrl5~8pYYty)#S^6?Dp3C6vE*{JM& zXjt3UDcm*3q^Z_O%CStVv%r?V@l$*|d_?)qzcfuuW~g%25Yt=6Rh5`ZGQBu)n&WEA zQE`)x%lKLsw(!P%PijBBO@1^^Ua;d`9ohuUyN9#2NI0U^ek9y=hdrs-x?hcoKtiicgLceucgj#mz&9)-R_zQ%}=) zb!hVs(wfssl8!e=$BNmehh$FqA?y??CdGD>3Hm=6A8id_@(-}R9Y4;X|1tVWYCA>$)F2&ieVn6|!dqNiEK6wB< z6|_j77jO++{h8&oYDe;U}ev+cBu31lFCi$(ksnaYh!W8 zBQBkRKcFyE_CFtj-hKSn^Gn0x8=ylxQeXEv>|dNNoc4c^pKVdvCjI~* z=yve3Ko13uN?wB~bt_>x^tKWxxfZr7Tqj7pM3q<7IVVJ2s5CcB9KJ$wy$#{Gj=1Jn zs1~cX7ZzQ+Id52(Ez`P}G~@u{0HOP7q)}b3SYn#cE}Qcrx`Ek2tRUj9SH%?)R(TWU ztAbJvm81j)@}g>(r5Wy3fw!%aTeE9n0jP4hKO$N6MDk6-RAG1WkAVh}M*VvDtn(~! zDBTU;aU#5QO+CAauag{7-8vq=)ZW1H6VoHDUnAlB73H6zTR0dNL|qVf*O3n!^ctqN z9`x=#R(w#vN-FygSUD^J=qy5wh9*H>b>`n9>^%bC1_;v2CEK-%E$H9DOIp3QApVM& z-yvp_(sK*PP~z)$CCKbulsh!i;laQzm6CqDQrdKQSO=?zCcU_Zjoa1bAi7QhK^%{N zowb5^saC_=(~8LBx$*n7?%&Y55rHSQza@-qitiHnul>}_*qI<+tCu=2gF!c?%a)7f zf;ir~I**#4R{4l+Oz|$|($DxSGY=;tR83Q*Y_N31|HGNdB+_a#YNDn#*@*6ce`bC% z%@R@k43i~D9FI_@KB8w)XeU^n_51%jbDlB(soA*}=E;L)M$%MqMkeB>p3+lD*TV{Q znvLOwYBRZ8h0?YQV@qP5zKasm^XHe&FNs6Imc(Hq9wBg)fO>d6PS^s0hX58A&o80S zc$oT+0A$oi@hJ6A02J6_A>)5IJNF{JwyBpbQ6=p62;3#043m6OC|f4m6=ZNi8T6kM z_C5h}>Y%Gu$NAuRw*ANECNt{cEJfS`z)SeKOa9v2QU7=55>KVBC!*%8X`0>UkeR_R y3p{H^Of$xme1V-W42JXesCxEDy@3D_M8n4ar@2G15>5Vher(@mmiJfo<^LZnZidnT delta 4327 zcmZWsU2I&%6~1%t?)vWkdUw4){*RN`UOUboA%T(-62L%8+y>mHY!lWuJ2&=jyz9;E zbrPakZHd#Sg{E@3g;qkKU7k>@J|G^TQdPBmMwO^Ug{~@rpc0j;@={eoRUi7DS(_Sa z>wYtH&Y77rXU>^(_LmoKpNXA{Mnf9@&iYzQA1%HSdzG;X_gCyaHtlZcpRogPO5+k^ z_qZSXrrQ}3#9vbaY^}(+R@C|AZS4}{3EqDzSTs~W$y2w?qK~Kfz%8xl=NX9*apiL;D<%ZLohsLXJ7$KT33Y(wkllF58mXPOL1tz^hPpii1Wzf1rFTmyny4Tz^@ zoNJd&vDkJ@u5YlL%n8c)1|xxG5(~+=jHRg;#)NQFLe3yLIRIKjrhrj!qT+R#MnC4n zF^b7J=3t&$DZzc<^h*Pnlrdlmt7@e*_X7tMwlgw^K~|DzP7))J(`ltag+of2VI?{u zLp*Sqc%u?ZmBwTSI4-lmT}om?_M^XBCh2UOOq_D|boTN-Zpt7x^I8t3?9ve@>j9ah zc3EqqnS$oYP1Y>(U|r{-)}S-Z!yAlgXd`m)W$hJMNR8;k81vXGumudQRd-{X*C9uz z&w>L(S0~Z-zvADL7>C5~Ab}05^^ zR(;!ajB)a+e|wC66ms{24vSV}JfGH}$RPR&v_ojakPXltT!SwTFBJTq8#AZe)8@~2 zGkBMnBS;YR13dppZq?2e`_5SW5n*xnT>NSGgLuI$CH|T(6g(fVlpW9CTyNOIGfSmP zz2cNg?yrN_SkXNcIx!kvZ`!4r#g{SJ3)k0crHhtRUa^~keA`_MU5DV`Lif6d!$%H? z9B73h&)2k9FL{B6C9GOg93^GW0E*^%vr=C!Etka<4c)(l$IcX!)^fR2Uad45#Jg&( zuG`{4nsN`xL~Cm`yY66L=b46GOf~AY(&wa0R@qq-;uuNat5$q8G8}q{hTl=cdL(mK zF+xYUW{DbKyH+nIw|gbTs&Gn;vLha)8DAlIir}pKR%9;y0J_9zejlKX=q#bf+;ntb zpAlgA*WHuRkpp^&`B=sX;6GuQEXI<$$#j+jCh&!HQ%^AT)ktSFJH#obji{p|xI@b~0z9DTxPirl|3VK8OpVXwm4al6A zjO&iCez)qPd{=ic$#B3$c)dd+C}$t{$K)+r`A;ILZ@3Y-52uzI{h_T=5Ye@(B6GN zNj+Lx1PJ^Dj}XwFF!vJm)h5A;V++fvtkp%^eRpsdOS}IZe68>e5?CPk7Qr_GUchQJ z>^kqQq4(B!*6og@qmHo3mrHhIt-R9Q?|v}!kB7-cdq-7^e4U1L$AYqjm?zi(>^!e{ zmZ&7@`73p9UoX1-BU9|Cd+*4SSt4NuZ|IJUvjz7LBg;shGo!orE-CqCv9{h2msI~E zF8q+^r{=KuAxPpm94zH-jP|p$?u(;G;*%f-@lhl;Kl8iqkA8KQUBBtzb+6+6c}Ux2 z7sPqJTLr1QOCnXB1l+6dQ{%*rU# zo0Jir>r%|;%ItE4P!isb7BwTEc5^ zHvPhHTP7rZD(@TO=%#UH9E(RJ-Drq7t-RXICtl)t8M@qu?O%0_4t%RR-DRCZmwY2n zzO|Fq25}H>6UN(RY)zT@!jvJ!)fM+`8LWmN^)}2uymevy82=GYN!$0F65I1U)EcC< z(GIEA^-Z5Zt{9zk%Mr=JpOZAPWa`? zLzp|O_@j5IDx^cuHY;?1&YT?CT6v7-D|63$k=7b*D^a?cEZqo}8OHl=<9HMH%P|CElJDw;WM{P7i;(^<8`5A>Z2=gHnLU+B8Izs}obW(bie$pVCl31G5lrQ3UKl{O^&8&H+e3C8AdkN^-(PSj|_Q z-6~Y~!m5YlUO6H6fTvy%vgSEB@|AgdF*=74W$=HQ(UIUL76cV%%lOovu-x^@y#ZCY zwA{|*jQjEAY><+Ts(Re6!zv9|6M z?sE4{?>Tsd$ZG&EXro@E%w>BKRg0BW4rqFQqOXdRBzeQVHa*g&w9qRsLN{!&)-0`7 zF56zbA?&6tuG*zbs5C4`oFlO(2&M_Bit~K6_0^JRTAY_v&uFgKM3LAk*HWb-C^e~M z=J}9IYrTx_YhDqkJuA(+wNO_DoA@rtsrv0U^&TU*3Gj@@+BH=}qGGjBw>63AFR!lE zZ81UpKH8sGt_lYu_qoDg`xZKdNtJ@VK1ep#QR?{6HkEe8_*F|(ta_Qs+Y)s4SVcTf zOR=h0L91-2@@22UY!OXe(T}KihhP)nB`#L1W~q!q)Gkwza9I2Vg6L0D1y(G+pVXOA%5wPwed>E4;0$(k(U z$LGTrpn^7lQg%?!xqqIWcgOclw2g6nP&bVP>fVqs8V<1=(e0nc`>8N1Waby=7sU*z zIZLpQV2)rv!2yDUfQ5zm#f3$2h{lHr?jkrsa5td977G3J^Q3Oo>{7`KmGB#Yue+b@ zOPzdxglqyTe8mld?-G#T2yz0UeCY?&`!N9ph}TzN>zTIYo}L>^tNJ@e+ydYce9cpC zYwo~A%!wFg#!Q<@bI>%+oN1V5fGMd0n=cHAEAEGLS^Nh2Y;IpbRg*8fGy9j=fcyOZ G@&5u{-a3W=