From 3f09b32df3ded1a9768261df65b78c37e0cbe0c7 Mon Sep 17 00:00:00 2001 From: Nicki Skafte Date: Fri, 10 Apr 2020 20:34:23 +0200 Subject: [PATCH] Learning Rate finder (#1347) * initial structure * rebase * incorporate suggestions * update CHANGELOG.md * initial docs * fixes based on reviews * added trainer arg * update docs * added saving/restore of model state * initial tests * fix styling * added more tests * fix docs, backward compatility and progressbar * fix styling * docs update * updates based on review * changed saving to standard functions * consistent naming * fix formatting * improve docs, added support for nested fields, improve codecov * update CHANGELOG.md * Update lr_finder.rst * Update pytorch_lightning/trainer/trainer.py * Update trainer.py * Update CHANGELOG.md * Update path * restoring * test * attribs * docs * doc typo Co-authored-by: Nicki Skafte Co-authored-by: William Falcon Co-authored-by: Jirka Borovec Co-authored-by: J. Borovec --- CHANGELOG.md | 1 + docs/source/_images/trainer/lr_finder.png | Bin 0 -> 17455 bytes docs/source/index.rst | 1 + docs/source/lr_finder.rst | 108 ++++++ pytorch_lightning/trainer/__init__.py | 21 + pytorch_lightning/trainer/lr_finder.py | 445 ++++++++++++++++++++++ pytorch_lightning/trainer/trainer.py | 15 + requirements-extra.txt | 1 + tests/trainer/test_lr_finder.py | 181 +++++++++ 9 files changed, 773 insertions(+) create mode 100644 docs/source/_images/trainer/lr_finder.png create mode 100755 docs/source/lr_finder.rst create mode 100755 pytorch_lightning/trainer/lr_finder.py create mode 100755 tests/trainer/test_lr_finder.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 257b99e4..7430c95d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,6 +9,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). ### Added - Added `auto_select_gpus` flag to trainer that enables automatic selection of available GPUs on exclusive mode systems. +- Added learining rate finder ([#1347](https://github.com/PyTorchLightning/pytorch-lightning/pull/1347)) - diff --git a/docs/source/_images/trainer/lr_finder.png b/docs/source/_images/trainer/lr_finder.png new file mode 100644 index 0000000000000000000000000000000000000000..bd1667b908bc68694049b1ab632225f29045c090 GIT binary patch literal 17455 zcmeIabyQT*-#&T>0hKfWX%%S+X_PQ%kZy($>68?ZP!R+KluilhmTpi4q;o)81*9En zh#}_gGrr&VeeXYa-MiNB&kM^n%sKn)v-f9z^4ZTeRQ<6$*#)`_5CoAaDm>DJApAfG z!uv@=1g;1V&n$uuLbnHs+9crT1A6WMwpr@6Y&@cr-Sf*so@p_!!v- zl#FxPHJ6ka+Luh^ATLaDYl;#Q8dFj|C7vVzzijJI@9dtQevqXMlDV^+mV{bAdaknY z!n%9;y=@CBbysrv2WdM$AsByec?tr2K=|Ph2?Uu55Xyir;nyL22+}9TqlBQgCnSO3 z67@p}G%u&cXMv!=4oV0Qf_Q@e|405`OYBi1LWX1a=mi%ir^b^fSq>s8lS@F7y!@mO zJv=;mdU~Q~o%&n&@y$eIV`Dkk*}FFG0NLAe@bGARde%%@Ccc7xkcEbZMo&q}C@Lzl z=@CFs!`CHOPImUZ-Fr^5gl$eMwEL<`N;kmLQ9dbtaTG61+ybxCyZ>{mM3^=)es!&@ z%xTEiNV%W)77rEB1DTb z_Hkr{P;DZ=EUb$6b1V);#d$TU$g{$vx5|JezfzKiR4nn|&coU%qrOTSXpCJJg%z4o zyO7R+aB3bmvV3i1B}r8)ZcOT##6#GBSIfy#vHz}_(~Dtm;s!Y)elgMU=Lt^FUnvQ% zLl{QbBnZjc`L(+x0{!K21Nrp1gWIb}?GoNlIu9ch9v+TPL%ooOA8M-qGcwK(H+8x?fgTY={?-{R+kSxy-&=m9WvtYK~}#LeP*~;Uuvb3lBXsD89_FIa9xzZ_((v zsp;fYs(2L-3RJ3jc@gfV@o$5&{dn>4Lkh`kN#FgII6aeL)3Y6#{rFyHpSA?(&|fMrlEx>=~8;9x-Ynf#ay^9O9?m>6HNWPOuiFEfJjPgpKzf6KSj*6pjuh{p&-7oL9n5Q zQ@1Nz)3q;^%?~F~INjoK--9Yhs{Tv8%;@XdVE-DJF;q{+&=`35Pp95+`9K8!G!4Ho zkdPdwXy@?Fe*V*D7KrHoekpL3v|U~2CLX^!{EjwvvfC!rRTx$(h~On`@G5%97BXX% zY{WASDmJ~#Xa5@Mq)oTUrWl6`!<`1M7Shc`XU85Y#BhHmf(?xyq4QckI#!9oOy*w$ z2_HynVFizr5OPKx+bA7*WxiRiH3qfz2s! zm1}(!kCYYHL0s2*Oi{hgCD6|U*mQCWo#u%KtYT6>1I@yRXPj7An0Xt@{U+OOUYleb z@o$_3Ef)g!&dy!<*Wf0U`|Q-T4L87o!n;0D5*`CT?)`;5`_OmNpE&74S~^}tGGIV& zaMF;aF#CxcP60WxdJ}(y`+Co!Qt@ltJss^?R>hvjp~Cf-*8k0@%Eyn(O;|a)LhH9m zX{Q8Ec0xJT{W};^S@8`ewV!ltGPn>q>Q0K?R}GCzbNq>qTRRmk?OBe{bu4}auIpVS zA!RA{;=kFG2?qAx$4%RuBw|q$ zc_W5CP(M!kczU|iJn7eE<$uHt&;{b*oEhH4XCkxtSp7+dVmw+KoHDzBetG{rZUv>_ z))zlg80xoNg`Us*9duh&GJ4QUxaUt=;z5-AwD=-=L6hea9lyP$GOTehSU!DA(u<9O zpr2kHq@4zu{ z0{8nR5CN$#6tyIS#{>VAyF*Qj(aI9oyVBC8pc`%APpEl~^HKf>ja~Nw&Q6XRuxG~v z1~(wcJmfsfpUM-;tq;*LF&=-CtnCIrN_P33eDwXR;pphNTvA@n-*y%Bj?u^Od_mZD zb8KS59yQmzHkvp1j5+FmT$zjA3Z77U4#CeR@7*dpB-hm|S8UL0$Yeq5NCxtd($+Jd zN@W6w#o_bCw?{mFWM&!~AK;5Ry1HDFk_g~HOUH{1K7nBHubBX$gOoCHo1Up3a^n(k zaIx{a={?5%(LmTSIU#`{+L7BAR?l0#ul!4DROUK8Ma&La!uA`zMig=nH2_C(`|YFtp1DL&kJpDZbbR*o1Q z>U#uU^G~DRdey`#Me|-u>b}I!1_w>3)!uGL6 zulWp|<@t}+AG$GxOA}v)g|rp^DL7#F+JBy)De`30d>F20FhCp8ttisS!0@l@-$l&< zNvHG>Z@Rpvhn%?&S|&EUdDz2z3xy zov7RagL(iDYRIBV>S3#oTMFyyeCM(at3CGIEnF`GSbbmhX=Q{`Sm;nwcyJ~LenXCbA&&rLi#cCr3r%^ZZlqsF#Lvse5a%QzE_QkQaByZrnpi*6=m@1>s5yvC`2h9~VQQ0=C4DWLl2}XsUnH3W4hviN11>xQ(03=;j#vRKl>55iBi1LQC z(a7l&gQa0u_%s`$tvStO?T~F1^@pz21y|X0>HnBJ7PVF}?RX@iQc5I=^WTx02B`6o z&86t<_0FY4W3m&RLPYL#20!s`3f`$qqx%^5)Kh$m0Jn%wX@jX>%!!;Rqf8IsEIjZr z{RR7~#w-(o51##WX_lZNPVH8z%p8qS^Eo$-oQ}(h1F7~v3Gs1;|rbm zh6qBqrfhPlZ4h_a*_ktrRh2n9^W8Vvn4*`s3?@)&TJM%(vXs#5q7o$utd=~bImkRd za$V(UHgW~wdOAX!;e}6Y8s7wYZNat%UF=6QPjM0m9uf7yL*TE+aT^wqXP(KaM9Gvd|KJK&GP)+7Wueh{$zz&Cz1t(wt7 zEdhlLo~78z3%I^J(E;OEqVdYh%hhFNWz!BON)bhOQB7I_N4=(xk0qJ?56veXL}tC0 z(5hc^`V;11+}JUu@ib`>}tOH0d!t*;lWJvR+!eNoLAbXkC359;7( zIek6J`lE)wzrR|-1yHw4x#gQ6Bbb1`$0jTZH8K#YD-gATh+_9O=zMhcn~OADe0*v} zdS&Cie67VaCeMGx&K0$^n264}y#DC7Z;+FneVAzoQbesTOn`-3%+;$@(I22ZPS$`a zrYQO0;o<0P%ieff;K`%2GarT&R#olA&e{K+sc+b*=oPZ)e8*IF845HT_dYen3dL~W z&yY<-w5`$#JPA5GS(aX}=;a+9$x%kP5l|>ZD=8_(#*d>Zwu{PI3i(D}AiExsuNQo3A%xQN1hH|`Ph z(M11ltAy83<+Wh*M8vn%EuUuBdSiJ+t)9AiMqQnl<l}@7#%-e-29J$#BWL7L=0o&o z`q^IM7P}%4dfEgN2+%zeY=J5a-?z%HzIh6zO8VB6Hmq^^p?v4HR8=F`9;fDi|E|C& z^+Lb(coyZ}#pso~Cwb}e<+y|dud&aEhtG3zav0r!JY6blG8Qn>Y=Qb33MK@HpfIxO ziRFyWo-hqSH2WTmza${u8&T$41ycOfjRg76jFMflwmy=zQB{IF;nq+w2Ur7odiuD; zL=HtC?CF6AQb(`bsn&iZTT#O0XPa#+swa-yH~#_9BE7D^%-()!N^i-#IRc$7bThRH z4v0Mj0@XA(&y43W$O`9eP>9yn&@i727srEE3L_!07U>DBfez4GQ2!((i#mVA9H7e1 zUuLt}S?rFU@*(JYcg@1p15~WFs&A^JzwPi*Q%9e+Oc=XG3`iar$Hm8cmf~?#_ivIuB+EEL3k$@*tH77p$h0{j?Jho=NeyjmGPVe&(h77k>Tvpm}%T zE}I@PB;zur*@q6Wxg0GSsFn8^S&C>BvL0YSj*aE#=E~_IJmP8`r*m6RH=V5#Kx_4> zZ-kUp|BUL9+GLe&OlKq`ci9B}d(^UP{i3^sj*(F>wL7R5MbWj~?im#te5c2RXM29w8h}lvnhze4@7>cR$krNKqvm5z7P1*aFMXm>)P$ zZ|Q(xeC@Kgv%}Xu4?&rhTIR?2Lw?uJWZ#~X$=m(LYzP)G#i@Z)1_`hei68jn@)ruA z!;{^?6a+{XuIGJ47T?jU+z{c6Vv>1l4>~9DWn9OmefW@ZxblAO_+>hr6p*+Xg-w-V zE$~UE;P;cbZ9uQe;1rRvleB%v$_LsOhNH(EmV-z6 zuBov*jR3W3|=xyi*zQOcS?KZMj*7iX$Knu<5;0hf_lq zQU|V+LSjm>Yj8)`ODpC|Uh70eXMX@zT6CWSCzanKwG>!#rKki!kr5!~3%G@>BV%lD z6(_}J+5h_*xsL6mgob6|*E~oxr=8M%M8WbHPNf4WJ%P>(M=;6C{ty-t?A~ytyn%bB zJaVp;%0MwJ>}OSE>8}`kNQPwkJ0?Jk4A+OgJjBqC4iGzs`=+g2q&BT&98D4^hSRwR zV*?cu04<4SpedNC5(Ba(!O15}gj!b_NOjPvGRc_i!Fk;Kt3Q`JUKAV0wii^mLoWbc z2E0RAb~fv+Teob1oAm*ns-O2FfO1z?@EFn-gz7H^(Sc%ZW5Yr6WUJk<)xUYX#3*lf zP};Q8dT<;Ve;+Ry6nNQH*thCH{WWWR;Fp4M{zga;@vP&j{A|G17h?3@sG5C$vM?P; z1r&AgAQx9)Nx=bbnig^Kl#Z^#B*O9Jf$}>(%|4se6JtX|8Qn4LDF|{1V&6kz0yyE- zd3prZS@u$^e!OSU-`}sKpfE)5+?4sNJ%st79)N*PwT-(iy-OFglR!eeRKzG#cSeyn zFfsutjl949uA|d!9{9ZKG*?(S7tpX!IGA+wLC zA@qBFg?`c5Sg#|Ougs6bL2)g4F!M>7xBk&72Caoaq?vdxT`e)L*Hn(@5oi>Gph%Dm z{l!j*R6qBVK=rxzGYn5}L{nw~jFy{=%Y9T;v;Y7JkMp#Sj*bi~o`&)^Z9Ra??Vq$L zXTFJ^w2&)Hc9PEm$8RG|^ zXW4p9Qh8BkkfYFP?8K{OY}!caZdA#fQ2}vj-)OOhL6={IIc=fSUOp&$8~*4DS5+0x z%@*pGWsV(kn>afAXtexnQAVjl>*zmmQlzSlGLS}8PK)9^*C=I`m2T&88KO2q(^ zdcDCHb!tGS8zQR>&%~#6e~-lfHauNnrmuQxl;kgXSB-O91E7f|-1#GoM!kD;<9Fif zAI~Y%(k!Xly#V842n#(~qo%jjd*GgZ$>K&V?d8m%6_Q(>kGMER9w-WYPzsCC)>N4m zco>{)^uo>wT{yX*qZO`4%tc>alC2AngeeWfJRh%hjd@W}cQs$ix=O#Q?pCDd6s4y6 zv1X#S3iCAo_MW007InGyx}<<1Xo2Lut^_{fsF`yB&X9BK5vMS(5|{${M57LJD61!! zT4%iW`-0=KX!W#1#k6C^bT&x{LV$`~R=&nBBDAX4S^*hxyU@wT_IPSVqq58Q(fhNt zW412Yo@?~Tv+7OiZfZ5CCHT78;blC2lFhruWgylz?ir_S7X1{oH7oS@J|ZC-REq9e zAJyXCUkRnyU)k)>F>GvATBhh8+Tr7NubcB5^Zh0#afO*NGsq_7(vusD<&)ZzLs{); zCt=?gb`INxpK}WGQoP^glQc#@xuG-r<3`z_I>OCpaB<8tc z<6u-NF;|%6{EX?a1a-t71-!!RGmy9YyUPPCCLh`Sl&s zbFY`Rz&TPevA~V&adhIG}T!Xz^!NpcJ;h! z*Spqz-$APLHkPPPIVOV*8z9@h#& zQlF3g{EvlPf-gU-H+GdsD=aC|(bUwGX{|M8&DV;Y0OJTJpLiDd&0;g zMjynH(w|$G|CU$P>20XH&EGvie0AjZq(SQ%Jpg3H7j7(RKTY^J|6`=^LW0p!sa4!W zz0NN?T}$6G7oB5thz{9p34P@F$%QQl`o{&xZ%5xYtodZz>lq8u5o0rTpI1A_5^?F9 zW8XKPlA&U_=CfIlKY?}R`1`^M_!pos9DN(T8;NVnUr<)RPSDllun-XSs*IH6fxXMO;^)o64IE-HgiaqP=V~1y) zbQ97#iUNrB5<(Zu1ug>u$}_NI#B<}OfZQ@g#PHEWUofG_s@&mEpnw1nHdmRWuB~I` zg(AA@=CPdFtLQ=pn_>Azigf_ep{qa{P4{7GEc<0avJk&&3{fuU$Mp*@nGTSY{DOhM zm{0x?@$^_#a_0D@b({M|Hj^{}5Ex{3j<*Ytf4&;R0Mog+h$&6uqT3zSKKq|e!y?Pq zn*>Kaw%@As(~cVri1`q|{JZnqkE__p$R`|lri*x#ZO;G#D8$8X@X0mWqy1Zrg=KX= z#vvdFK$l}(zT`#ag^J3x3HGb|uChNIlkah;js`!q*=Z9Z>sNV-Pw7GzR+zGgz7`f* zRpj;L9jF0RK$LSg$Q6{Uls^DDzn>*-a0^e-sPATD&jwG~Is5qqk=L~;qE_# zYfQsB_mzha^3pPiyo%s5jCmt}WhM$^nxeEUaEM(#CqQOES&cVAPx{S}^BegQoI zhgc<$CjR*jYD_rqAmM17%@sJ+6v{OX`&;jxKX;g=A7JVx13wy7{RY!D^#GXexEI^f z+RlkoSm@It34M8x`^zH1W(T3nAq5^~pe+at9=nsL9;1uKC9|eEtpT3uD978#&aZTi zJjKJ#b91HkgFWM_F~{8OaliP3S_$IF00dwh3l~VBrsN0JnwbPJ_)r|^oJ@`+oA~PS z#%HPWH5UNEE15SZ+?Y`l9o|Wpzx4!BQ*-1bqRL)X{mIN}F@1|1hXC!BQCC!ZHxBZ{ z6<;l4TV&}4Nvz^asVP5wlmqAJU85c_y^6DinU@V>d|Kq=3J((vnm#{@7P8ll?~Y@c z0YW6a6*8N<8?rwR>}kFCPQuy2!YL5pBYN_2o4?$qHb7No-MZCrC_xUs*_$BM{*4aqv1H~I%bP~6g}%k%mDyG*lc&$rN(D3iWFw& zc5hsIIr@ZF&KgP;2K2fj!0-F~cQEG_3eASu6aoR{CdfJDXr9XW@s zR048~Gp7a5x+``qMqO+7;g(A7-Teh1Xz04L_TvibIVppHCjWOBF$-)$%a#XTBC(<%0N- zd+M3PWLl#fak~VJks-%orwgb@SfttdLqY&C$msVdMvQ9BGZh8QJ*xqhr5n>#)RnL~ zzSpI@R_*a-=crDK^3U#bKOZmUt}SO%0OZF3h$N`G5+J98cL&DdM^nf2;~Hl$-RmjB z%WCO1r!Va)dtZ0>H651&Y-KHuvZVX<2mr9C0U|=V4<^Xli~3~grG}JDAk}%`9wc%V zxwe+4w_bXVUPreV{gPk4YYVk>x@!7SFcq0S%zM7hcg$HGj|Lt;>(69AwVJmm!7hAt_<7Dt@hV1uN2SJ zpb40L0}`xGoZ{wAF_@Ra!XwOQEM!bPIeE$iL{)x__v$vEt44T}-B&QGuY3Tg4RAsA z0|uZ8U~XD@@yW*1=CkH*OnYmDF0=JJ;B?*Pk1;RrDj`OqZ)tkgxKAKPTKwc&pU=0{ zEKE*14iY>bKA!HqrI--e7u(;t#Rv1<+D71&j{57{hj~dce^mD(!G!BmL4jl4!TQ4% z#;)Y#>|WlRgui2-H%Z-|AEt)*=swuG3YXsiMb%0Xnjir?p^XsaxF6tS(L$PjJ!R1A zHx1^63m<)9KiaD07zk6Z5Q`vAgxH)e%xktBxmROy@r*nZ~zpT zdD=pKng*Pkv%YKh2l%H<^mLZe8z@C;s2jBnr;F=2LI#+|dI6ZLL1jf_mU}Sgi^@Od zYub8s=gyrf|5KkduQ_ivopWv4X!NYFA0jQbGm`w!3gzrff0ib~iV*m)aeC2Bzkzsb z-Ok?rmSRCc0h=Cvo6!05p!Jg{gJUnt(PuULldbJZrx$nOb_!DIYM{@&0YLZx3?0S; z@&64@ctiNNZmZG>db1VcwHg_P%^GOz^`A zY^_W+codoVp^?rn2L3l(QA{Z#^TT0;BH(!spCIXKnFa4ejpiBuH<0sUwS#HX)Ft zsg@Gy%0aPvubEF|oSV_{y2bjL8yg$WFXmvJVQOk>z{uR(94aEDF7}0FFg(0 zo7}LAd^djCk9|E1xIO}BN4?jL$U)f~S8m&yo-Jij^_{(`2Z#e%`w)fy7D4lNCp`v0 z+CpY+ftRJEB_s?1mrs5c3d&r3x|)u4SSPz@KvFerSv$RGvXz|sjv4SXBoDUz`wwTi zPTl-IG)cw3xu5dT!>L4C|8{}nsdSqP&=j}j0pvXP%&U2)hg(TSMccq22@nd$#>Sq) z>n-AzySG2B>EK{Qd4e%1VeEQhk>fR|(UaPX9wW0rpOh@b#n2UvOj(^y*r&L!GE0qA zY%bl=XCB&#M{y*WD*(vd@>Ums1$YJk3bLA)uIuPnlB=4M3)o!lbJRTguvWeL(wC&E z@nP?9VWjP)ytIgm+4T1BC#zZE5)^7+`|$vkT9|HfPg_xdYRbSQ5BuCUt@Ly^57QJV zj9*f0vArA~JbFyLJ4t%RBk9PSpbSu%>gsAzVuD{lOR2p}i|mcoZV97*n0zTKaNJv2 z1)I?)?8i$^j?-F)BOj750SOd(%`ppYaXd!&m2wFE}gM zyFJK}H7#tF7}>zeAHx_K&<%cIx}}7L5Vr7x&w7ZK5U;2#*q-eVI8a7lx}x@0=ahis zP4{k+IyFCXQXBz5M1{w%DB?Dld`QXCUKw4rwcqWtF{Pmt9yJ0Ia|CvS4lWP5)@RuR z2!gU(%fvJMLg->aRy{pH|9poA!UIriAOtGzhCjs)ep{xa?jUuUbQqzh5A*nf(zX57 z@WAR~NK$eqwPZ_i$pNaBmoDoYEp>pL3q08Ksn0%O&>Qz|W6s3s#@Q7ClzMuwp5k#@ zXZ=PJr47u0hP-~ZRk1jbW-#NnifP;_%c(_Cx5IE1!uif@`l|#7LMoNW`3*EOD zm{)6+TeKZxx;A!($b2b~>nMA`H1K~m=_0;id>rx<27Y2>#Wpa5yF|9D2rXfBpF2QN zeC#?*RJ|7oB&(lLO!vIW51E_G)J@O4)_w<0d4XGJfO%T?Zvdc+Z$`poaOQb@@ZRHC zKo)C7RA^}hG(H&gYclO;Dx18O6%JA*h2GX;^`P0@bW@4G4c9;m5D!`2{%{cKXZTRy z<$q5$;IMXhEuA_K6jDAvOE!}+(^q}YF26D>w?||4QffUFqMcGxi(al4gCT>MAIP#; zJ}uHCe}CFbUnU!i zh;cjr3Qrd#JlyG+N1(z_zEcNM#rW+ocb7Pw^4}+_&<3@X4Al2$8?)C*)8XUV!32=7 zSk|0Q`PCXu0ysnNkiSfL+7Id$VYRaH(j6aA<&1qrsp?+3>rW{B-xJV>BjnT&H6ev| z&ABhJbqc3lu`e5rHqMk-NHDS>0k3T-2Uu8vHYn1m`z7TKx<6YbyovewpAo!L1td#R z(fAcGg(|EI06YP)7wbO&%?gR@;r6g-r9i>**RWi}*^yTVlqVOR(u{0&Cy8B1lDe+7~X;2u>h58`JYcU zRZ98y#u|-0ro-3wR$G(+c@DiyKX~Kh4eftRLz^}@6Y0y1PM|md$^FB$uHDKsn4!~8 z7nhT$Fn;MCY^t+6C(AJB2B3PU=s(XuecQiL;S%q$v6Gg+g#KG1nZ;m7=FNK8 z-hnW&1WjCg_Qj9`xbCnK`jtBx{2Gb*!c7k+#CzXtPEQ*z3$qFigFz(!_aIz;qOTcN z7d;3#-V4~+{M9-=nhciHC$7^a%kdA!1mD$|Yp zPal9K6Qs7D51jo>sS9Fcj@&+gP&7aRY$i zTIeuc?1zkBe$iL1QcX=~IC2&uDX`mbE^yIbT`gZ-P9Pt>ZGBMdpk`C|>pz#5lbFZzWMioM=Q!5cKbT<2y8c1lxbPGMhk;tlZ{+Orz~g%(wUJ=xnRGy7$#}h z&}^(}G~~B4Z6OrA8CMlKpfrqdk%sn*(#j=zccD$MhW2LAuP?=bmckMN?aEiG&bB zsnIZS-I>R9AOPs+pgtlaE<$&Fu>%WhlZ7} zH$0ncIBmM$qiC7ZH4R759z8?W^nab&ttK8Fpt#8dYyrL<>o<8=Q{( zj-%XE2bjX8DQ5t0#P+C{S^l^Fi%F{9Ht}m(E{}KM9tG^cE@M%fTE*8Rm)+&$UXea@ zt@T0vMi+owp_94{5H}*d`#%7NB8;;ore29yO@LiYO~XU>|Cly3JF$jd-F;;T>cIj) z4VW_2EOMO)4RN^3+6D6GEz|_;;ctzo;mNb!l-i(HPJK z#B{(x6@a~3C754GF<&~IRt3pB*aQU`XKcUKC@59XfMLaqPK{wa#w;)%Ab1m80pmws zKR-J=+iAk5N#gNA{rsNTX5MKPcCky%T{?Ces~R=zKgIlRQ}llygDLv8g114`I9m$R z?}20TkcC8w_3lW&)9tc(H6^9B4_E?*vdB@tzoL?lk8PzlrH{5R&r&Sy2aH-&rSB8l zAD2w+PJDR=xXH8Do>ae{0ybCPi;eXtwM72=puC`vJ}LiOzsw{nFMnK8^Pg_81xyzY z_L)(Kp?^bMk6Tv-rqqrbeWYopoUfAN5pyQ=-U6HmPy{d?UE~-W@&|XpoAXX}9Rf0j z2xz|nRxg|qza_Dc92@=byIT}suMhC+I+L|b!OHzV9>p>LdC-b&mU~AF@6xWzOJhFd zh}6!3icc38=3#Hp2=Us$LF%A9z^yF~%mGq^&# zbm>wbZ-)@5FU$5fCO?54Vz@J_gvDpS`K5sCZgA8F}~K(U_K60hlyk*A0IQy$l$MQ@pJKK0Xb9`cv}!PcfX_+{%Jh zeM1@dU;9*?pZ$&6XSV(I^$ljcwgovG?Hk)V-RgVbbc>aBC5*dt=)Vr#pr0PiQol`a z1x7bkpe5{o@*F1cN$U7SkvuwH1|1*mE+1%jOL9?emVs5z*Z2fzTP5iyi|lL2B0aUZ zdB>M0r?sG*_ii=;#_e9a!F2ii{=yhI6;zy=dH<7K#E6!fq{mDb4Uc{XIFw}6ib0m` z!Q8j8e2rvrw_GkRE{XKfXMVjC&cOZ)OG^O_uQ}({EBvc8H1>ezj@fFbAV4<7?06mD ztuy8=F=;kTF!3pvpLXXunLQgH*Y(>Qmd60jUbdR_skU+gFB0$^;Rh~!8^nSGI^!7h zdiC?|Cz5+Z52u>E-DIZ5ubw$!Q=D501PAYFS%37l%1Cxs!r$Iid;EAD*j-fs)*mi@ z0bi3sVY=aldMg%Cz2GhWRlb0uaQMaeNdHte0=%2u&le&@BPB%4zY%V*oSYBh#7eC1R{e2@9GJExt@o7l-gf2*g~*7WuD^>T)W z$sb~4pUwnePbCe#)XswmN8bI*ZvswiIU!g1&g=r*+$seud#)_xB$%vyHaz{3E{(m) ztxF5I{~l3ewzmLbKR+i&c73930x&Hkq9!*z`x6Boru$=T6XtG$!(o80stfvPGAYw7 zH7)?hB@T-gnUI$^QIzl6TiHi1Z!$~!H-Qz@Mj&pd59Yxt+a64|tJf`A$UZd>$RrI7 z5{iy7fl`M*2_k3(zVe~>t3KLb0h*!%Tsc@|*i;)fYo_;lP0vP5Zyn@+vO?3oJxlq8OlLMU7cuhn$QfyFpqTv*^9=fmW*m?+4fs+KS_uWJN!EB=5 zz6Hn6B(nkOP7Uxk=#7-~xWSs&JCvh<2CJab@u+tKqb&u(cVn#)|`WmbG zojuf1tHQRuBuKtv873D&mjhhXN{IAHrd%Y$Dy~TtHbzZq_|>_KRW`#}13Be24wJwC z>sjl*#7z6LW?fX%Ze&Ge<@u<*z|~jgKfe$&kG-alvhlo@`%EJ3zQ?Ru2JA`VQ}MVI zfm0jq^GY>sGzVd^ArR(Txf}U8n;hOfS(bXm6hc_ZeVSDw6m1f1v-i6Gj$4mlZp#W zaN{#_UAq8&Gm;;JYwPRP)`K4hj*kye6X4Bp$;pMl{jC0NM2+{0&KEyh=(Gi4(&uEm z3kg&Jm^1QQp)WZcs?+zqGES@Jl{P7StkfQodvWU)^RqI#H6Hy*xtqTG?Xnszpx+3B zp}YgUQ(AB?PoV54lia`P2p%^2eEKD_IS`J5!7x+s|NE(7eM9WoDf)0BDg{9V5nj3b z6548nl}lJy*s#jx!(fPAdi>aFJQBnn*riZripo;SI zEd>Ea2@5ppINP9;VC)(SqR?2Qm-8&}2ofLwE-#>(4?e=Y*jw!E+_!J5!bbx4|ISVy zJ->hdB|wS+)Rj@7mB+=#Cdb!$QsGp;L=B^|J{zCA#z1%goMBtwQRLs0l~LeLJBXYD zCM*(gp|8-;JARY!KBx+q)&cPq?Og9T8{UAhx4t0#RJJ#3mGZCNM&!Xc!y? z1ZcaCu&X}X9W>~#7x`A+(3psstHLL8#oKcDq`Iog z9i7e;7Z`t}!z>HuKg5_rJ>FTn1RL{$B((@hbPas7&?-HNTm&=3Jhn|6wb;|MtI z5c~f9kT-n%j@5&F_b)g)45;bN@F_aLP(-8VO{3knet!8hXFpT-tg8w8mg*+(D-x)K z)lvDYSFd7Qn~#YV*-00mGaU>;lIG+&@kFxgFVEXjpDS^^RQE_oHb|Z7^tFu*iQU^z zpQTqA^tCO1UecVUq}JB5mU|t+Nl5;{xa$1Gsek5GM)!FNh85pAx50H53e-ZcUvkXE z`t6n)eV+13hYtIv268Nv>V7Sajg8s$XWXs*{i=tj{v@vgr=|=a-Jt4kl4_{0zb}y% zjI_J$Zw+*jZB<##&)fM;%WDczAfsEG!6`o14oGODJWjXlQ8C zJG+;6w~kML%+4ZHFBATn_P#Vc--ZAFt^3mMk<0VvBniE5!Y*v8X^N!0;pu%V_h|)% zB8F1a)3Kh(vQD-1U0q#;pP$}2%F&Nz%>*WY`SRtC&r#gZfv&4xkrn>4dsJRq8!tH5YHACs ztE+jrN{<3)FUo*{g+9VhOG_)Xu(Uh}BpS_nU*ifXDQU*1PXvO3f}jw5J?RiU=^)2L z-Ht?Fy~jjhWo>;uEX!#Z*hI&dFL;rWkw4KspPQQzGv|yE2!w92i`YALj-brU>w9~9 zEGhkkx&(M;FJ3SJ(Sh5Pgm^$~xw%And3i1pX(VqR*f=|P{rEw9;SIhXVQ4<kWK~pESJSAeG`jWSo2zGM zOhrXSK(L_O#Psy8k&);*TPmu!cR>8$=k+f(O5CS8p2DL?kDl`LK%sBnz9kJl$3jI# e_2Eks!Sa*m>|~c8z61puq$vCNQRxG-SN|95>rx5; literal 0 HcmV?d00001 diff --git a/docs/source/index.rst b/docs/source/index.rst index 0424bcfe..6d1bfa26 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -66,6 +66,7 @@ PyTorch Lightning Documentation fast_training hooks hyperparameters + lr_finder multi_gpu multiple_loaders weights_loading diff --git a/docs/source/lr_finder.rst b/docs/source/lr_finder.rst new file mode 100755 index 00000000..aab0c754 --- /dev/null +++ b/docs/source/lr_finder.rst @@ -0,0 +1,108 @@ +Learning Rate Finder +-------------------- + +For training deep neural networks, selecting a good learning rate is essential +for both better performance and faster convergence. Even optimizers such as +`Adam` that are self-adjusting the learning rate can benefit from more optimal +choices. + +To reduce the amount of guesswork concerning choosing a good initial learning +rate, a `learning rate finder` can be used. As described in this `paper `_ +a learning rate finder does a small run where the learning rate is increased +after each processed batch and the corresponding loss is logged. The result of +this is a `lr` vs. `loss` plot that can be used as guidence for choosing a optimal +initial lr. + +.. warning:: For the moment, this feature only works with models having a single optimizer. + +Using Lightnings build-in LR finder +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +In the most basic use case, this feature can be enabled during trainer construction +with ``Trainer(auto_lr_find=True)``. When ``.fit(model)`` is called, the lr finder +will automatically be run before any training is done. The ``lr`` that is found +and used will be written to the console and logged together with all other +hyperparameters of the model. + +.. code-block:: python + + # default, no automatic learning rate finder + Trainer(auto_lr_find=True) + +When the ``lr`` or ``learning_rate`` key in hparams exists, this flag sets your learning_rate. +In both cases, if the respective fields are not found, an error will be thrown. + +.. code-block:: python + + class LitModel(LightningModule): + def __init__(self, hparams): + self.hparams = hparams + + def configure_optimizers(self): + return Adam(self.parameters(), lr=self.hparams.lr|self.hparams.learning_rate) + + # finds learning rate automatically + # sets hparams.lr or hparams.learning_rate to that learning rate + Trainer(auto_lr_find=True) + +To use an arbitrary value set it in the parameter. + +.. code-block:: python + + # to set to your own hparams.my_value + Trainer(auto_lr_find='my_value') + +Under the hood, when you call fit, this is what happens. + +1. Run learning rate finder. +2. Run actual fit. + +.. code-block:: python + + # when you call .fit() this happens + # 1. find learning rate + # 2. actually run fit + trainer.fit(model) + +If you want to inspect the results of the learning rate finder before doing any +actual training or just play around with the parameters of the algorithm, this +can be done by invoking the ``lr_find`` method of the trainer. A typical example +of this would look like + +.. code-block:: python + + model = MyModelClass(hparams) + trainer = pl.Trainer() + + # Run learning rate finder + lr_finder = trainer.lr_find(model) + + # Results can be found in + lr_finder.results + + # Plot with + fig = lr_finder.plot(suggest=True) + fig.show() + + # Pick point based on plot, or get suggestion + new_lr = lr_finder.suggestion() + + # update hparams of the model + model.hparams.lr = new_lr + + # Fit model + trainer.fit(model) + +The figure produced by ``lr_finder.plot()`` should look something like the figure +below. It is recommended to not pick the learning rate that achives the lowest +loss, but instead something in the middle of the sharpest downward slope (red point). +This is the point returned py ``lr_finder.suggestion()``. + +.. figure:: /_images/trainer/lr_finder.png + +The parameters of the algorithm can be seen below. + +.. autoclass:: pytorch_lightning.trainer.lr_finder.TrainerLRFinderMixin + :members: lr_find + :noindex: + :exclude-members: _run_lr_finder_internally, save_checkpoint, restore diff --git a/pytorch_lightning/trainer/__init__.py b/pytorch_lightning/trainer/__init__.py index ec2ab71d..2863be8d 100644 --- a/pytorch_lightning/trainer/__init__.py +++ b/pytorch_lightning/trainer/__init__.py @@ -135,6 +135,27 @@ Example:: # default used by the Trainer trainer = Trainer(amp_level='O1') +auto_lr_find +^^^^^^^^^^^^ +Runs a learning rate finder algorithm (see this `paper `_) +before any training, to find optimal initial learning rate. + +.. code-block:: python + + # default used by the Trainer (no learning rate finder) + trainer = Trainer(auto_lr_find=False) + +Example:: + + # run learning rate finder, results override hparams.learning_rate + trainer = Trainer(auto_lr_find=True) + + # run learning rate finder, results override hparams.my_lr_arg + trainer = Trainer(auto_lr_find='my_lr_arg') + +.. note:: + See the `learning rate finder guide `_ + benchmark ^^^^^^^^^ diff --git a/pytorch_lightning/trainer/lr_finder.py b/pytorch_lightning/trainer/lr_finder.py new file mode 100755 index 00000000..ff14e93a --- /dev/null +++ b/pytorch_lightning/trainer/lr_finder.py @@ -0,0 +1,445 @@ +""" +Trainer Learning Rate Finder +""" +from abc import ABC, abstractmethod +from typing import Optional + +import numpy as np +import torch +from torch.optim.lr_scheduler import _LRScheduler +from torch.utils.data import DataLoader +from tqdm.auto import tqdm +import os + +from pytorch_lightning.core.lightning import LightningModule +from pytorch_lightning.callbacks import Callback +from pytorch_lightning import _logger as log +from pytorch_lightning.utilities.exceptions import MisconfigurationException + + +class TrainerLRFinderMixin(ABC): + @abstractmethod + def save_checkpoint(self, *args): + """Warning: this is just empty shell for code implemented in other class.""" + + @abstractmethod + def restore(self, *args): + """Warning: this is just empty shell for code implemented in other class.""" + + def _run_lr_finder_internally(self, model: LightningModule): + """ Call lr finder internally during Trainer.fit() """ + lr_finder = self.lr_find(model) + lr = lr_finder.suggestion() + # TODO: log lr.results to self.logger + if isinstance(self.auto_lr_find, str): + # Try to find requested field, may be nested + if _nested_hasattr(model.hparams, self.auto_lr_find): + _nested_setattr(model.hparams, self.auto_lr_find, lr) + else: + raise MisconfigurationException( + f'`auto_lr_find` was set to {self.auto_lr_find}, however' + ' could not find this as a field in `model.hparams`.') + else: + if hasattr(model.hparams, 'lr'): + model.hparams.lr = lr + elif hasattr(model.hparams, 'learning_rate'): + model.hparams.learning_rate = lr + else: + raise MisconfigurationException( + 'When auto_lr_find is set to True, expects that hparams' + ' either has field `lr` or `learning_rate` that can overridden') + log.info(f'Learning rate set to {lr}') + + def lr_find(self, + model: LightningModule, + train_dataloader: Optional[DataLoader] = None, + min_lr: float = 1e-8, + max_lr: float = 1, + num_training: int = 100, + mode: str = 'exponential', + num_accumulation_steps: int = 1): + r""" + lr_find enables the user to do a range test of good initial learning rates, + to reduce the amount of guesswork in picking a good starting learning rate. + + Args: + model: Model to do range testing for + + train_dataloader: A PyTorch + DataLoader with training samples. If the model has + a predefined train_dataloader method this will be skipped. + + min_lr: minimum learning rate to investigate + + max_lr: maximum learning rate to investigate + + num_training: number of learning rates to test + + mode: search strategy, either 'linear' or 'exponential'. If set to + 'linear' the learning rate will be searched by linearly increasing + after each batch. If set to 'exponential', will increase learning + rate exponentially. + + num_accumulation_steps: number of batches to calculate loss over. + + Example:: + + # Setup model and trainer + model = MyModelClass(hparams) + trainer = pl.Trainer() + + # Run lr finder + lr_finder = trainer.lr_find(model, ...) + + # Inspect results + fig = lr_finder.plot(); fig.show() + suggested_lr = lr_finder.suggest() + + # Overwrite lr and create new model + hparams.lr = suggested_lr + model = MyModelClass(hparams) + + # Ready to train with new learning rate + trainer.fit(model) + + """ + save_path = os.path.join(self.default_root_dir, 'lr_find_temp.ckpt') + + self._dump_params(model) + + # Prevent going into infinite loop + self.auto_lr_find = False + + # Initialize lr finder object (stores results) + lr_finder = _LRFinder(mode, min_lr, max_lr, num_training) + + # Use special lr logger callback + self.callbacks = [_LRCallback(num_training, show_progress_bar=True)] + + # No logging + self.logger = None + + # Max step set to number of iterations + self.max_steps = num_training + + # Disable standard progress bar for fit + self.progress_bar_refresh_rate = False + + # Accumulation of gradients + self.accumulate_grad_batches = num_accumulation_steps + + # Disable standard checkpoint + self.checkpoint_callback = False + + # Required for saving the model + self.optimizers, self.schedulers = [], [], + self.model = model + + # Dump model checkpoint + self.save_checkpoint(str(save_path)) + + # Configure optimizer and scheduler + optimizers, _, _ = self.init_optimizers(model) + + if len(optimizers) != 1: + raise MisconfigurationException( + f'`model.configure_optimizers()` returned {len(optimizers)}, but' + ' learning rate finder only works with single optimizer') + configure_optimizers = model.configure_optimizers + model.configure_optimizers = lr_finder._get_new_optimizer(optimizers[0]) + + # Fit, lr & loss logged in callback + self.fit(model, train_dataloader=train_dataloader) + + # Prompt if we stopped early + if self.global_step != num_training: + log.info('LR finder stopped early due to diverging loss.') + + # Transfer results from callback to lr finder object + lr_finder.results.update({'lr': self.callbacks[0].lrs, + 'loss': self.callbacks[0].losses}) + + # Reset model state + self.restore(str(save_path), on_gpu=self.on_gpu) + os.remove(save_path) + + # Finish by resetting variables so trainer is ready to fit model + self._restore_params(model) + + return lr_finder + + def _dump_params(self, model): + # Prevent going into infinite loop + self._params = { + 'auto_lr_find': self.auto_lr_find, + 'callbacks': self.callbacks, + 'logger': self.logger, + 'max_steps': self.max_steps, + 'progress_bar_refresh_rate': self.progress_bar_refresh_rate, + 'accumulate_grad_batches': self.accumulate_grad_batches, + 'checkpoint_callback': self.checkpoint_callback, + 'configure_optimizers': model.configure_optimizers, + } + + def _restore_params(self, model): + self.auto_lr_find = self._params['auto_lr_find'] + self.logger = self._params['logger'] + self.callbacks = self._params['callbacks'] + self.max_steps = self._params['max_steps'] + self.progress_bar_refresh_rate = self._params['progress_bar_refresh_rate'] + self.accumulate_grad_batches = self._params['accumulate_grad_batches'] + self.checkpoint_callback = self._params['checkpoint_callback'] + model.configure_optimizers = self._params['configure_optimizers'] + + +class _LRFinder(object): + """ LR finder object. This object stores the results of Trainer.lr_find(). + + Args: + mode: either `linear` or `exponential`, how to increase lr after each step + + lr_min: lr to start search from + + lr_max: lr to stop seach + + num_training: number of steps to take between lr_min and lr_max + + Example:: + # Run lr finder + lr_finder = trainer.lr_find(model) + + # Results stored in + lr_finder.results + + # Plot using + lr_finder.plot() + + # Get suggestion + lr = lr_finder.suggestion() + """ + def __init__(self, mode: str, lr_min: float, lr_max: float, num_training: int): + assert mode in ('linear', 'exponential'), \ + 'mode should be either `linear` or `exponential`' + + self.mode = mode + self.lr_min = lr_min + self.lr_max = lr_max + self.num_training = num_training + + self.results = {} + + def _get_new_optimizer(self, optimizer: torch.optim.Optimizer): + """ Construct a new `configure_optimizers()` method, that has a optimizer + with initial lr set to lr_min and a scheduler that will either + linearly or exponentially increase the lr to lr_max in num_training steps. + + Args: + optimizer: instance of `torch.optim.Optimizer` + + """ + new_lrs = [self.lr_min] * len(optimizer.param_groups) + for param_group, new_lr in zip(optimizer.param_groups, new_lrs): + param_group["lr"] = new_lr + param_group["initial_lr"] = new_lr + + args = (optimizer, self.lr_max, self.num_training) + scheduler = _LinearLR(*args) if self.mode == 'linear' else _ExponentialLR(*args) + + def configure_optimizers(): + return [optimizer], [{'scheduler': scheduler, + 'interval': 'step'}] + + return configure_optimizers + + def plot(self, suggest: bool = False, show: bool = False): + """ Plot results from lr_find run + Args: + suggest: if True, will mark suggested lr to use with a red point + + show: if True, will show figure + """ + import matplotlib.pyplot as plt + + lrs = self.results["lr"] + losses = self.results["loss"] + + fig, ax = plt.subplots() + + # Plot loss as a function of the learning rate + ax.plot(lrs, losses) + if self.mode == 'exponential': + ax.set_xscale("log") + ax.set_xlabel("Learning rate") + ax.set_ylabel("Loss") + + if suggest: + _ = self.suggestion() + if self._optimal_idx: + ax.plot(lrs[self._optimal_idx], losses[self._optimal_idx], + markersize=10, marker='o', color='red') + + if show: + plt.show() + + return fig + + def suggestion(self): + """ This will propose a suggestion for choice of initial learning rate + as the point with the steepest negative gradient. + + Returns: + lr: suggested initial learning rate to use + + """ + try: + min_grad = (np.gradient(np.array(self.results["loss"]))).argmin() + self._optimal_idx = min_grad + return self.results["lr"][min_grad] + except Exception: + log.warning('Failed to compute suggesting for `lr`.' + ' There might not be enough points.') + self._optimal_idx = None + + +class _LRCallback(Callback): + """ Special callback used by the learning rate finder. This callbacks log + the learning rate before each batch and log the corresponding loss after + each batch. """ + def __init__(self, num_training: int, show_progress_bar: bool = False, beta: float = 0.98): + self.num_training = num_training + self.beta = beta + self.losses = [] + self.lrs = [] + self.avg_loss = 0.0 + self.best_loss = 0.0 + self.show_progress_bar = show_progress_bar + self.progress_bar = None + + def on_batch_start(self, trainer, pl_module): + """ Called before each training batch, logs the lr that will be used """ + if self.show_progress_bar and self.progress_bar is None: + self.progress_bar = tqdm(desc='Finding best initial lr', total=self.num_training) + + self.lrs.append(trainer.lr_schedulers[0]['scheduler'].lr[0]) + + def on_batch_end(self, trainer, pl_module): + """ Called when the training batch ends, logs the calculated loss """ + if self.progress_bar: + self.progress_bar.update() + + current_loss = trainer.running_loss.last().item() + current_step = trainer.global_step + 1 # remove the +1 in 1.0 + + # Avg loss (loss with momentum) + smoothing + self.avg_loss = self.beta * self.avg_loss + (1 - self.beta) * current_loss + smoothed_loss = self.avg_loss / (1 - self.beta**current_step) + + # Check if we diverging + if current_step > 1 and smoothed_loss > 4 * self.best_loss: + trainer.max_steps = current_step # stop signal + if self.progress_bar: + self.progress_bar.close() + + # Save best loss for diverging checking + if smoothed_loss < self.best_loss or current_step == 1: + self.best_loss = smoothed_loss + + self.losses.append(smoothed_loss) + + +class _LinearLR(_LRScheduler): + """Linearly increases the learning rate between two boundaries + over a number of iterations. + Arguments: + + optimizer: wrapped optimizer. + + end_lr: the final learning rate. + + num_iter: the number of iterations over which the test occurs. + + last_epoch: the index of last epoch. Default: -1. + """ + + def __init__(self, + optimizer: torch.optim.Optimizer, + end_lr: float, + num_iter: int, + last_epoch: int = -1): + self.end_lr = end_lr + self.num_iter = num_iter + super(_LinearLR, self).__init__(optimizer, last_epoch) + + def get_lr(self): + curr_iter = self.last_epoch + 1 + r = curr_iter / self.num_iter + + if self.last_epoch > 0: + val = [base_lr + r * (self.end_lr - base_lr) for base_lr in self.base_lrs] + else: + val = [base_lr for base_lr in self.base_lrs] + self._lr = val + return val + + @property + def lr(self): + return self._lr + + +class _ExponentialLR(_LRScheduler): + """Exponentially increases the learning rate between two boundaries + over a number of iterations. + + Arguments: + + optimizer: wrapped optimizer. + + end_lr: the final learning rate. + + num_iter: the number of iterations over which the test occurs. + + last_epoch: the index of last epoch. Default: -1. + """ + + def __init__(self, + optimizer: torch.optim.Optimizer, + end_lr: float, + num_iter: int, + last_epoch: int = -1): + self.end_lr = end_lr + self.num_iter = num_iter + super(_ExponentialLR, self).__init__(optimizer, last_epoch) + + def get_lr(self): + curr_iter = self.last_epoch + 1 + r = curr_iter / self.num_iter + + if self.last_epoch > 0: + val = [base_lr * (self.end_lr / base_lr) ** r for base_lr in self.base_lrs] + else: + val = [base_lr for base_lr in self.base_lrs] + self._lr = val + return val + + @property + def lr(self): + return self._lr + + +def _nested_hasattr(obj, path): + parts = path.split(".") + for part in parts: + if hasattr(obj, part): + obj = getattr(obj, part) + else: + return False + else: + return True + + +def _nested_setattr(obj, path, val): + parts = path.split(".") + for part in parts[:-1]: + if hasattr(obj, part): + obj = getattr(obj, part) + setattr(obj, parts[-1], val) diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index bdb23413..0e2e4c15 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -36,9 +36,11 @@ from pytorch_lightning.trainer.supporters import TensorRunningAccum from pytorch_lightning.trainer.training_io import TrainerIOMixin from pytorch_lightning.trainer.training_loop import TrainerTrainLoopMixin from pytorch_lightning.trainer.training_tricks import TrainerTrainingTricksMixin +from pytorch_lightning.trainer.lr_finder import TrainerLRFinderMixin from pytorch_lightning.utilities.exceptions import MisconfigurationException from pytorch_lightning.utilities import rank_zero_warn + try: from apex import amp except ImportError: @@ -70,6 +72,7 @@ class Trainer( TrainerTrainLoopMixin, TrainerCallbackConfigMixin, TrainerCallbackHookMixin, + TrainerLRFinderMixin, TrainerDeprecatedAPITillVer0_8, TrainerDeprecatedAPITillVer0_9, ): @@ -122,6 +125,7 @@ class Trainer( profiler: Optional[BaseProfiler] = None, benchmark: bool = False, reload_dataloaders_every_epoch: bool = False, + auto_lr_find: Union[bool, str] = False, default_save_path=None, # backward compatible, todo: remove in v0.8.0 gradient_clip=None, # backward compatible, todo: remove in v0.8.0 nb_gpu_nodes=None, # backward compatible, todo: remove in v0.8.0 @@ -271,6 +275,11 @@ class Trainer( reload_dataloaders_every_epoch: Set to True to reload dataloaders every epoch + auto_lr_find: If set to True, will `initially` run a learning rate finder, + trying to optimize initial learning for faster convergence. Sets learning + rate in self.hparams.lr | self.hparams.learning_rate in the lightning module. + To use a different key, set a string instead of True with the key name. + benchmark: If true enables cudnn.benchmark. """ @@ -344,6 +353,8 @@ class Trainer( self.reload_dataloaders_every_epoch = reload_dataloaders_every_epoch + self.auto_lr_find = auto_lr_find + self.truncated_bptt_steps = truncated_bptt_steps self.resume_from_checkpoint = resume_from_checkpoint self.shown_warnings = set() @@ -696,6 +707,10 @@ class Trainer( # only on proc 0 because no spawn has happened yet model.prepare_data() + # Run learning rate finder: + if self.auto_lr_find: + self._run_lr_finder_internally(model) + # route to appropriate start method # when using multi-node or DDP within a node start each module in a separate process if self.use_ddp2: diff --git a/requirements-extra.txt b/requirements-extra.txt index 8b720f7d..5ebf92c7 100644 --- a/requirements-extra.txt +++ b/requirements-extra.txt @@ -6,3 +6,4 @@ mlflow>=1.0.0 test_tube>=0.7.5 wandb>=0.8.21 trains>=0.14.1 +matplotlib>=3.1.1 \ No newline at end of file diff --git a/tests/trainer/test_lr_finder.py b/tests/trainer/test_lr_finder.py new file mode 100755 index 00000000..d5b19ce6 --- /dev/null +++ b/tests/trainer/test_lr_finder.py @@ -0,0 +1,181 @@ +import pytest + +import torch +import tests.base.utils as tutils +from pytorch_lightning import Trainer +from pytorch_lightning.utilities.exceptions import MisconfigurationException +from tests.base import ( + LightTrainDataloader, + TestModelBase, + LightTestMultipleOptimizersWithSchedulingMixin, +) + + +def test_error_on_more_than_1_optimizer(tmpdir): + ''' Check that error is thrown when more than 1 optimizer is passed ''' + tutils.reset_seed() + + class CurrentTestModel( + LightTestMultipleOptimizersWithSchedulingMixin, + LightTrainDataloader, + TestModelBase, + ): + pass + + hparams = tutils.get_default_hparams() + model = CurrentTestModel(hparams) + + # logger file to get meta + trainer = Trainer( + default_save_path=tmpdir, + max_epochs=1 + ) + + with pytest.raises(MisconfigurationException): + trainer.lr_find(model) + + +def test_model_reset_correctly(tmpdir): + ''' Check that model weights are correctly reset after lr_find() ''' + tutils.reset_seed() + + class CurrentTestModel( + LightTrainDataloader, + TestModelBase, + ): + pass + + hparams = tutils.get_default_hparams() + model = CurrentTestModel(hparams) + + # logger file to get meta + trainer = Trainer( + default_save_path=tmpdir, + max_epochs=1 + ) + + before_state_dict = model.state_dict() + + _ = trainer.lr_find(model, num_training=5) + + after_state_dict = model.state_dict() + + for key in before_state_dict.keys(): + assert torch.all(torch.eq(before_state_dict[key], after_state_dict[key])), \ + 'Model was not reset correctly after learning rate finder' + + +def test_trainer_reset_correctly(tmpdir): + ''' Check that all trainer parameters are reset correctly after lr_find() ''' + tutils.reset_seed() + + class CurrentTestModel( + LightTrainDataloader, + TestModelBase, + ): + pass + + hparams = tutils.get_default_hparams() + model = CurrentTestModel(hparams) + + # logger file to get meta + trainer = Trainer( + default_save_path=tmpdir, + max_epochs=1 + ) + + changed_attributes = ['callbacks', 'logger', 'max_steps', 'auto_lr_find', + 'progress_bar_refresh_rate', + 'accumulate_grad_batches', + 'checkpoint_callback'] + attributes_before = {} + for ca in changed_attributes: + attributes_before[ca] = getattr(trainer, ca) + + _ = trainer.lr_find(model, num_training=5) + + attributes_after = {} + for ca in changed_attributes: + attributes_after[ca] = getattr(trainer, ca) + + for key in changed_attributes: + assert attributes_before[key] == attributes_after[key], \ + f'Attribute {key} was not reset correctly after learning rate finder' + + +def test_trainer_arg_bool(tmpdir): + tutils.reset_seed() + + class CurrentTestModel( + LightTrainDataloader, + TestModelBase, + ): + pass + + hparams = tutils.get_default_hparams() + model = CurrentTestModel(hparams) + before_lr = hparams.learning_rate + # logger file to get meta + trainer = Trainer( + default_save_path=tmpdir, + max_epochs=1, + auto_lr_find=True + ) + + trainer.fit(model) + after_lr = model.hparams.learning_rate + assert before_lr != after_lr, \ + 'Learning rate was not altered after running learning rate finder' + + +def test_trainer_arg_str(tmpdir): + tutils.reset_seed() + + class CurrentTestModel( + LightTrainDataloader, + TestModelBase, + ): + pass + + hparams = tutils.get_default_hparams() + hparams.__dict__['my_fancy_lr'] = 1.0 # update with non-standard field + model = CurrentTestModel(hparams) + before_lr = hparams.my_fancy_lr + # logger file to get meta + trainer = Trainer( + default_save_path=tmpdir, + max_epochs=1, + auto_lr_find='my_fancy_lr' + ) + + trainer.fit(model) + after_lr = model.hparams.my_fancy_lr + assert before_lr != after_lr, \ + 'Learning rate was not altered after running learning rate finder' + + +def test_call_to_trainer_method(tmpdir): + tutils.reset_seed() + + class CurrentTestModel( + LightTrainDataloader, + TestModelBase, + ): + pass + + hparams = tutils.get_default_hparams() + model = CurrentTestModel(hparams) + before_lr = hparams.learning_rate + # logger file to get meta + trainer = Trainer( + default_save_path=tmpdir, + max_epochs=1, + ) + + lrfinder = trainer.lr_find(model, mode='linear') + after_lr = lrfinder.suggestion() + model.hparams.learning_rate = after_lr + trainer.fit(model) + + assert before_lr != after_lr, \ + 'Learning rate was not altered after running learning rate finder'