From dc4491c2a043fba6145d9664f6fdedce7f7035e7 Mon Sep 17 00:00:00 2001 From: reinerterig <54500223+reinerterig@users.noreply.github.com> Date: Sun, 3 Sep 2023 05:15:58 +1000 Subject: [PATCH 1/2] Mac mps compatibility --- __pycache__/test.cpython-311.pyc | Bin 0 -> 12217 bytes __pycache__/utils.cpython-311.pyc | Bin 0 -> 1626 bytes data/__pycache__/__init__.cpython-311.pyc | Bin 0 -> 196 bytes data/__pycache__/dataset.cpython-311.pyc | Bin 0 -> 9174 bytes data/__pycache__/musdb.cpython-311.pyc | Bin 0 -> 6727 bytes data/__pycache__/utils.cpython-311.pyc | Bin 0 -> 3492 bytes model/Untitled-1.ipynb | 78 +++++++++++++++++++++ model/__pycache__/__init__.cpython-311.pyc | Bin 0 -> 197 bytes model/__pycache__/conv.cpython-311.pyc | Bin 0 -> 3317 bytes model/__pycache__/crop.cpython-311.pyc | Bin 0 -> 981 bytes model/__pycache__/resample.cpython-311.pyc | Bin 0 -> 6191 bytes model/__pycache__/utils.cpython-311.pyc | Bin 0 -> 5442 bytes model/__pycache__/waveunet.cpython-311.pyc | Bin 0 -> 14178 bytes model/utils.py | 15 ++-- model/waveunet.py | 8 ++- train.py | 28 ++++---- 16 files changed, 108 insertions(+), 21 deletions(-) create mode 100644 __pycache__/test.cpython-311.pyc create mode 100644 __pycache__/utils.cpython-311.pyc create mode 100644 data/__pycache__/__init__.cpython-311.pyc create mode 100644 data/__pycache__/dataset.cpython-311.pyc create mode 100644 data/__pycache__/musdb.cpython-311.pyc create mode 100644 data/__pycache__/utils.cpython-311.pyc create mode 100644 model/Untitled-1.ipynb create mode 100644 model/__pycache__/__init__.cpython-311.pyc create mode 100644 model/__pycache__/conv.cpython-311.pyc create mode 100644 model/__pycache__/crop.cpython-311.pyc create mode 100644 model/__pycache__/resample.cpython-311.pyc create mode 100644 model/__pycache__/utils.cpython-311.pyc create mode 100644 model/__pycache__/waveunet.cpython-311.pyc diff --git a/__pycache__/test.cpython-311.pyc b/__pycache__/test.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..fa4ed36174d33e15ca992a5ecac0290203309d79 GIT binary patch literal 12217 zcmb_iYfKzhmaggt{VKYfmoWxRVS{lyfFA@L6F&n6KQOTajwi$uQgjv2xar1Kg&(9l zVWN!OSV+^(5}6r!HN7KGhA1A1l`;}q%@1bQnYFb0qY9NoZ%LLct+K26M>ZMlXnyUv zw;o-E@g)1G?mk_0>)d`6VpI`}= zj1xoX-w`J#$sv-|(i}rh{LVGx!vF4}BK+@}^bYxS9`BHsb$&z)`B@jx0P6-yu|+_O zSr5=4OAZhb?;nxK{s^hM#Ov&2)TjYUb;HjIzi<8l$bBM4q#Y?@0`9eo?|~=McrktNbIiH(lrCljk%2{Qw5;^mYnFKeV!tZl4A%`SCciDMTj+FB#@XQ}i z1T@K?LVZ3r>jd?6KpRQsG9>wFxVn;)Q_~_R(8*~LZs32Eo=mb_oW2`*|+5zdIS#ir#+#%8ezs;XaPx zk&^dcc)j~{-@yJWpvIQVE&ZI>a`Cp9^DPE$oCd)1M+?!mO z%mkhVwn?sS-+*w!HviB%f#=^g(8wL)PVt@KeL^%xH$lu1qGcvcnD_SC$G<*De4p8r zb^xCX)nfBypLIFZKhIs8c#585fCM|wU0B`H&V++?q@8obgvZFiT2_j?#VK;aYrb(6 z){wO8MR`)L8IYcKr`)Xj8rBK#VBHC{2DxWIr}`V*QocqkGfF3j3A{3+=Q)BcVxfdl z$C!ilu--?$Pl5k)Q#*HoW9u5Hi!#J0$@)L?Le0^r!dLg^wg+wP3Rv3ZvAv-RZ*1ws zg>>8|?~B`!@}#^kZii@DXD2A5zV@DKCrCo9Hr^2gXr{72MDxX91_qmI8+qs~$`T`T!D8GG{ z!#zD^DW{gNfK{CZ<9o(I1?Oou%VBprQP;oRp%cC{uFaD!Q zeF#Y$pQ+Z>+Ilpej52ZIcx&Nh2QUjS0r~SRAq5At@5=6s@X2cnXBIC%c>CA89_^FM z+m-Tm^EDp^ezjEBFX)gLH7%u2ghrOhB=wA{DAXCLr+0 zB%++^#-L9v8pM-H^~8kX=rqglXk4kz=+rciMi*a)9*c0N3w$kpH^s1g5xVmvlEVlu zt8Si2jB%=eluu3$w>OvwTA$~Wo@IFLZ z5pzI8*TZ-gP(}xj$3{if3(|*2nCLY%2pvAgiKb>t&>N^hLL7#9BXYMyHEb{!QuIvP zXqkcm?T=3L{4gFcTAxukLE{tn2&&XRAOdQ}yR_vxc@mFtnD*K#$9Dn8PQ+RcbKuW^ zhMxLUknomf8~?WAL5=L$qj>fpjJz{nT6w?aSO9&a3Q!J-nACq zC5Ic7aKn0d-&%N|9Bx*^%^7!oTU{mq;CEXskjk2c_i_gp-j`hE8%0EU)xyzSNGUro z+m{cO&%d{LaQ^)qC5M`4&*aN%a^%MsvKKPfGuLye<#D;`IHY<)+4>Jtv#F=y(uJbz z^x|~x`qJLqq_p$pl}_o(kQ9Df4!^C0-_E$7RqR-JJ6E#s&TN0)n3OR6Rx6ZKZ@|AC z98iJ-v%OG@s+|ugsqu)iy;E{kaQo;G2Z2Mxn zwEJbLqet4+EA2S*WJbD{kf@|gB^4?;d-_jBgg-Q2J6AVfw?X*bXB--RN;l@1rJXAk zzpYtKDo5Uw=_?9-WrOhf5C@_XRr9QB=i+trk#-%D+PYS`$8l+ZQF=z?wh?L9h*USC zR7ErAK*WwZr7E(tPp)cLs@m79POMd(SS97E*OaQ)@cTWTD~+o={%(2Ay;ud$xKwP?$#x##a;a6X@?r|cj?Z;#Jav+%?rCYIQo1X5ST21@DSZh>bf6T}<7>=stGyrpIDRj&kjV5x z;%vWwmXyaQy6VrgIG?n*&K&eSd6@(nMK;5@&;o>>zgd7K5jTfYlhRAp`zELOZ1#IlgnLbk?*Je-m zIq%Q+ee-3>A09{-&G}MAg+8(Lp{1vE-`11zu&!&km-$z1KATTOrcS?79y{_z-U91K z)SgZfcG)nxNWcqZ$`jb+ts%ankUwz&*8t1mdu=rc{O@XD8!x7>x2awf>q2$Her_>% zIiP`pw{FyJ)A-S})M_15#wgWCm35SxBh@Y$jqJ^1)sUBH+ayub%HG0k*g}puo9uG7 znpuV%n}!gA#Fst4y)LL!39q zV-TVQZ-vH5Qw+pUV}j;yHPh38%EoBBcZEwI1RqMEqsTvk*elNo%p|y4z>7-4kVO}S z!kUnTJsPNs!GDEzF8r zJd9cf9u{gJ@&RMFAW8~Vz@8#%x-|!`?e}rVARrl=gr?{ULr&LbrwpMtMJK#ke7$9^PLUpl>$f)jVF&2j!>OIf#Js>sWwm29rz`zhc z21O4;8$J3c6wtbloC~XWVL0Q%LK0%_Tlio+WD?qeU>9RV3Zn0mF7Uwy7lzg=_O4az zU8-CzdsMS5esOcPNUrErDta|8pxDQlDZO&)jsCO!y`6Ly>;O!mu{z_`1|%oURGV)v zGAxT;YS^A0xlQxiKr)Za3BMgkq(pUb;P&#!@!|V`s4j3hSrD49Oqv_Io(q{L?=`yMBMp^C3Ag{wxew3%sPCh`**-w1RkaWNKM(yZwB%nFR~r{Y^7b>z_A|1pV)kt2M!u|O z;Y9ABQq~B#KUg`RTx66Wjj_0w`Y63fu6e2?Pu0`%Z414NXO~W`xTNx9a``c({8+}9 z@jVMvE*@M9)JTDvr`vZt@Z@|;-H$G;9F$%qWetS&wdrG$i{LNaCJY{^=10<#U?j9+^My)M#=bo1w|ds9bRgCcyzgWm6Q&?CRcYU)m?=uRV{WuxUg0nk%}XE z=!S!9Wet+?wdQ748>H%9xw==W?$vUCRavD}>N8cY>`*E@GM=n+j+&#MAqjp;cS81!6FbBP`0L17<4}-UyX<@b#ZB;Oi!lN7=SCI+1k#P@!SQ55!0lserl6Iq>c8exFA|P^Rmj^%=G*Sr7 z-wnM$S_q`MzFQjPH=zb-!)$5J1lC6V(SOQ47qSETtDwttQQDLCrhREYY%T}vz}wz7 z&}S2NW2kg7Y(7{v4q$u3`sB4e`K(W@`++MJNEIWPIu9UhYv`T|*d>$1{J%r+ya;{l zEqHP8JmmsNA0wcn`Z*LrVzL41cS77Sn~$-Ac_xjq>2+}v}_cR;VY`g{1(ARfam2#lJxoPPsSRdO)m z#8?Ibljd%?>fvr7{K{<(BqNn+f#ezl4A2mQR|X>!h{H3T5F3$zk8ARft&K%(A$J@AMS0>A@wjzO@sp67sf=C3w82OL5@5Cq;D zE^6PEf!x*Q-AZGJ()fy*s0D*vs&kg!TF+`9K`D=)k99CTjz^cv|drNhpT^OE4Z^wWLVIY2}0t9>N3mQ?#up2^=K=PG4H;9lA9)M(e z0Wf4+01-4kY?Vt7$kaiFIw(;G0VM6DAN=6?F=VFPCcE`v3GqM}q7V)Q3imsE&^^#h zKa^6lcgeLhw#+?BtZ;I5w^H4m=|d=_(3#MytLK5p4Hp%l)t8j&OBnNg(J5R344MR%9kWAHQy0T^2=we%r%$HSUx}VnW$-O7nwkfr3D55z^YV@8D!kUBm zJrTe<<15aUuHSl}?0ypcm%Y-}Vd>qda+Q^8IrwH>*};4MtUphc{G{)0-(tB;?Nq3p z68>(~Aw@Y2)nWa8jgfCe6otKhSJXN1rRr?yO}1d$v|L=6@~cvFa*05BqRCJZrTq5sXmh;`b* zhgeeJhSOj$DnB=u8LRvmfM<_f_V?`o3vAncCbY0i#itS~EOIuL1CZfgF=X&y3NQmC zn*fWyobu+(r$Q;4+m`l)HCu zLu%7m3w!$k4i{~Gf&c)OZDat{QUPjt(dYX@t^XN!$_*ym$25aj43y5o=`dK&w>S$x z2Fm=mb@Sd;rk!;lk+z z-nvJjD6mxYULiIai!*TIFlinFF!Klzt0>flVPvQ}=|O9A_8Nx`bR(u`(or8Sp$GQF z{rJ@1un?Ppt#V;}dUP}nNB+Pf4&O}j*Wi?%7ULk%F09DJk?n$O;{;z0e>3G>uqBD$ zAUZ|ZNx#?G`toSHKT@t4WE2a*NJ{m0Lk0UFD8li43Z=e}2sYhl=~JC3K!NS|b@bv0 zB5_1;P$YoWIjn`4DC4jn5~;fJ?!FoVj6JMDQXKkIi!h{x!-vyT(0(pbj4iqzDls*w z`b_PqexxT9x8ZozkK2mReX6dh5r$U-I@CjZ#wvMB_2|2oCCzIk%}a0PUpe~QP`~Q zhj*XjFC_T5eydU%9dxlXqNHc3hl0pE;Ge^^DrKc<7hz z+_mLXf0vZ$4u$TJt2ikS%2XRAh796l6fLHi$8gyLb89K5v7vx%@fot{yZZ}B;%VW w>LvTvWZQ6+kd+&RN%pjnWgCP^Cdk9+*(5w!OO|X9COJvcr_@p}?HWSr)n-q$3^I(S58757B)=7Xd|G0+jR!U_{4v zP)>frmXtKiR}MgePK0*$J-`lf2U?CBBfq`ywCAkoWxi~;7i1qaMjAvcI!EInZ=!Xn8@RIr*ky|p`qd+Gxte>Yz; z-2Ck=*Ctvm-%uUbAcdA|)}1TOEjksRV3}^YVvz(K!!mc40lMf!k{^EXMd|a>w@cqF z|D$C8QnJ3X7$}PesX)1O@NS@7d34@a-u9(yfpo1OerK>4jolB+m*3AK40e!*Hkb>S zmsOEhix-Ltb#jR5WK$(-qg*9+tLZ>>B6ia?8-vOaMlmY0;?)UB!$gMW%m8%JapGK% znA=r@L}qU~NM!p`mP~?0d!7#?23-Klc!VoxJ=zgG>10mz9Y z#l}oQan)hDW-^&!Sk53U7EK(QlLLUTlIbA%`YFPZ84f*soSyFC-i0rtz38bJst2L6 zbV6Z!@M6lq3J6v3#5SbTk=+s95Q1=Jru9S}x1l=wFg_Y3+=()~7i~unV%AQ~lXSt0 zcoLg2t&UxG#!xm!E8S?QgS}(pBp?DbwqbSP_xO5Freedi2<(lhsM!vURO+_sQlV)Qb`iM=8s`Ku+XkV6>C%|d zw6$6pMzXf0)eVQrumkW3di68bPjY>^Zdg^fMklluA!V+jQ9fo%En^d09mYo>z)%J} zNHQKUe^X;DT>zIPH0Nsopo8+-?gL+0=&qflrnf)ZE&4MT4`v_X!&m*(Qjl8et{zXF z+b-_C;VYN@smtB>A19SxA{L8`nNy!SJ#8p zb$@yNZ&eG_ivL!{zgi8hR{ea{pRM`HTKC3@G|^k$uK3biAkFo|k5djR7_nN7<`%CZ z|1A;=hCWYZzY S0*2Qxf|!iwPtl9O5B?8wKv_Zn literal 0 HcmV?d00001 diff --git a/data/__pycache__/__init__.cpython-311.pyc b/data/__pycache__/__init__.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..97083b5d92a83825b0b195193c49d6f07e395f3c GIT binary patch literal 196 zcmZ3^%ge<81TLq`lR)%i5CH>>P{wCAAY(d13PUi1CZpdMv>NwL0fVsb`iUTVBgYGP4dW?s6!ufL0{Pq2P?Vp*zgsIFgXiEcn; zNq$jshHh?RaY<^CeoA6VqJDgQW?p7Ve7s&kO*YF z$B4CXmamKmeGUJ#XUNV{gQ(T|9R>$lDYbMk&Wjv-GAxJyS$YCuK;UrIBwk~q_#24_ z*kE=z`Q(T_-vQvKz&3|4t3w91B!xm0Yy8*-ESo+?LnhV?r8$KrsG$;=)}mGna$^BL zAco_y>{J82L3;sp8`SXx&;Q+m(C5^r=u`SO5~)0fXvo*6(N_R{t+%B0S{?_XuK@U3 zZvahE)5xC_0*QFW>@%8M50tR7hDBpBbqa~r9411eCAoSUIDPp^SQ7@;GzA^SRA6qM zO2Pt=K7Yal5ZNo3TEYyFu9-v$K#trj2^&C89V!K=OoxmBmFJ*!ICkaw3%}E7ROjp~ z(D>K7)@Lc<*-xRAuojOO==HQ*j0YNS%5ya8&eaP!W6j0-+=Fh>pgDT9GKbBU9E2+9 zudvNR+KfeQR@^4G=Efn!+Fs&kd%jm}%hwZB!ce1w#iNR$9yCi`I&cwoX40N?i0%3y zlyI=6c^V+2zV=YpSDv+BCA)gjQGAjKMQ5^X7$wUS#vXL(aMD#wJ>mLaQm;ssCCby@ zJjeN3qM}$z3e6l5cjx;@!P03@qMUUumKCoVa9loGo5Nu3c|H@)SGQ|A5DX{KFR5ud zU6&hM_!ZyCLkW8^f5lMJop7f&=CSiNw!Dx7s5j*L63#`}S0*?W@K1Xvv@Vs3LiwPc zr&9h3D(lFca2M=pH`%}IX{`{m5@XI%GY4&8>F}>u_Y{iRn@|T5Og9V@#`EX~MWOSE zg5Te9_UK|-`0N#(sFH$}G+iY&R%kE6uoPgJvltu2BRQQSRD^BKV;8Sr!l;eq%f+Ml zv20Z_)-S35A04DAQBchfeA1| ze4Gz*jG+4g46FGZI{pFdAw)^@0p@s5FB2GN!*K#L@hhWTP+W@|j*X3rejz-?9bg8= zqgObdiH8`{$_RnzSOjJc#d(Gc1V~8Vw@S{ zxiK!rGEgN`tw-@nAb1rBp3U*KF2TUqYw3-|$FedGUJZ{iIgn{%I z66uT%oBW^QYGW$l@&>oe#p&A5UwFXB5u^1N-RC_oki2Qhzi;2S2&Lb@0 zM%z}QFy0-ohS_b%Y{l!!$Ag?80MD&HpHDSL!xP_OZ>ZJ4V<${CMdPtJkMCvdAqQZ3 zwS?>eq1C9Gg^|D*rE`k7&-3Bm!2$4#vU)d-8DnggQ9WBnO?J5{R?j}807?1V~%1bmM$%*`Lq zwVxF@UTEjJaC=WYs4ZvvTX3kuG0uOI3-B?}O#54J_Z&Mp*#2(dI@fl#ZGaQoPECq& zJ~+}A4G6H|+VP5z0wNvX*yN9;2f;Xr!FY76d&>7xfa3G}!(hSue&4}JJP0SC8!9>c zbHJK_(545pX;7W8oY)YCz@u*Zp zH3?iq*DAjWTKPsOeyZ%}VT`a13L^U{$%Ehf9rM)unL8I28kG9|vUmUTJ2I9$+siir z`gi{YPwQ4VC*xxr&*PXwwdRzp+TjHIvu9B?#0AwF4P52eFfXX4u>hPL(`Y;#Q!NB7 z%&YbSBSj7;j{_qf-&ATjIpTh>-!TDfU{tk*WBw=?jq{T{4w%#ua>{`6$eyuD)zk~e zjIYNX<}p4T6IB}*gI$J?8jqKV$LB#U(d@Kph>fXsxEF%Z9m6Oz?sWut4gmi!A0Hp% z>u^(9wkazWUxgb@cn_(i=K_&&?ikO<;Q#=e!<=dalfdIGua>|T7si2W0gQ#{2V2Nv zXS~(JZ^hk~oR(~0OGRS$iBxN}!0!TNz5^H7LFYU1SR+1KmaGL=ZD8O0L{}>1SDo4w ztu=tf(m>Tw*sL0d`7za^SucTSFio4^AB+S9!S5G50KcER{|ml{Z))2OM^#fB)<0@^ z35DN*r~ShZu(@7%kge?7vW7G*TU!)s%Oh+1inU#~b|}^k$=dOxwjpIwtlr1g(i!_~ zUz(My+hyx^#kw6TkE?6uZro1%Dv|D6IKNyiw;h(Nk0{kgQl=TlliC(CFyp2Hv}C*t ze4H6i9enJW1|ROM%hWUiSeD7+z9>a5_sord{P5O?^X#37T)tf?-~J5M+m2IDE34+Z zZg>BxJH17&+@@4+TSb&}^JABL&OYBKyP6eOv*c=iQG(o_xtjUTI|mmk<%%|?qHPsb zJC9RK6&a>+mYTE7n?7@9T70BhD!ZC8t(~OqlwB=PT~(6jxa{gzT>X-(|8aBMH|$c^ z@_uFaG1=Fv_OT)E7J%Z8LQqP#eB)rZN7yYOV<}ay5IT7F6s0+<+NW4 zvMZ-qXF5DYI{HXkT&8{St~Oyst~&7oY0IRQ z&&ZN}Ps+mf_3E&X*Xw%GiOSqF=V!&aM0!wmwkponRb;kZHY~i8DR<4>m=$M|X~Rl+ zqg39Qp{wTl=1(Z}mW6#Q^lpjX{h0R7SKO&i*DE#cGQCrwcP{;Wg?>|_-+b(0Rt>1? zlHr+=s=jRa9^n$UWCfH)RCSnIwcuH_xwP^bnfH5A7PhZfCsM5?q9MQVzCtRK?HvlF z%H5N2A^vsz+q3UU=ZEC87nHLX{*-)}ghttON%34t^}#l%euvU8j|Ud0%AQk-=M?nW zFtJsXFQ1Y6_XRD5y)V}Z^ZDV02rnFHL-UtEyW77oEI02}n)k{zT}n+Cu>+nBbFuVk z#j_O~i4^>VePFEqSlR11vBUW19i|gIEq~rkLHU=8Bd(J^<3pe6E(jeli zDf^7#@#FzqJyXFN5;&U)efdlv`8)AO5l?QO{s$+1gP(2g?@ZEwJpM4?JQ`=CU}Kk}Qe2h#p@}1=c57g0sK{ zv)I;X*pjfonh+*=x2WfwfICaEhLcv-l(fVkKb^F`o@UW}1=&h0PtumKC31VfitA*Z zOB1G~jnIa?^iN%~+1B(cv)4syS3#5j#TvQokk`+0ER-Zl zb9KH>#6DT}$>*IDcGdxT{L+wN*uXm3vc>Y^3}Dj!Kaj8>EBk*h;Vs}dC)M90;k4uR z65ev8P|6W9uoaNq(z85}r9-3oSt-`>+KkmLhggt9pGK^^c;E93(8B&Ad45qdjnX=l zYsKG$ecDkxQp>4>#LyEmKm=4dpqe1n!0|Zwj3XbG4_}7}1KU^<4H5U51fBe5Xj9$U zfJyVON!X-1v}95?hQdJ^1pR~pE+D4>*%dJm6NdpJ4+LJ24q!R8JlEyFj-7geuZI4j zFwX)FB+NSvmCp^gQ5c8QO+*kX2}XL_1SY_kKQvQt>=R(*@o`(*{z<`+_$R|mGYP1* zn2&@&t!5JN@qI8>HIdW+4jTwp0yg*Hm*3R4BQ+xB|Jvpy=nr!y_UIu}ILS1v>aY0n`#LC{!)@ zivu}hG%%sXgvP)Xh-n}ZfRi(p5M)Bc-MdudSb*jC<7r4-Xc-9zI31z!)ykRQ+q$=`v6d{sXow4xpac%N2XAWN*qeG~M}7s_L3OH+ycW zGu8VT;K>Dgv2IDYzgKEFAvc^*8cw8o6h|F6UsW}81CQwT6}o-NEYo`xdT+{>vD35H zZXHS;%Fs3QJ$D8c_Agz#cSxz*FVp)|$6%1R4pJW97#I5Ql`WPj)m=*Uq2&>!x-aF( zK$mwbxb)}mZCu=_c=svZ?&XgZ@7pOSjB~hvMeKz2P>&ln-I;vUuydth=TgPL9bI-l zFg>(OXM)OcR^G=-!z1#(5oz?Q+!;|iBXUDjX^2XWI(UH$vpFr^{b1=$Wy?{y;V5pi zqAJoR?e0^y9GB_i1wuUa*3S2TQFf;+)t~X!r213+Pw)rmR5qp0+`S-m98)&;%I@9= z^ut=|Y*6WEWjC9$WbE$Qi!1g`l6}*Y>W%3-xw_b@BE(h-i7|932FObsqx5zQi&awXb4_L6nZ3e?1|f>xEsIRou2${ z_uX#Ay=!T|;y$z-z`y%754xlyr={-GihJ;pduYWy^zDU9$^}k#hZJ`R3zS_}eL|0S zE;Rpc`udV>d7sqYBefiRa6%fnDtRKZC!%;FAZYHWYiIb|i<}f1RYKS0i#L>uH=t2= zPblt*l;tUCt5d4kCfm0w_U#aY*nUPWnUtM9(#}(v`c04OJ67sDmNr9lUVlWXKQcX# zIyT$&q@11|PJJj@>xpG3vZL4uCXO;TcelZY<}D_+g79|@TB3v^2Uo!8jCO)7ew=aj z`=@GNb$UqGZG3V{p}2`M<&`rZNY)L6gI2>pFZ=}qWU6A2Ui16aQokSU{&)n}9e)3{ z@jyiD;qisXM{yB{LUm>nBS9__@%w?fQ_xDdI!w4CDKJuCPg4j&0f#pCTQ$5*nIR;6 zZF$AeW!|`o^n#}-Gn}T^6`H}t9(JpPZ-Wo9wW=91USpFyF{3Vl{Ba1}YhnOWth}J* zrCcPN>Khls5y6L(P`n8b$9IZaieasbWXp(H_mDQ(Ccya7dQj6^H$Mp=UK`;a6styx zqB6*Kll-0_`%Uu8Ah(ozGiZxc{AN(IRQzU8l~nvbLA6rxyIOM6K!J8%UCvvck)H3@ I=_jK6H-`CeMgRZ+ literal 0 HcmV?d00001 diff --git a/data/__pycache__/musdb.cpython-311.pyc b/data/__pycache__/musdb.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..9794c871441dd3cda44f07d9944664427a39e8fa GIT binary patch literal 6727 zcmd5gZEPFImAlL3@@px6=!>!~du36!{7OcWh)dugfaw#PIWL_s0iuNyAk&d$7v52x{u6&x3bS2#ItrCXlt_tm zf*LQs>2aETGvhY$X2)50GYNZ=8|RXaaR*I_w)@n$Q)B_UM2E=Tr^nqQeU0jJevK!N zbTKBINJmA}bqB>2Y2r?FE=F&Y#4jiOZSa380(?X%HGQ8#^Jnlcmn<#PgH)cr`MWHm z;1SzJWK?Dz>Y@$9PhtFL)j4QNu~+aYYP!NIrEoQ^Y^7eT!L77WluD_L$bJeGKdVx; zQpwt$vualH*34zIjwJjX%-R*-4$dyFInh2y-R#RcR1oeASpi~HM~W65S?7=9!>&3- zCy^1z;j^kNvM$x7a;oDDEMe)TA2*}u`qVm?`x)c$E`P76ZmTrptYu%_s(r5=y-j63 zq6cJhGV4(4cP7O=v{XMuUi5xmXWg~;V69m`>s38PavZA2+ZC0!u7{?UPASbhSkru^ zuI#VWKtDXzQKC=uTj^%&l-3%2w_2xqtr#?=x>ZgL;5|^;XRMC1(xKrjED0+83~Y7} z?2lD^FoW~|Ayp8{U$uDAL}Hr@!T_QC z2gbA$A|{I|%sVfE6~dFwI|SgTr>DoHk9FBiTQYvvWPu&ow2LT{gjZUbl90*GrDM^A zjIgel4!{CNq;J_Cuwe6nXe^da&PG%5q?A%fQ_BsoZ+L!7m!PDDE*tW|90;oq!qy;i z0jwe;Ba=<0Q|bSP7ARZ)R;j#UdM$Hw`M`vNBux7(il-Ely`4zkGMTh&vY1_yoleJ7 zCMTznB8eu8C2exjT^J!FY%okNIy)<+M3YI)nrti)pEYfACW)}&gSowUhpXFZI^@hP zkdF)#l8VWJ6k{_obd-0&&m>kkw(r&V{A|*{H)}vH3DVc7fS42llCooj5 zxiQuR=`!u)g#vhuE{s9GWo!(wg4UF8-|D?COGxfTQoMIC9m`;`%e|LDU*aig;*u0a zSYdiEUl|;_bglPBbWS>c{dh!Dj{oc)=wfW@crprvQ7`7TH<^*eTRpS)OwVmenIKt!{>q5nI%>3z=G`UXSl&@!AA0j%AQIZQeCLt+ky^Z@ogC7e zFBr`ia>GyR8<$^xboQgO#ZJBc1*0CQIJ~De{Eb@Ei))wkrgOUgyx~8u`OlYH55bHF zPH$06;5Y&K!3`l?ys`S$CpXuWFFQ10SQmy3Vff)leyDJ=)Nt^thK|P#9YslR=r$U< zHS%m6IJk20_XDf$f4=g`m9@9lPw5BVFb=$tcWwAv3Z2D~Riyh*8vc`-|K#S;-rU9f z6`c=n^7S7^9z+T+>wLSxw`+VmR5qQy59=P(%eJfAbW92s33IP7xCjA#SL&D{K1)&}Oc8Wi9aREoOG z(P!D`%x71y+HDtgm*FOq*es`V)79`qWJUWwGCR)hL2xlBa@bjB9Z++Snp3H|D+qn6 zV~dUorP{aZu(}KMrh-r9b`hN_nNxJ_<&aee1QzZU1_u(w8cr&sDJmD7q(#qtr^pXd zR#{r9f*Z6y!=1TyV=gU$MDQMyl65OP;@Cwa>w0$Ig8*LB&Hk(hD)w4rlx~@zYB<@uXRq8OBSG^F{fKKcLi@=q4F`#;gbLS!0BDJO0l(wC5(^Y9J z9XqYL0LaaE-&w0Wi-BTI(6fX4Jh?t!>Y`r^SXT@IOD#`M2w1+)lT&qxLBbPPBMqNr zq(k-XG137e>+wj71R5&5z?iyuaJq{Dv0+EzxW%diS-(X`80CZ|MT5<-_1KGAgm+NTX4b>a*?*sjq_I{Baw^!!nK+B z>@1G$NaQ2ThXgsEib+CT5z29v5KSN{D&Fg2NVtd^friB!Vf%~E?bA*G|3p8c;*eP0 ztjp3Wy~v0Zu{<=fJOuulS@OZ|FeW{L{4nAVRF}Dl?#YHZVUCP529cCrf*}W{g`QyD(A9aI2Rd}-Z|zE znWX?QPM4?chj9=+zXmUa4F>7~An#%shsMugx!OH}$94h0EQe>x&w9P~?)%!U>BsL* zYwu1M$JcD1OhBdl5RMV%0AcGycD1~}nCbbmm!o&%$xKp+-jdUaj3No5G#8KJG)I_3 z>7*d1GYGUKDvFR9%*B;^IA@Vz^B{`v7ht>;g%s%jG#Fyz`R2hZks)CgNsyIc<=j4> zZCxSLPJ&vK8T<*t*3b0eD;ON;T~4;3`gD zO@|UipiwgF#i)2^k$?eqV(2gclf^llX#;*By-*M&R#8Z-lE7nEJ!;1oHdp`(;{n*o zp(7X^#T__Ni_Mr$IE^GrK*vcdQE)Vb(XgtP*GuT1!2K?aUd|*uVtjIPLXO`~!9m1I zf?fw8H$Gl0q+zZ@C95E{2lZIf`5;Hsk7i+l{9ys?(Q?(eG3Cm z>Vh>1!HYN{0FdKKfrhUF;m3h+@ohcOZ3Mclp?=eW!i9(LFTcM<9djN227}ztX3N2q zeZ_&|)#AX)q1uC<%aB_a z(4R5<8O@*B^fy1M`>3vv(fu8UzeDqPZ2Ci+^`Yhd+#uu-jV%u^EMLeCZM3xhrfH=K z)*3pqMcKWO5&?kZ1eQLaHy>RM!}`~n^yZhLE!4W=Es93y1taug-j3t^aM4}_n%#!K zTl069+B@^^LZ9vzN(T?;?RoK`C+`7bE&Ep-zqeISCsh8#emVZ{6Mvl0hu$`Z-qz3E zFwWh8W<5M(02C62kkC>{Q|{=YyGH14Ze)XR%uf|vt9?3u!r)J6{E1S>@!U|pSLY9d zgh1||J(2de15(H0r6>CitQ;w}td8kzr`Pv=Ic&UoS#P_n?~54wBDtS#@BuA&^)J`o z`ODb5#@K}R^BH|CVT>j8>q+B!QV%9|K4tJJjZaxIf38In`dHE!OX}BC#`TmQOzC{u z;L{qPuB19I>3o~Pw`qJERQ61Dp0F;>-LA7?gAHqJxP&twVx&xF_{|vTZna>4gYD$g z5ac-@aE~Z3x%3RdK>ph6N_BTs#iV?-*KptOhwIRq+yinT$cJ_eIAl85IBb_`yarN= zHLaC4h_iQPJ8)rJgI6;e^Ygu;+C;m^eeSSEuMfB!2XV9WzEg!f%^DYH?TeghS6R`u zVgonhj-YEWE5dU58YhJDSAo#iKZV+bF4}Z}X-!09(%w!Do6xFLbCx~%9Q_5S(%{5g z*C>R&7XjNY%w(km+>l`J9ZyYQzhrS`=p5kBw^ywaxnPlNoOjieWdTvf78QFk-)Q+M zG_3Wl&ugbf;K{S)H$b8uumGI)nYKAFkN7UOizC+^#PAV2ZgL3j;?haD#gX6=iHl98 z<1wkrVY+baRKB`#%)t!|I5*R=?P<)qGAl4EVaIggo3?V*S!qP$c(s^(*$i)6U*9f1 zk%ZCk|JMM3v67=MAO6)Z7k;@BY%Rp~V5br6%sFz7C&B%N5j}Xs2*O25sqsMJ&E@fh zNXg$^pqJeXBTrcF{>TTBUq$kLIvXp+rptrQ!zzjIH0`NOa`qmPA__sP z>Bi#e$tZESJjD86_Llg1sVr~IX<0c=;TnnDoM)10F_VzaqA6&@swn3GY}sgZjf z`IRWw0{NAwpjLTGlxKnbHmM`pUe6ZSN!M>tRqz&l7Jpa4Ifkc)=`HGckTCzh&MH30 literal 0 HcmV?d00001 diff --git a/data/__pycache__/utils.cpython-311.pyc b/data/__pycache__/utils.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..ffa1a9db44de3efeb582bf3893bfd21ff5308c37 GIT binary patch literal 3492 zcmbsrU2GG{dDgq$O>8GPB!Tpx#}Y!ofskz=Tz=-n{t(j$5ZM%}t_oy2-xBHM*zu*n^NGJGO@uRa%W zCy#6dJ)sPUueU8gZMV_0P;{#Kjf?qLGliUVHIZRrDZYwdoS zuWOSxJD%uCTR^kjrn?LBZjf&*)#VPUE_KTE|4V}Wra@K@B9ZxN5o0h9;j_T&Ju)}H4 z@@z};7~g`4fLrm+(fD(#b2%n6>87#nuz;0ivVYi2$yp|zB9oG7Gs_}LYG>okbBHMv zjn0(jn92TWAb$&hc87IQ!!h-omS!TdR)8GM3+=0=i5reVoe64~6SH9R4EqZ0wz#v3H)%{H4D-cBM9UrF!H_;l0h7``7A`1Dlsh7amSmBgYCeub+n^ z#Z#LnHr>+oYUp?^bUZ)vJbbt$mOdz_s^Rfkc)TKwzXE+ovK?SBXs5RaFo5*-O;g*M z4bcfSLeN0TCSPWtkUqmq-(FlSl00$iz-IDep3ab?dZ{pmC0Kj#+ zNc8Ytuh9i`8(sLiv7ZS1xmbDYhqW{B{Z+0kYPGA0%9%vvbfR`7=^OnudYe+1u1JSz zKkV>w_{W4Sv7nX9WY_-*2;P!t)mDSj=@+mXbRH2GvSq8Swfj6-2Xw)cyxVy&8%#>9ju-BvF@|>jqpX z6s@*~VaS^&ckY8l%u{?{SzzItnhA*^o`Ud>1zEtMOv>1n&2vfIUC$a!wo@s`a9Nk0 zqdarIM46eA^)o}Piz#U{L2XB4QrgyZCZ-LG1-a)?JPf8lNr2?)uprMROm+>+u_>&S zvYNX}j{*_>Hd4agx(UnYuc zrQ`S%GPTIjg7BphF0A|0&m(=sYo)&Xi^WC2ha-iR!c;{Xdsl(* z7PcgC1o-?WfSDnX0}D)ke`Y>@{$fmI-B5WgJ!Pg17Q9VO*I>d@Or&G5f(3bW^Fwsd zw_3-;x;bV1wswnt2k>vh=coYSkffe_i<`5hlaEi8-zwj!N|QBdveJ5ez1s4?3bZ1! ze)T4p9+RCHXgXcM16lWmhDY|l33cAQ=yyG`|1E6`w%NT4miI~lGk*Iwns$J+=JDT2 z12>}mgM;g$3lO@n)2{caXA z9n*r0X(bJc-pb@99s^9)4Hv2~6O-8-9l{kSa2iWFR#k(>kU@}MK&)cJ{Rv`$$q-TS zr)Ll=Lh^O+5Ww=S(3pZDMesfNoL>OQle*HoIqG|K{lGvW@;8t@5A0_(czuwa5Xwoi;g^tj_pLps?oEx=-EQ>MX0Y9dZYAq`P7q( zPvF@i(%XRh8K2xrRk!B<_=kUny0<6U|g<}ph-zXH$@1wp8j-aLQm z>P{wCAAY(d13PUi1CZpd3OacWVq zeo<Mv>NwL0fVsb`iUTVBgYGP4dW?s6!ufL0{Pq2P?Vp*zgsIFgXiEcn; zNq$jshHh?RaY<^Cer|qBYL0$7&-6`olx$)#xOhh)+vBg;RM?J}{Xu#l6cP8~Z^?KrZT)>4p^K)YbLE1MR% zlxCMnM6V2^Fi@gU0Vi-^IxrA6FccXQ3LSN8jsbe%3PejRU_eEYgKsRwKA&MqHk-h1=r=X>+*d=&`v5GYfB{pf?gK<%H{aEiNC*`0#QJt7mCW=W1FDaAUM z<#MhhPYFjZ5ZQH`$h<W7z7 zOA0%m_I898ppMaZ|3(SoXsa;W^}QK)#4y}4KX2~=dlBZ~EyCPocT=NrZkgw-E#d}K01Aik#+|6o_7fa7IwvUfw+KMF>Q`4UfX!=50rkYwJ4lq_e| zs}^@vWw-*1`#`;VBmCzJ|0v#m0~~~0KBuU9As6el+*)BlVGMz>1WC%M8C{YrpUzUM zwvg8p%cJQmBP*7FUSX<|m9$Jzv3zOBHc4NCt_Y#khuJsl^oh%sXD*Y~6=r$R0Bzp| zmM<(M@+@cZDoX#GDLs7+6>&M{vv^I(&RK#A6Idwdn&nfad_l)53&H?w07VEzKaLkj zymRHUH2K!KMB*2hW-Om1rL!qblOzpu;Bc$l`fhgoT}@%yI8!p?Q~7iOm#vLoPNlD9 zR7JX^q?igDHGcV>DF`$EUg|?-^xe^fqL2P^Nzb$Nwb5KkgJq28^0Ja0$3>4VEHTW~ zVh~qmAGk4R+i(-L|AP3>E%I0h{4rAwo5D~{82Y<#Vp}+23bC3Hs|vAt8(FM z2!HbLs1m+M8FZv&pKQuP!W_ z@?tg0%qXixS!rrzx-|W5bnro}^43~rD_nc+;y16D(JQs+6}0X6!)1QkKUnn-)MwGPH_%o!XUUOvH) zI~bCK<-xO&%Eml)u3PlH#WN*aU?O@vhaNpyB{jb!p*H4XQB={7Ef*XX4d(`#&kmu| zJC|o4rI=hC*vr$F>LA!P4a6;CXIXi6+kd3$KT?@}_}fi=``GE~vD1%(p}WV+;V&Ye zM=DcaO@BGPN&l31m@p#~waA1SJXH&xDovIqcS1k9d#3XKc4)Nfyx8+(S^wh3=Qq~k z_isM9X$Hq?!Lib0y)RhmVTWPduIxZmapt>QAiA?d8!)I$k?aAHHhM*%e9(e7V8s~} z_Yv&v>_AJi7g#!SjST~#em@%0dMn?WwEk2ti0+IDc5dFzGsl01JaQ!%GLKESk6boM zH8~5}%K^&Yq%?;@3yj&-Zs2;aDteOzZDUc~!>`+N93pEgAOOa`(8^6Di;;R{q_TMb z(AvnxOQv{gi<;tVc!GDt;nKy1hdeh_;cJIqHp4G(#?A0V$p^Xv{YW&N@8@-}4t(;P ziu`ci497QrwHb&*na3ytrbZu&5+Kjxq6*tA#ffxG3a{Ld=CHsL!RS@^V_ZygM zH@+-__;&5(v)~Q+@q^fO;IIAP@Sh>6~nDgLyH@zJMaTOWDdL~vrAonN!BX(lTcEdI^v$MC7ktCMec#M`-}~kn)36=3tBCip3e`HG;WLt7MqWYI>7Iw8f@Q^#*Pg@uL`9{Af^F#LJ=v zI~AY<`m?U)b#cK`1En|UFhlXQ@V{@zOIxfp6L#Y*x4ElxX!GCp7moTkLB}|A>G4WKy`0(j- zZG{oaG)gS3>U9jXSXP3MvJ)L1jAs0IK{nS(9ZgJr zn)4R}HNA7EdvoVr@7mtl!DRl2n)gdVI@i7N;bBkuSnfnCP_wXc*DoE7s+~7mZNUZ{ z0U&*=<+py1Pwl6sj?{GLnOI#ZTDN;wddY5aKQncpPXAK#KUJWgIZ$Wz6Ei0u4K3xR qNbFZca{%c$xYN`eT+OYM`}8uv=pNG$^c6{x`tnujdLRAoqWXU?P6hq| literal 0 HcmV?d00001 diff --git a/model/__pycache__/resample.cpython-311.pyc b/model/__pycache__/resample.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..c562864323721ecc93f54d768e599e1763729d2a GIT binary patch literal 6191 zcmd5=eP|m;7N6BhvizaQk{w&JbMe}7Q!9<@bL!^9Nt!rzoJ-TVN!*@WH>tYTu58)T zN|{|biBC?=(Q<0=fdeU3xE$yp%{jz9T<9M@=s^!!NGTk$D-S~dAQdy&@`VSZQ1}1|NMIsp zl%a3Ss0H4Zh$YI7vJB~CBi5*G)XpFaI)()6J4mpJ>|}*W>=ujA9rzn*)FIg4L8Fxd zGmQKd_etZBpV2Ln1bI~`AqACiOyVQK{CN5@XFu2l)!%?bB9#UtsYsy^Be3sOj=p*~3M_^qVo(VdWCXG`@F#p8WEq6sXFfm* zyv?Lo%o0LLbPH;C3N4vJsW=9qTTr=EXkbv%mP7$_9Vlsmn$_%1vQRVk93`z#vlsdp zs1@WCm?S$>jnHgG(jr)gPJSj-j1g|4wei&xNIU zLLsxGyi(xP8|N_}<$@EuB#9BZmm5k%$AC#J#DzeD1aRm2IbPvJJ~%_gm^h_ph@Hs2a<=U;lMd%A}kYnbQsjc4Li4!lNBB-bh$Vdr6b5T z!|8^VFff;SLbSm3Uxt#c9lB3U?mp2-OiGwot!|^@0=8NCi(whwPAr~JgyGNwZ=+#W?69K}LCkkBhq9nD4LBt+E&i;mu@0 z1PU;*aa>O_f*e#3W4CpCFrmakA^bDi9tZ>@ppJonOi;r#I&SXA*ScN>vy!{8815Q~ z1rwm0irjS)ls_zqftN%cOR)D{Ctn^o`qFUM2tO@$zS=n?DxIe;(ZzH|d9V-M6^#jE z1gIFe{%rgbCgwgmh}0nEtx*8_NhhpYV8iyd~&Xaw13sLO?7R{Tv&d3_33@8`P_52W%e)IHTRyZ zdr!)Kuc7(c#MccSs|_6=ZO;1l{iW?w{})F-fAQ|L`r7O28)wxw!|KSSHZYYPm{Ow& zH94#HUelU@4<2!^JM=jseoGc*)zXz}bC`cUxDM2C|L!Dkb zw0?k+%p`3lnZWo8+$3q<%-U+P*YI*m=%s|U5s}RchG+A_X;J%Vz0OZHp9`%XbVkkolDvX*I9eg4lbDP1KtK@hB|4}DH{&l3ZWEw zmjqXk(=o7$q`ewWg6!$LmJnNVB35ZH2_i~%gSdUR%(WNR1&)^(Ax;0w8cE1Dgb683 z+&d;w0D(4QiWKwwtbq>{8tmgZ23YT}281J^LIHz)Tn<`1OztX8+F z)or=P=4|8Bi;ULjU+mKwI#b8K@zh;=D$S&iWk%HI&gFqmUQ#`Mnx`-8=}Yye`qw;l zss434s;QlSVI9>}jxo6wn0o#4a{VXG+SUUr=e4beQ$smd-Fz~AVb!%=b#2eJb}hGO zcOB7Mk9>AoYaLwJbk9?l^|WV(7u%PfSrI;+Qioqr&+*xFld5M*^GszuQ>p%(Z_5o| zru)X0#Z8)Tck0EQ%a`uEayfN5*VvlryLuTAs&R*EKDh>8x|RO=(&RU}*m^y>IJ$UZ z#eI87-P)@*^=2FPFR(wX)%`4Ey;-f*J+pXzweIJt`H*#Yq@N>wu;>$a-Rkg2cK9vT z!)qQs>)~PPE>EfolT9rHAY!N#fYzWu07ziUqFPimE4&-jy%dDVT!fp+9Aae9OIt`Q zg$9N|1GG}GqOC9r#zJwJa(etM(qBf87#CA&T#uB+NoPHV9~cii^Up%qa7KjCPl8AU z0uWGK5;zJFa!NA&{0S?UxDHV4e4t2zQL4{=8d34wcz9Zr^6@4K2@7QN(Pq&;O5;r< zJ~bj+VjWz^aFCDic&8D2=EG4A;-i=(c9w`!66j$nH-uJ(%!E1DAyYENQ4|>{M)i9A{XhYg|+U#&%Hao@pwv%nf}9$>a=Vb7;K=)ixOlo+3&P zRtg4IRoPZg;W29M{>SNk zIe9Rxfoyw^*4VRrUTfT|*6v+vY*lMppP=@C1MxxaOVnI;mnd{W-HQr6YoY#;V1Gf~ z!6O^S;G{mGsDMD-Fb?X$p=KoTL6(Bi6nQ-4K^{2@I2lF+G~uJO%v>#?0GT*I>gvEFQm0j2qSSh2Fo7}EZ<(AG_#{L7HxUgln(mBm zgYp-ud7oclH1EFD@zn7(@3z!&!|xQG zmp8AiUZSnuz2*S08%yQE+~%zr_tkS5;pUW9-@SZ#wSKp1KDj1hMxBcTOD8nno|SIR zdmwc@=klhzugs?C4=Sp|OV!@b4$=>yLAZp0j~fX-K;Yx1DiDB}ArT?vnn2)uf{z$I zG=|6Ing{ogY%gTGGyg3S6eE#9z|Y`TXqAcH9;W+7Kb>esVnaC&83g*~Z{5;l8?~&X zC+2&-wz_px$a>tiwslm*z@tT`Lu#%)OCVlAQ2^daxQ|=U_R+%JTM_zn&qN zHsG3OEu~jZMOT+Z3Xm{U(dw5rT#}SW-Ed6>J9MXD-2>>{GKo{EO@lOf#ApzhGU`^7 z4?&zHLvuE1Jpj$~t!45RqTioDp}>`_lh7=ettb`khp4gvlqz_^(K1=95GB$>I1NxCvgmZAQc~8A8}dj8~r9rQ6`@Wn9!P z@vzRq4S>!Dp_i2tQNPXjaOcP@=A`j?&V*%hetfEEX;X-vD{ zJCx>Et6Nm_$vNs4`1f3Mecw1;^RF#TedJkcQk^}TvnT89nH&7Z=~>vXIa{*MmUU#a zy~!APipf>iEHDe*SFKm}rS|2#zTY3YarAm+W>E9)NDVG@zXwQl&)0mTA;WxlpW8p4(6~dI?{L<4__Myd%-!xUt$+Q6<~y0&>`V85*miUKV(Uj^ zOX6++GJm@hK*@hVYde^2JE(0rsBJ#@ne_|&`Se}>?`K}go*DTvtd73%WmG$MHhb)> zcKEEe`RsQr<9m}a#Q?0_RR8|iwNUz(SEXx`=4_uE%vHPRP8yh2rQ7J;2|WqiNOq1` zx$dAB6wgW$CMShENSpm61v&$x$R^)L23WK~`A5eA0|#HgC!tD=Ri1}zon;s%hdiqJ x7&-6`mz`xyv6#Qc3+-mXuZe6I0O-?8uTWMR63Ja*eQz#Lp^u9X@nHxw#5|u!c z?wgvVB+m+(^hjQyUWt>K6>5_9pcjzjTR{>lQ?upvXx3Nm_Ldi+JMeeYq)(zIP}Khg zDV~ngrE2mbEePLTpNEXCQ>4;O2GXIfoMtQe*-F2+pzcnk96)%{gP>43!~B)gtUIev{O(4e zkY2GcLbDYs8Yov~#j*ql-Gc0$O3|Rc6_-cy-1LH_Yi5vwH3#d&dm58Tje6~dWJ;0~ zf~JeQEJ#XRwApQRXY?Db@3Gui;0As zuzk?$LPEB^30bv0Db1$mMSaF*B?YTuQnr2aMMcv!n@fswGE@SCCN9c$U|WtTZ8MH+ zw!b0}CTkf#21Fw(gopI?+5WSdjJ19&EB&!ld?6{Ty4HV6jL#^lES!`@tSahs|EclV zv6B=1ufr(4XM0b}dhaWj9O-+LqNdBZ-x;faL01x5-~1(z1ULn!61TuuWIXNfK$g)a z+qk;tmrXZYtsQ+P+i$V`MMNHvtTrT;NLZ`txfq+!(($ zo{O20s1=EtT(`w_=eh1p+Q0Hrp}lLPy=T3>$87Jj+WQK8bb}vQ=LbxF(BcORjm<@t zZtxWmAspgDU;1gL?kfZY>Z-|I&6G-`HtbDdeK)P&F~oBkEz;anx_2-HbVvM zh?M1I(2{3Ixqxb+gF@<$b|F~gUWrbyNiNM$$b$@~b_1n-5|d_Y`^h~NRBGGB(Dw)5 z1goUEwBO)z6=0|c0~v#@l?b;OoK6-{Dc%2FLT^`TljaP+;WM~rfcwhOBQ4;o9fLyW zkm|kYJ%=t(kAjfvLc8s1458AX?qngcXg*Dfgf3?`6kd-PN^ zh$E!KCPE0v*gi>KQ~>7!@fkTjH=j~e-R-5>0fwrPgem2sEZM$!EJMzw=4@)w5nUtV z5uhq1+)O$SID)$95`~IJ_-k)~Km2d#^3TTRwYF^R_F}$q&}OEbdsIJGR9>w%QFGgsGt|J7hKN z$}kxQy82#gwk>zu+VNDrp(n!>f(^?jH8RP6**VsRK5H8rKHBU(9XurGH0sa5jse$dwJ-)YXL6Ouu`(EBYR3b+cK}e5KD=b?ZMgX+3p*AF zTI4RkOq~;5`QIhxwD$h9RRaXfnU=R zg1}1DZ*!VVY^2L}kc8RTjc09+Oxy(#OEyg$81YOxnAm3G2{63POo2nuZF)|=q&fbq z6)?xvv z|5=A*Vuhx*)w9{b+<@7%&uZG2IbI02uO806mW$nd^Y+ju10Ro?;loz=a3;1@zmvEk zzR(c4QFpEGYW-UMYW-FylBIv;|AqhO!J9!dv@3I>5NgSSPioJ+yv6U$=^unP_@~$T zr%nE##UI?@N7nfflOMJC(ag(P+6;B%LmgXO!}Vj?r_9J+E3((*_GY*Yw^iSep$pAB zSLN*StI4$_k^S+F$nN#XZkG>a@jpko6x0YctiHJVVs`HKZmYBZla@~-f9NnHN36&Z zGjt>q1L=b;gtr|S-#_*4sdvV&jAzET!YwxrUpsttWNqZbrVq8-Bj%Gs){{eK*CDIx zkQqJ%Vg#GBTt3L>X&wW-qB#63sq7){c0<)=@M|)Hy9RR8z+})A1@{g9Nw1TmSGY+! zgDz5&3}BJZ39Q^BVvU>hK|hY0`~hUHQ?;jq>Z=3;)w;KbyESl|5c?bSEpSA4D$d9~ zqJBi3hSMpdEzHZ<5y$2PK~WW55TZVa2j#?+8~Fv?7)?+l&8~A-rcesoypoc%{elor zh?*t{8WGsS!%$Hw5ToCiO2 z4C55G+0rM+=ut0@krp&L%;OV~)rgWjUoz}VP*`fc&#;xg&lCO$AYj{~$LsIg47Xk# zS{vF3?_CeWmNdgXR=6i0>bbXB-$6D--!|C{w`B)zj@+iraIY2a&4)@)(F<%4NNV-a z?@#&N=!3_{eBRF(3Mkxml&(8`WDyaWLBW7PPWNz00r`+1ytN=E+%L|uz(fVRuC#XI zGGL<+rCe3XEIsG0pa5i@)k!vl_9hU}sr)N?C}x5nz&BAJ&47B}jyTK987@j4kFuED zQn{xga=^soAmL!`_TlIwSR16496G+{o^cg^78BzWEwmLNU^4XGQIs-6#kNAm z-i5cJQf-%BI1LwYQSKse71#@Ci_Jm>T2zZg0r#)Y2os1Hz^#D7MN<@oQfGl+|7iQ( zH=N-i*-5r(KaSpf^WEqBe($|ce^FLurNI5`(tmOb?G*LT_@Oyd*7Iy0JReaU#nEAE zNPW}e%n(DPp9verjYCG_H-t^&<{>lj8^a~z><~*+40V*^Om`^G9AvIpw89@4DC!~n zscy(hQk6g|_64a*Nh%AZvc4ddjib*|p3={;=zb5alm>#4C?E6%_=!pQw4MzL{_)9h z(93;7E0#kOk#|n{ZwC25wk_nvWIvw01fKVy_d#li=BObiL0zRe`VKo};FvqqkdZS0 zH6^GU^pF{9H7bq^lce}iWaMCYA}|)vy9_0vA3Ocb2^D-uzi)Wo_e{h{zyZ#G(Sg-o-xKuGbJRK8FIm&>hg0%%d}96x~A7ySkjDP#>mk# zB@bYkJ=B&OH9{v!62^q#s(~}#VWSnfa$;tVo`%Kx0LJQ}K3;uAy`P2*g(;w2`V4%1 zyJieA2Fv)`)|Gr`ts@%{`^izTfY|=ev2%*)YUoCg``iKk^RQ?wQQjXBCMSeoK<_(5 zVb7}I@7Ycu9~wTQMqu_XRmbW46*XgopMj%hn8LBX%golCqduW|Yypo`@QF!by0{1bLx{4~BaB zCIVCAu-}B9)BeC{C=&FY3i|m7?1rAxXZjAGI@j}t|DE8L3tRew(Jimvj85=@(JkYC zAsXa+#wWO7xaT^=Pep>!?#Y|Kx9$mtglJ%5d~)A(SJB4m_W44QP}Jw^-UIv09~Sn3 zM@#L5#tHj@e07^zr6})FCZCFvZ+}|9@bY3v9JnaHkmXhmj?=kGHPQ|zRPdy9A7jKGZFNwa(;*~L}H!Syt z^V6t)9t-l6DkfoSGRX5Rd?~Cp%ses$#XL2MOM}OKuP~8_Vm&>BR~0q_Q;Ca8u}1VXd^J?T<3uWEAhwbzg zxZBhdw(Lg}si?%Z$!yy)yJd;pBC$O(+at0)8Ao-B{$M&meye86=16Y(U|(Y2lZLkR zn{)j9&c*X$!wIS3gxqi<(KmZCaq?+HYx>yS=tAk@_Z|<5SFVZ;BT~bN+%SUCPg~kP z-U2Id_u`=V`mlKEE%9`2FY;dd<6*75OPTqTdK(w?*8c~~gd#PGN{ zaYO38DfiwKz0=SxeMv#d>-s-?=Us@3(+9Ec!^{zER~%Xbuj@2v1isYV5C&yHyMRKw zkWzHpacmb*Y?oUk4_TD~9V@`B;arnr^xwhr5wN^8yw|!IU3e744KeDP&OZC7OJ#8* zL&c0SLvgxUdsIJ?U}PCcQE(8)O;D2FA}GnLmK4>G8aVj$5@M!;(%{f7ENz70=toRp zSW$~iv%N7B91n(m!mZ&%FmR!9kC}Htw|O@bY!1H>h+-mKRWStoQAp8>LG%rU)pn#} znNlH2kW-AfI~5i<>=oeEoPgUJb`dzTgOrl16Z~~Q&rR1B@1kr(17r~X1jub_#aZ{! z$)B7|(|7yt^(QRwDLXf>QnaO$kVIdm&XqEz27l_f??@a^9RHn}a&Ac7n5&nnw#il7 z5-e=2>U!DPg~>0{DurCNjmhYr9ln1&edE){g~7+HhUesE zisp}*tt(nF-!?acM|Ho)NM>Adj0B_Elb9UL>Jq@&hZJPSt%RdyPwl`I&#URYkVL?- z(C3&(>@dBdcz%m2@L;|PAh)TEv+mx>WoOTlvuEBaIbW5XuO=+nxg+zv8>enJneXMg zuBE!Jxk0IJmt40CXFY$CUzl*kQ75u>>XB4}$56yS9`yMXtIr2O0gfE-+kC#ar~F|x zhDUbKZ$pAzQp&Q2Lm(Iq`+R_NNFQ*nUsI04HNP(NNZbc)`3tEMPz^V7@cpt@G-yZx^)c>%Y|7Wvkt(#3C<80<(vc!<80I7QINumo0iZaK$E=<%|95!TW8y2Im3$=< zgmDDl@$KMS7X^2(S>rE6{jXs;IUMZ8{Y&XoR@X*xh^KpMJjhbHNO(^K|xQIMuJRa*j zj-<*|5ry*|NScr|Bf$;KZ$RQkf@hB^BCs4VCL$!XKv))-UtSB$8=gP8yD;M)!vr8S zSF!nL_6=#L#CFJR$1>Zq#P&$6S7yB;>n#B1D~%lp9qTeRu7q*66fj0>dlqXC6-j&I z)-Yo2!bP%nF?@DC27&LWXP00U};7vGJf2MGcq#Q;i;ECbS@F98zM^=sE? z7Wkfme8M;o0KQeu1PiDr3wJwUmMk4v;F>yi_q}`XAua{r1IW}s(h-R3!gMZ3S6P#6 zOA2>A_dJP1i9?vK>c`_BjHeDscDHPICk!i&%H**`Tx4A@n1p|Ww)|nZwm>?`lWPix z_S+-gqH=M$5cj_=H>OXoiN{G*>--vNbGkW;ym8`5a7xKveNsbF|FM34ng zgouwq__TZN)E8EB5mE{70#PCH>1F5ECFj=pM#;HXcJ57Bo;qEr^K#P;$+=T@?!?L7 zMMxGHZ(gd~Ja=5G+bP%WL|~jvzFl(ekexfeHhCr=9`i3_PhLUtDv}%$#;)*tkn9B_ zdPYFedq;${*3zKEzN4AvHaKiwI!zGtt2CtuCZl@dj} z>5c^y?Pj8A2jjqyiDQA9IWPeXm2fuBdWRljiSpe7!f~0xqH+Wu$;J~AQeKkbyCF~K14LceZ%yqj8xNa(R`pHpmJgiXIRfVap_FFJXX$`vM`>B z(UFoFCRSdr7ZR_CIj-rgN$0_QJqZ=XLn`ZH6@~fn3+jg0T$!*QN@LcUk{J5{R>ebo zWfafX?Ca6I-hfcCwm-0g9_lAA904bxKaOkqXgsIA+qBk#C57S4v2q-ny!J6%$u|o7 zru{gz{w#tm$&^5$kp7aEIE$TX<_;*|w$ zLXJ!vjE5;V60u6oP$bAXd@e0*%mFi!+ZT4^<4Jlp0ecg1_!<{87q%P?vA`EC{h$k+ zoZ=C(5Wc7w4xBp02-u$Tj7_TMhL}j&HD>=pnQ#^=6Sf2S>NX{oPiCLDd>_`Q z_RKMI33rdMblXi#B5gNrk) zSy%c(fv959EMtOVR*f!15L2p9-qI97zAI>S7Ld@c=4-nnu z+kYz7ERzw)B>WPVzb;(F^wSOZH$X1?N?Vs${c>_RIXu5( zepAA_($Ka_)mpZGfh2L_Npeii&&Tnf-TL?z zeziTVZb)svzkhB(YV4M(d*te#Rmxa)3m_lTB%SHlH0S;K)%4YrHD#UeT6p_cJxS9_ zd5v8D(sH?XsoeX-*`5x<-r4xbw-A2peLV2@+vI&`h=?5&N5W#{x_C1tIpeZ34%VHD zTPT_W1zf9B5YB5FQ{EJJe`21Nn!QqWuUy@`N-<^qAWtKOc3_pA%jG>w6F`cF4z{&KBj+jUfzA;@hNd|SiE>i8oVT*0^QD-I35*=DGCL35Y3M>r%-IK zK8>QKdztKvp_%Ezi@PMK-zdz-guN` zTk-UOklmkpMY46kv%=bzS=SOQhbB;RY?axqqV}v6{}2%pJo^K{iYS8vl7b-B!8w6| ze~l|PZCr%C0QvND9{d_Ir}9l&99?kg<(WeY&mBY~fCPaSrPoVcCq2P1h2e0QC9e+f z>-iB+a7KbPv)C$o*X70ONk4N&oOMc6Q{?J_KIEs$Zx3g_kN6p9fH-p2mh?|IUl4%n zy5I4KL!28K=tOub8k&d*?(3oGsC(MxR*wU>Kf<}ES+~Z}$nlJG0Jn%LH8B-9Y1DU> z2h)dAo(%)`oee3=hLCtcv8f`k`sFF%kzDYmFL92MwtntS&!ZgW;B2mQv zycP`XqyEXD0B3V6#uEmKEy#W3jkRzs?a=iyk?V`Q=)2jLNzw$}m!?O3VWIz1aF4@2L+2BR-l3%_w zE)7QH!H8s!h;P3u+257en9RmRHnvjkyyr-_$t}H7dGEYma_n00E;;sy+LP(o_UP84 z6E|JYfcU*Pr0zE|TfD%dqVoPT(v~yg`3q9l1@X;`@|&Yl*J#GGMfSX$>Fj#sSS>Ns zzPw5Sfp(TA;J0?Sq2z;W;~CBsT!{mXmj%{G8ot&#FdNf_MqM-*a)v1ywa z*MNZjs=6j&T{RjlE>J}+S8iLX+%|tgeC@1Mc}}i8mpGJhRHoWR#|C&ZwGByY##WQs zICCa(V$}k9)KSuR43r4yR~@DQP+)M)z&j1so+HJrsWHD_-?ZMwSV@c-!+F#03>sr4s6}n1qPnC_QEe1o z8fOPVs))1iu^g;n)pwn~PTepJ!W|D-#|44}E(;XGpJ6H4>uJU~9i!p$8Qf?1gyLfm z0elCg9EBbe^0^K{`7f`(e2aF)dUwEn0YwUUqX~e&`=I)y-%eNPA7tecG3bcL%D;sz zFpdRp!nL9x0QDduao!K1#GK3@Ltg;NQ6!gvDE0#a+T3vPFG<1WAQ1*kZtTE2{4`tNohm zd(^+&wSTE=ztnX=?m7UWlKY_SKDg+W-G}9hBMI{=Q*Nox)U_mzKWW}HXO=r(lbR36 z%?AJpmF*?ut~F^$UdXsMF1x%-F7LcYa_x~_dywu{=6sdA;p#43+y^L+Kh0oq8~0y;>IWUA_tJG39dd@zz&>ftTmLj3i? z&O;lh-)^uT+HUx5j|phO^&;Fggv%D?C|1gv$y#~@yHg}X3Lk5Itti@6c~JZ=Oy zhskNxfH$?rO<*@M6>|u}v5de`3~91gwc54!TKRvr14`Br3Y1(HFM*Ovic1FGX3sYn zL$2rW8=Z5qS{va^gFD8B8lIwa4>>23)^(Cvd-1)tB`~gl?o=n##*FDaUaN&8zT^y# zjK%fDjG$pKO@C?Cw$%Z5)?VYDgnNQ+xIH*AH8SegKp$5>TOha=lE1O8iI@A zd4Ym^JTx*Ibzcd(r@)@j<(`NHu|#(eZnwIHi77r1^yrnng5vwQ`^GEo<9MAB^4=(p zQ{~xL+;IK%`UF4bR=?fR39SlJI$WUSLILC}x%_-EI>kp`ai1fVx`RSAG!7g}&{Gh- z3UVEAtADF`_Int`Fdm8k@!#;+_=^xrWNtnRKDg?M$d7O{{t8AB4tEQEg!Kh65ev7E z;jV3DL@~j2%kfE#&k-vcvN*+}T`m&{4&i@*flxNw0VM_&{`(NH4!=|E+3CsN>6xz8 z*?nQeBIc)HzrlT8J5^@?ao-1hN$0&r$<`#>ni6!zS(C72?2d#1z-7WeJ2GRRwI}Q= z)r~~7Wh@{h14$gt)HkQ9Q^HRh?l;VwNE}L5XKdBWHrJBPmAW9=T4h_SNFKbW)F?ZN zG4DB+I*=9{+tNqmwmnkAUbz8Trv=1UAlZ^DYm@CMLuxS1&mH>N4Y_HnROyu~y_%J; zdfDz?vb)n&lD%EFw}WlaQ8gRK0(vr))gO6&;z=9lT9Y2BvRkg~CO9l%T?@BxdlPnb z2Zjelz#P2d=zM^LrhpvF%2&}tI3B#j@vQ({sWg_j#-VtbV^py`AXc~$L+~FM1@I*= zS0Ss^{xv}p-@Zw`oFUh%!g}Mz5J0;8YyQw)Waq>c)(niyLl~+V%+{d1gn3dzxi@Dw;q&{B(Ko zmik8s75)hb0EdRg`jmjcT!p`r-rT?+t_^O=op=>h!B z8Rj^-d!N*~Z_%*m|4qr`?efWCY5%a;IV`ph%Ns8y*_89HJ!#KWH73rGt>Q7Nwm_o& z?pJRGSl6|eqL9m@$iyQER1!77@asP(VVK2PM4`_mD0(cGEDrnLw@L4r>(J)$TN%G3Q&=(l~WiIj} o-HnOYk`!~_@&$%}nTxy(7wTyk>h)=YpS}47iTm<7DaeTY4{Pc`ZU6uP literal 0 HcmV?d00001 diff --git a/model/utils.py b/model/utils.py index 4c92c95..bc7738b 100644 --- a/model/utils.py +++ b/model/utils.py @@ -13,13 +13,13 @@ def save_model(model, optimizer, state, path): }, path) -def load_model(model, optimizer, path, cuda): + +def load_model(model, optimizer, path, device): if isinstance(model, torch.nn.DataParallel): model = model.module # load state dict of wrapped module - if cuda: - checkpoint = torch.load(path) - else: - checkpoint = torch.load(path, map_location='cpu') + + checkpoint = torch.load(path, map_location=device) + try: model.load_state_dict(checkpoint['model_state_dict']) except: @@ -32,8 +32,13 @@ def load_model(model, optimizer, path, cuda): k = k[len(prefix):] model_state_dict_fixed[k] = v model.load_state_dict(model_state_dict_fixed) + + # Ensure the model is moved to the correct device + model.to(device) + if optimizer is not None: optimizer.load_state_dict(checkpoint['optimizer_state_dict']) + if 'state' in checkpoint: state = checkpoint['state'] else: diff --git a/model/waveunet.py b/model/waveunet.py index a14aa55..2391f5d 100644 --- a/model/waveunet.py +++ b/model/waveunet.py @@ -100,9 +100,10 @@ def get_input_size(self, output_size): return curr_size class Waveunet(nn.Module): - def __init__(self, num_inputs, num_channels, num_outputs, instruments, kernel_size, target_output_size, conv_type, res, separate=False, depth=1, strides=2): + def __init__(self, num_inputs, num_channels, num_outputs, instruments, kernel_size, target_output_size, conv_type, res, separate=False, depth=1, strides=2, device=None): super(Waveunet, self).__init__() + self.device = device if device is not None else torch.device('cpu') self.num_levels = len(num_channels) self.strides = strides self.kernel_size = kernel_size @@ -111,7 +112,7 @@ def __init__(self, num_inputs, num_channels, num_outputs, instruments, kernel_si self.depth = depth self.instruments = instruments self.separate = separate - + self.to(self.device) # Only odd filter kernels allowed assert(kernel_size % 2 == 1) @@ -195,9 +196,10 @@ def forward_module(self, x, module): :param module: Network module to be used for prediction :return: Source estimates ''' + x = x.to(self.device) shortcuts = [] out = x - + #print(x.shape) # DOWNSAMPLING BLOCKS for block in module.downsampling_blocks: out, short = block(out) diff --git a/train.py b/train.py index 391d379..3164f6f 100644 --- a/train.py +++ b/train.py @@ -22,19 +22,21 @@ def main(args): #torch.backends.cudnn.benchmark=True # This makes dilated conv much faster for CuDNN 7.5 - + device = "cuda" if args.cuda else "cpu" + if args.mps: + device = torch.device('mps') # MODEL num_features = [args.features*i for i in range(1, args.levels+1)] if args.feature_growth == "add" else \ [args.features*2**i for i in range(0, args.levels)] target_outputs = int(args.output_size * args.sr) model = Waveunet(args.channels, num_features, args.channels, args.instruments, kernel_size=args.kernel_size, - target_output_size=target_outputs, depth=args.depth, strides=args.strides, - conv_type=args.conv_type, res=args.res, separate=args.separate) + target_output_size=target_outputs, depth=args.depth, strides=args.strides, + conv_type=args.conv_type, res=args.res, separate=args.separate, device=device).to(device) - if args.cuda: + if args.cuda or args.mps: model = model_utils.DataParallel(model) - print("move model to gpu") - model.cuda() + print(f"move model to {device}") + model.to(device) print('model: ', model) print('parameter count: ', str(sum(p.numel() for p in model.parameters()))) @@ -75,7 +77,7 @@ def main(args): # LOAD MODEL CHECKPOINT IF DESIRED if args.load_model is not None: print("Continuing training full model from checkpoint " + str(args.load_model)) - state = model_utils.load_model(model, optimizer, args.load_model, args.cuda) + state = model_utils.load_model(model, optimizer, args.load_model, device) print('TRAINING START') while state["worse_epochs"] < args.patience: @@ -85,10 +87,9 @@ def main(args): with tqdm(total=len(train_data) // args.batch_size) as pbar: np.random.seed() for example_num, (x, targets) in enumerate(dataloader): - if args.cuda: - x = x.cuda() - for k in list(targets.keys()): - targets[k] = targets[k].cuda() + x = x.to(device) + for k in list(targets.keys()): + targets[k] = targets[k].to(device) t = time.time() @@ -139,13 +140,12 @@ def main(args): print("Saving model...") model_utils.save_model(model, optimizer, state, checkpoint_path) - #### TESTING #### # Test loss print("TESTING") # Load best model based on validation loss - state = model_utils.load_model(model, None, state["best_checkpoint"], args.cuda) + state = model_utils.load_model(model, None, state["best_checkpoint"], device) test_loss = validate(args, model, criterion, test_data) print("TEST FINISHED: LOSS: " + str(test_loss)) writer.add_scalar("test_loss", test_loss, state["step"]) @@ -176,6 +176,8 @@ def main(args): help="List of instruments to separate (default: \"bass drums other vocals\")") parser.add_argument('--cuda', action='store_true', help='Use CUDA (default: False)') + parser.add_argument('--mps', action='store_true', + help='Use MPS on M1 Mac (default: False)') parser.add_argument('--num_workers', type=int, default=1, help='Number of data loader worker threads (default: 1)') parser.add_argument('--features', type=int, default=32, From 5cdd4409f59c16d043d46d382758ec7a0211bfb9 Mon Sep 17 00:00:00 2001 From: reinerterig <54500223+reinerterig@users.noreply.github.com> Date: Sun, 3 Sep 2023 05:21:07 +1000 Subject: [PATCH 2/2] Delete model/Untitled-1.ipynb --- model/Untitled-1.ipynb | 78 ------------------------------------------ 1 file changed, 78 deletions(-) delete mode 100644 model/Untitled-1.ipynb diff --git a/model/Untitled-1.ipynb b/model/Untitled-1.ipynb deleted file mode 100644 index c7c365e..0000000 --- a/model/Untitled-1.ipynb +++ /dev/null @@ -1,78 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": 2, - "metadata": {}, - "outputs": [ - { - "data": { - "text/plain": [ - "torch.Size([8, 6, 16384])" - ] - }, - "execution_count": 2, - "metadata": {}, - "output_type": "execute_result" - } - ], - "source": [ - "import torch\n", - "import torch.nn as nn\n", - "import torch.nn.functional as F\n", - "import numpy as np\n", - "\n", - "class FourierLayer(nn.Module):\n", - " def __init__(self, B, concat_original=True):\n", - " super(FourierLayer, self).__init__()\n", - " self.B = B\n", - " self.concat_original = concat_original\n", - "\n", - " def forward(self, x):\n", - " # Applying the transformation for each value in the channel\n", - " cos_transform = torch.cos(2 * np.pi * self.B * x)\n", - " sin_transform = torch.sin(2 * np.pi * self.B * x)\n", - "\n", - " # Stacking the transformed channels\n", - " transformed = torch.stack([cos_transform, sin_transform], dim=1)\n", - "\n", - " # Reshaping to match original shape but with added channels\n", - " transformed = transformed.view(x.shape[0], -1, x.shape[2])\n", - "\n", - " # Optionally concatenate with the original data\n", - " if self.concat_original:\n", - " return torch.cat([x, transformed], dim=1)\n", - " else:\n", - " return transformed\n", - "\n", - "# Testing the FourierLayer\n", - "dummy_input = torch.randn(8, 2, 16384)\n", - "fourier_layer = FourierLayer(B=0.1, concat_original=True)\n", - "fourier_output = fourier_layer(dummy_input)\n", - "fourier_output.shape\n" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "torchGPU", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.11.4" - }, - "orig_nbformat": 4 - }, - "nbformat": 4, - "nbformat_minor": 2 -}