From 0f85142981197ab9139b58f26e80527c6291d9b1 Mon Sep 17 00:00:00 2001 From: Less Wright Date: Sun, 18 Apr 2021 21:21:02 -0700 Subject: [PATCH] fixes cosine warmdown --- src/Ranger21.py | 31 +++++++++++++++--------- src/__pycache__/Ranger21.cpython-37.pyc | Bin 9849 -> 9863 bytes 2 files changed, 19 insertions(+), 12 deletions(-) diff --git a/src/Ranger21.py b/src/Ranger21.py index ad1a429..347e615 100644 --- a/src/Ranger21.py +++ b/src/Ranger21.py @@ -89,8 +89,9 @@ class Ranger21(TO.Optimizer): use_cheb=False, use_warmup=True, num_warmup_iterations=None, - use_warm_down=0.65, - min_lr=6e-1, + use_warmdown=True, + warmdown_start_pct=0.65, + min_lr=1e-6, weight_decay=1e-4, decay_type="stable", warmup_type="linear", @@ -127,10 +128,12 @@ class Ranger21(TO.Optimizer): # warm down self.min_lr = min_lr - if use_warm_down > 0: - self.warm_down_start_pct = use_warm_down + self.use_warm_down = use_warmdown + + if self.use_warm_down: + self.warm_down_start_pct = warmdown_start_pct self.start_warm_down = int( - use_warm_down * num_epochs * num_batches_per_epoch + self.warm_down_start_pct * num_epochs * num_batches_per_epoch ) self.warmdown_displayed = False # print when warmdown begins... @@ -202,10 +205,11 @@ class Ranger21(TO.Optimizer): print(f"\tclipping value of {self.agc_clip_val}") print(f"\teps for clipping = {self.agc_eps}") - if self.warm_down: + if self.use_warm_down: print( f"\nWarm-down: Cosine warmdown, starting at {self.warm_down_start_pct*100}%, iteration {self.start_warm_down} of {self.total_iterations}" ) + print(f"warm down will decay until {self.min_lr} lr") def __setstate__(self, state): super().__setstate__(state) @@ -217,7 +221,7 @@ class Ranger21(TO.Optimizer): dim = None xlen = len(x.shape) - #print(f"xlen = {xlen}") + # print(f"xlen = {xlen}") if xlen <= 1: keepdim = False @@ -270,9 +274,9 @@ class Ranger21(TO.Optimizer): else: raise ValueError(f"warmup type {style} not implemented.") - def warm_down(self, lr, iteration): + def get_warm_down(self, lr, iteration): """ cosine style warmdown """ - if self.warm_down == 0: + if iteration < self.start_warm_down: return lr if iteration > self.start_warm_down - 1: @@ -289,8 +293,6 @@ class Ranger21(TO.Optimizer): ) self.current_lr = new_lr return new_lr - else: - return lr # def new_epoch_handler(self, iteration): @@ -301,7 +303,7 @@ class Ranger21(TO.Optimizer): if self.current_iter % self.num_batches == 0: self.current_iter = 0 self.epoch_count += 1 - print(f"New epoch, current epoch = {self.epoch_count}") + # print(f"New epoch, current epoch = {self.epoch_count}") self.tracking_lr.append(self.current_lr) def get_cheb_lr(self, lr, iteration): @@ -497,6 +499,11 @@ class Ranger21(TO.Optimizer): if self.use_cheb: lr = self.get_cheb_lr(lr, step) + # warmdown + # ========== + if self.use_warm_down: + lr = self.get_warm_down(lr, step) + # madgrad outer ck = 1 - momentum lamb = lr * math.pow(step, 0.5) diff --git a/src/__pycache__/Ranger21.cpython-37.pyc b/src/__pycache__/Ranger21.cpython-37.pyc index 87b415c7875499708b15fb380813b587365876b3..fcda430508428c065efee610c791e42b2458d50b 100644 GIT binary patch delta 2180 zcmZWrU2IfE6rQZnxcTyM;n&Y0;7pLbUQD2*uJup;Wd@XNBI}X&`}!yb@*%l-wqm;UgrbX(!+18-@p56oBY$+@Vt%`_$A=0 zflr06*sl;`5}^pMa0{P^GP>j}=_XwymLjOoNAkcV9Nw4I#7BQI}qdz8AW?{qProzL;Dh)ltl!dcCU6=6G(p$oqiL_sK847CkEW_!7O|RVy?a4Wd0t*`k@_W_EVM z$!D2WPNCiyKXa+0Vu>`{2sF)X_y zv4u0y1G&P9Jn!x4dWJFqm1+NQ#9Dm`zT$QOK^$-sjBX|h$znwK(I{ccl)_(NSW8bi z$_fW>PPL%k3}*#0j}eGc9pLq_7jMH1>CpJAI2k}Vgm4g`s^w;9^92SW=+Vwl7y4}X z8S1A2cv8P2N2B+*97VqUdbas~oQ4pv(eRxJivV!OILda0okeNCoQd_)KKVs#oMxpO zUqc7v`uI)lJzRWW-jBD@!Ats!^vQC+-miFWg8$lfQdVEPE3+=%c&N`w>8C-7nl zaa@r-jXeu}WpyqEe6N5mQmO^LqomvR4t|B{qIRkVcD!z?mLF{@)as}ZXjl4H@;_a} zRw6u%PN^#v>z=Vl*sF@YdP}&CT43D@m>$-+NY7yw*Ms7&eB7AP9UPjI)y|S|+MJxsAPSAeblAO=wj*`&<7N=NCPs*PmlU}2fE5;gmI~P($V$`Z zUESajG`JXQHVS;CfQO)87Ou2nb(b>WdymM7v}l7hh>$|w%%w(1S()280@Y#lE;lz@ zH25~=S^%JfR@Ir)X2&M^L+Uea1l0%ROiSy*Q55Yu-+@vV0rR=)m^Aa#cF98zo#STJ zIaw^^3f!)Tw~&U#$}b@7MmPshRcB}Bp~3#r(ofIJRJx%7_aCi$BgZFmg^7H2I+rb; zp03hK{(&4$XVM>{}%>&0LbH(h$4Cna?b8@E8%RiP^)7|v2{65{Byo|gn$Wu8q z^G{KtXShW5C0Wxtq#eN`_R8_rcbpCv^=k4#>nHGML9}(!VR@r%EP_t*^9Z)%@1k^3 c>X{w1UuH8K9EAE!c`?H@yVqTnPcmo!0gt^V4*&oF delta 2197 zcmZWrTWniJ5Z%4^`r7d;j$_AmoH&jXr*)gAFSWEK6;S$sK0@1srlb$o&6<###7WjS zq(p0lHmHS`2*Z{Nl@{4mzu>D1U;XA2^#k#cT=7LEKvb3B3xq0|*#xN-mhRcvnX|Jq zyR#!-I`r2;|1O`;P2l(Vx;XQD{~iBhgNCKo@w02Q`u+5r{M*sIC5jaIHN)qGPlvDA zs}W)op$V_36Fw1UbkR|Y+H`?bG(kmF?l*JuAwjZU(a4N3gWnR3F!R}d80$ks_(NoY zf~KB100Tk?Y!H6Hpr{9gRe&LdVG#kjNuVYokROGeNj9did4zx*SIidSX66{$T7_pG zhJ*+MwuvUdq@uKo29Q%C$jF$hOnI-}QR!q(<`Q+x6(dm+CshnAT`)xuN9Ddt5UZg~ zivz5#pfPtTX0KwN1xg8wUJ?70+|*WPK48L}`ESnasDfzgq@6z@k!6#!g-mnJDCWms zubERho|`E0t-$36<&wAZ>=1}U;lj*B7U?rn*-<;sbN1+Dp*ZB5Ddv@0dEqbNsvQW& z5f00$Z>#wZh@?gS=<9O2s7W;%kY4{i`pwdP{|!pNTRIm2-EVR(_$YiBSDiwbK^Fgm7jzo7kCr^KR6ixQ|r2C`P%VCziXk17hZLv>jWrrBo2 zX`7ZyXi%lZs>5=KkPu;xIEj#wt}4|nbnptEnb}V=&bdf(c(OQJ$R~^TTMgs+Bpw1D zSyG*+{36u7#ly20>3)O*2>7m+5926GClE{o2f{vte(4PNtXB&g;9`nRpd^}`8k%1- zDIAn;=v^Qg+FbBxt}v44n>O~ysc;(|kmYca9+r2)p?S;{bAF--j0(IW%dA3GF#;JF zeY^oT#M_^X_vL4k`Kig_H@cIv=+CuqhEFjZ#dlH~0k>6Yeo5c)V=u6+f=2(#onTfkqYUCI8bk z_Y&cO!qrz+tbS@mTv;_`y)Ei2H?VF8Y!7Q$pckNNL5tG)UHPab9n<-rq8S=nD{Pk$ zO;uWIL08rq+1;9%e_9_0eMZDYlqD85=vWwlRyOzmuBfaPt=1Y57cH!f1r~Lc{F>ox)OViY_oN%TeLyT z4YKw{9ij+U5>*0MlMoG68Yhs2$s#fEUP`G<9lXyDLe#`+WZp#$jyzBa2|r6VstDv~ ziD&k$gNTR~h@hrNBvJMdjH`l{4t(3CG{nAMq(w?}f(Qj@m&{ZF_@7l5e^Gug2{F-xke6l#h^ECs;BFC+U>fg-q z(Oh9TpPk5Mi!&27I?6B06RC9S3c7nA;RAqL;OuCwm>r&k$2^Qpc@r