From 6dad8d4295040189fcd6d3cd2578526ba7f09d53 Mon Sep 17 00:00:00 2001 From: Less Wright Date: Sun, 25 Apr 2021 11:36:19 -0700 Subject: [PATCH] adds clear_cache, option to use la params at validation, move la step processing to it's own function --- src/Ranger21.py | 70 ++++++++++++++++++------ src/__pycache__/Ranger21.cpython-37.pyc | Bin 11947 -> 12744 bytes 2 files changed, 54 insertions(+), 16 deletions(-) diff --git a/src/Ranger21.py b/src/Ranger21.py index fd13e19..b68ea21 100644 --- a/src/Ranger21.py +++ b/src/Ranger21.py @@ -84,14 +84,15 @@ class Ranger21(TO.Optimizer): self, params, lr, - use_lookahead=True, + lookahead_active=True, lookahead_mergetime=5, lookahead_blending_alpha = .5, + lookahead_load_at_validation=False, use_madgrad=False, use_adabelief=False, using_gc=True, gc_conv_only=False, - use_normloss=True, + normloss_active=True, normloss_factor = 1e-4, use_adaptive_gradient_clipping=True, agc_clipping_value=1e-2, @@ -139,17 +140,19 @@ class Ranger21(TO.Optimizer): self.eps = eps # norm loss - self.use_normloss = use_normloss + self.normloss_active = normloss_active self.normloss_factor = normloss_factor # lookahead - self.lookahead_active = use_lookahead + self.lookahead_active = lookahead_active self.lookahead_mergetime = lookahead_mergetime self.lookahead_step = 0 self.lookahead_alpha = lookahead_blending_alpha + self.lookahead_validation_load = lookahead_load_at_validation + # agc - self.use_agc = use_adaptive_gradient_clipping + self.agc_active = use_adaptive_gradient_clipping self.agc_clip_val = agc_clipping_value self.agc_eps = agc_eps @@ -247,8 +250,8 @@ class Ranger21(TO.Optimizer): 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"Lookahead active, merging every {self.lookahead_mergetime} steps, with blend factor of {self.lookahead_alpha}") + if self.normloss_active: print(f"Norm Loss active, factor = {self.normloss_factor}") if self.decay: print(f"Stable weight decay of {self.decay}") @@ -258,8 +261,8 @@ class Ranger21(TO.Optimizer): else: print("Gradient Centralization = Off") - print(f"Adaptive Gradient Clipping = {self.use_agc}") - if self.use_agc: + print(f"Adaptive Gradient Clipping = {self.agc_active}") + if self.agc_active: print(f"\tclipping value of {self.agc_clip_val}") print(f"\teps for clipping = {self.agc_eps}") @@ -271,6 +274,23 @@ class Ranger21(TO.Optimizer): # lookahead functions + def clear_cache(self): + """ clears the lookahead cached params """ + + print(f"clearing lookahead cache...") + for group in self.param_groups: + for p in group["params"]: + param_state = self.state[p] + try: + la_params = param_state['lookahead_params'] + except: + print(f"no lookahead cache present.") + return + + if len(la_params): + param_state['lookahead_params'] = torch.zeros_like(p.data) + print(f"lookahead cache cleared") + def clear_and_load_backup(self): for group in self.param_groups: for p in group['params']: @@ -384,6 +404,10 @@ class Ranger21(TO.Optimizer): # print(f"New epoch, current epoch = {self.epoch_count}") self.tracking_lr.append(self.current_lr) + # load lookup params for validation + if self.lookahead_active and self.lookahead_validation_load: + self.backup_and_load_cache() + def get_cheb_lr(self, lr, iteration): # first confirm we are done with warmup @@ -441,7 +465,7 @@ class Ranger21(TO.Optimizer): param_size += p.numel() # apply agc if enabled - if self.use_agc: + if self.agc_active: self.agc(p) grad = p.grad @@ -609,7 +633,7 @@ class Ranger21(TO.Optimizer): else: p.data.mul_(1 - decay * lamb / variance_normalized) - if self.use_normloss: + if self.normloss_active: # apply norm loss unorm = self.unit_norm(p.data) correction = 2 * self.normloss_factor*(1 - torch.div(1,unorm+self.eps)) @@ -740,18 +764,32 @@ 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") + + # end of step processes.... + + # lookahead + # --------------------- if self.lookahead_active: + self.lookahead_process_step() + - self.lookahead_step += 1 + self.track_epochs(step) + return loss - if self.lookahead_step >= self.lookahead_mergetime: + + + def lookahead_process_step(self): + if not self.lookahead_active: + return + 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 + p.data.mul_(self.lookahead_alpha).add_(param_state['lookahead_params'], alpha=1.0 - self.lookahead_alpha) + # save for next merge 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 21608df9b62f4ef688a920b853ecca6398c20f16..631e95e017451ea8aad4a414311b107ab1244943 100644 GIT binary patch delta 3613 zcmZ`+Yit}>6`nggvpc)qwb%A~{mR;F$8oZb?Iun|Q<65XCMl)~trL_mc}&K8r}nOQ zXE%3d<43N$q%IT$fht1;rJ&eQ#G_S7wGx#g#E<@>5<=<^1R9X&Kls59qDlw}@i^yB zvbG|+)_imBz2}~L?wRxQTbF*FZ#|kw#3cCrxACPb+ZWz#{kPJc!7==tHvue3Csfj9 zI^_+k)p0sbCti{Cq}9IOu_4iYbdo;y3Rxq1C*4E$!b<8@Ns;ufGg9tx?;Z7KJ&Tgy zHw;e%o)A3FFNq|%5|?==Px2JcP;ygU&AMbwT9-ND*$tUUJi;X-#Jd0$o(5F8R17(X zsSI@34$~06xds>k8Iqy#7QiTv1I9S6ws6I}7jCbIfo|goz$9-4Z0Bu&9lRT`lS+WF zzu3v{@idhdG5xX>XD^!b+i*!mtCr7@nQv0kZ1>X$a~ zW_#^0@5Q$`WD0Hk8gbsHvMtk=)gdme@4KZPB)a#M)@acqr%Za1nG`oSh-XKT`x)fcYEHqPGv@Ma zC(a#5IDre^pIQbd^$t`aUn*3q$OQ`4EOr`sK8-wya;0q9Zmn#%%T-JFa`6Gs@!9yc z*yoVP5_0T#dY>MzIhL_#vNEkK+Iq)B%5YqhxkfebvNO2jAVMDDg7;p0qU|t{xZu15 zAdPq*$A^gQbtXoL=Iu``kwf0QiPUfj)_fI2xmJ#_DvVeHp%vjOz~35wNgR{4OxAI& zsx#i$ojgkF$8hx$0vbXn%Pt^QKxo>NP2;d>7dDH-iwMsloI(&5Ig8Zi5zZmZAfWD` z5-J#|w5gZ60F(E<`|TI%Do`3csFAY-;7Syd1d*0z!NJHg(xvmOWNC&ID)&k-IKVxmaiH0NtE{VBE{4I4f&ZX%b>N(IHF*wr z>ll>*b1t$nlFgSalQ~&;-pV$eD4RF)^A^npE|*>D7fj$2kFrcoPO>sw!A~~#8o{QO ze!KE0cectbM+X<0JiXH0+z=cV=wZ=Zg!Q1nxDaE^u}ZD#Fa}es6X#X8!t(QeY{g;~ z$0!vREMI{rGIJ^`c^`H3UDl%|^MMYJd<4HmT*DBT>w+xn;7L+QBXQD7Vnl{_m}o?i zHPS;?QZuGKXR)chlUufV$2#cK;>L+tGrv%)J~%3;uxqHJP;^-+ny*xs4VV}1u!4)% z5xDnZU#A!AOp{S>sPpa-6voiuSq&f;-XdoaXVF6#N@R6h{{(plQfV^Tm`PnGLwE|X z1=e!r&8o$ohbdpdd?d{DVb>ISqM@X1a{SA<63&;*a@7!y?JIMT!i1yys+eb}hOaKp zTg>vq)N;W$0$n1(-L8A@bx)I14Y?;mj%>!%i?EHoif|cW4yA;wrK->d?kViDA7>Q| zR1E3m5~P*s-aW)@TPP2BD3|=DTqD)8nD8dy4 z1HccP)v9IF;C3Q`d=}Sk0OaK0cCULInbdZdfu|T&wUVEAzKL#!nMwkeOM<(#%B`U| z@$UEk@i->XARUR5zJNpgimV^u1%x*MwwjK;j) zgvZ0|KVYl_Xe@Xi2xr-a186I1k1ELz4(2`U*HlIdDqFxyoFypJSVLP#A>keg}C)WxIovNci6c z$`4g5i$P5lbpk_36j47!3)k2I9DW`7Abh-+ay#mNx|!>YDy!rzXbI5c+4q1eHzr)g zkL00Ms6m7&*fKEGTW`B&vVv*nEu(B2POa>d0>eb&lVXrd&}Ko(c@nAH2Z5q)+F|v|qCUUfT9yr8L z0mKv{JR3h4qyfo^Vfb@s(~E?XCnQ-ipb;$MBQL1zE26|;>Fy06biM9A$0XH6zQpRK?aI! z$v~pUiu&*%vi^Pu0{d;yITs6hm(Kt#8p6%;Q4gxdJzSYa!2dEnsfgf`=vD-XR_X(o zON79jDIg?BRSG04sF3hrwH%gL2AhWce-56ss=3TrHin{i1Dy5N#$O<5?|0)n^usub zy&A*5h^-=g5#a`ch^Sjgi5SG76S2ZXAg<%^M+l-YMeIu8S;G52@TMmQQ=%ypm&gI% z3_Q+dZ*}6p6lC4yxT=k4?V6@VwX~Lkr$>v!6VkLWiPs@qo@cl#M^5Cr*rNB*L^pJ~ Q|C|^PqwQJ48=JiGKe5Sz3jhEB delta 2831 zcmYjTYit}>6`s3~+1d5(`n7AXUs*f0Yme>t1yV_!Hfcy|a2wO4Y16uOC*!?id)K?O z+dH$ilU~4?l{EIOi^j zyV87f?z!ijIrpB|{x#k0N_~sG72-wgp zgU0})ya6!Aakqho#P9Ux?KselJOS9mlYq^<4X}l3fS_4uVIT4|)z%F5uAide9dd>E zX`bF8xW9r@ZQwyFg;YDpw|jpAIUVJA7w|H?0}!4<2m3GY1iI4)exZ}6K_|7H(v+!g ziq7`X5a1z&S$+uUUTy&P@hIS7xD?ImS2&fY0ZrZmIHDv* zIk;wxsc#(N89dG|VK@Epy$2^~l*V`?jdf@ltxMa-b4|KB&){2qWD2|u?h@|}8gK)& zVXK#Gd($Fg93xUZZPXfMEVs-`#j@v#m!k_t1UdFfuJ)hk zAOWquhVPhtGUPeMWvN$eW|utn6^Ok|pkT3FUb9yno3c6Ni9a-S46v^P$sWVCQT6if zinWqsL&%#&_!@2qGoBrnE$|&HSIk!`D79`Es}B1H$~=TJiBh@bxPG-{`Ijn=6zlOr zPsp?j>m}RItva4nahT;)%DGk6fm%-@te~or+OS!vTCwuJ!)!lac0H+1i)>onaHaTL zJk@j_b&er82rr0G;z;93AfJNA;{aN}IFaZRxkNwF#kIsHdGz&Us^5hT8TNeJcTB>z zV8s#$NrYDb{-KF~WmC1slj8x>R(rMSB-uWVdoLp3sLAkv)NM|E?It zGipkUMe zjJRh{XFIs7%!K_i%dN`TWrvkLtC(MN*j3zo4FJs|{Y`{x2-n3#n@Kur&$Ycs`j7)6 zt?+_-U$r6jXORwJ_L^aE>67+p(qC)Hxa9QfxTEKacBx{iyON<5NP2a5a0=wJ8_Lqg zs>2+qQ^&U|;Mem=pF*$^E&@p1u2dYCs*{+3y6bnu z-+OvnUj{8t%@2|!O++WjfXE*D%W15xdb%o83%JCPur7pi2yX*?@gCS6uymSE98I88w!r`WPdo3ur zPl&b~d1>6dfUn*QFlF{0;c;a100d-kUc~~hQ2kh03=NNveQ|DhlzhAP%J3X%dlz<0 zB7^0s|9`npkDJHEpUr_G%vbh(02y_l31U-oRB6dDE{pZ7fhQZ4E&pk%O8-l$K8S=_)1t&f5r&9Ac=T>b#%pgP6S=-@Vc99akb z@Eus1qbf&bBnRWF3UwR8vBl6`*|=`AyzS;3t7KbVwIoSiWsxL>dN;wVQ(cH1MQR@b z3yOV!@Kc1J0m!DMyzN;txiVu;4(G-+`#ElYgqsO$BNp6do>uLmOnvdA)Im*$O3^r$ zE>YCMW63~n4IH&gOEO%=W+h_>wq-2ux-i95clsr&{0dcq45sW|q{wCVF|vPwi@IwE z+E5nMamytszB|@7Dq|J*-pFLvQAzdwg8_g?zJrB*s`mS_X)+Zv0?9y<=z%yCd5jz; z-GLanoH%E@D-N5U8pr-+cAsY>XbTHX+?+T|QsS=@gHPg!RCA8uRAJvjxP_nsdmAa0 z3K)Bav19Kc{1~BzfaxqFZn=K+eX%&1O{p=O#9b2ryNP#R>`u-MCpY3@W56(ssL^Jm n;K=};F@lDn6D4bs$ILXlAwHjMhyN_;sWCl|W_~SZr(XRZJyF%i