From af26da8cad0c5445b860cd0c2b686e4d35ebe202 Mon Sep 17 00:00:00 2001 From: Andrew Gundersen <53452154+gundersena@users.noreply.github.com> Date: Mon, 17 Feb 2020 16:46:35 -0600 Subject: [PATCH] Add main folder for repo --- main/CSR_Net/__init__.py | 1 + main/CSR_Net/__init__.pyc | Bin 0 -> 192 bytes .../__pycache__/__init__.cpython-37.pyc | Bin 0 -> 189 bytes .../CSR_Net/__pycache__/blocks.cpython-37.pyc | Bin 0 -> 3395 bytes main/CSR_Net/__pycache__/map.cpython-37.pyc | Bin 0 -> 1453 bytes main/CSR_Net/__pycache__/util.cpython-37.pyc | Bin 0 -> 1396 bytes main/CSR_Net/blocks.py | 81 +++++ main/CSR_Net/map.py | 75 +++++ main/CSR_Net/map.pyc | Bin 0 -> 1663 bytes main/CSR_Net/util.py | 34 ++ main/__pycache__/ds.cpython-37.pyc | Bin 0 -> 3733 bytes main/__pycache__/prep_vctk.cpython-37.pyc | Bin 0 -> 3588 bytes main/__pycache__/util.cpython-37.pyc | Bin 0 -> 8958 bytes main/ds.py | 160 +++++++++ main/main.py | 141 ++++++++ main/util.py | 317 ++++++++++++++++++ 16 files changed, 809 insertions(+) create mode 100644 main/CSR_Net/__init__.py create mode 100644 main/CSR_Net/__init__.pyc create mode 100644 main/CSR_Net/__pycache__/__init__.cpython-37.pyc create mode 100644 main/CSR_Net/__pycache__/blocks.cpython-37.pyc create mode 100644 main/CSR_Net/__pycache__/map.cpython-37.pyc create mode 100644 main/CSR_Net/__pycache__/util.cpython-37.pyc create mode 100644 main/CSR_Net/blocks.py create mode 100644 main/CSR_Net/map.py create mode 100644 main/CSR_Net/map.pyc create mode 100644 main/CSR_Net/util.py create mode 100644 main/__pycache__/ds.cpython-37.pyc create mode 100644 main/__pycache__/prep_vctk.cpython-37.pyc create mode 100644 main/__pycache__/util.cpython-37.pyc create mode 100644 main/ds.py create mode 100644 main/main.py create mode 100644 main/util.py diff --git a/main/CSR_Net/__init__.py b/main/CSR_Net/__init__.py new file mode 100644 index 0000000..1ac3f3b --- /dev/null +++ b/main/CSR_Net/__init__.py @@ -0,0 +1 @@ +from .map import MlRes diff --git a/main/CSR_Net/__init__.pyc b/main/CSR_Net/__init__.pyc new file mode 100644 index 0000000000000000000000000000000000000000..8e4b42fe9e6bf19d1653475a00b96e8bb086fd9a GIT binary patch literal 192 zcmZSn%*&OqWJ+u@0~9a;X$K%K76B3|K*Y$9!@!Ws$PmTIz?j0s5Ujxrl*nWR5*i?) zgcV5m<^-h{`)PpmmVl&l6AOZX6oUpTQEUO^>xUMn78UC!rlgnVr2tu}dHOD?#n~nK y1^T%;@kOb{`o%@b`nmZjsX4{^@$s2?nI-Y@dIgmw96-%BK=IO?R6DTEAj1Iw04DAL literal 0 HcmV?d00001 diff --git a/main/CSR_Net/__pycache__/__init__.cpython-37.pyc b/main/CSR_Net/__pycache__/__init__.cpython-37.pyc new file mode 100644 index 0000000000000000000000000000000000000000..e300e065a8f95c0ef6fef6b58fe2ff88387edd78 GIT binary patch literal 189 zcmZ?b<>g`kf?X>Y#EJvy#~=<2Faa43KwK;UBvKes7;_kM8KW2(8B&;n88n$+G6ID) z8E>)r<^-h{`)M-WV$Mx0C<5tP$xy@sq`<^4SN+i9)S_bj#FX^Xyc8fSHBa9qwK%&Z zzd%2^C^I*)BvH4xv>>%ewb@rJlkS!Kpp+UvTBjZB9s>dgIKA_r~5IDrwTHwdC3HoAJ!^o6mRN?A3*Zz<}%A zo!`Iz6Ka3s!{oA|@;TgM00lKD$&5CEUo$iNR@>^^ZJVIo%ACI2c8T$|L2c^XGpNI? zUGJ{du7Sm+9$37};)A6|eX#hIB>+o6>tLx@mO3Rjjj-_uW;M1#Q@L5Z&$v*HG*4w3 zXXyx@LZX73!|m(o9hP0*NU-nZLcj467ocF;=e9{n+oIlmqeG}g?R$P3Jf-eEqwU-` zY@_YMj5XzbnI{Ehd~}}aifE?!U?`jMkfudb7R{ZsyTkY;n)X>P(jphiD|_-}GOGiV z*5IzA2j4;=4N1nvU31s$Kn%xLO2+oj=1+te(%QAhWNcDXI=e3T1@SvD(XRt%Y($>4 z9@MD!5Z;b_YM_T(%Fg3HQ?`Jz3xmV;&lQM~-Hp>6%#6`TCKy9Ud7U(q7`J+l@tkFm zfWRql5L24wUA0nK=Ypy{A$dxf=>GZ(UVmH(1LcU}fN_qWrD{=x-J^)(9LilU^4pg; zxP=eFH7<9GZjr~CB74f;WN~l*2HU))NDSV9YgtG6c{K5AVOj;Z1o!lSy}g)sac%u5 zO0rmpDEiI#cfEB>z=yPA+8ySwJc8w|4JLZB7_<_e_G1}e(q3NTOtktCiPrkf8_^b% zt?jHxdSY#`&sSiA1-Jz+jzP|lfSe+u zb0`=^T>;Jm=tN;9oBcw_Y+nW2(?Jf`)_36Pi_sS7ETA>5v%1!qzomKPh{~MGxDU@9 z1IJO_p~D+kN2>TD7E4$xLs7l-*Oi%)8vF#dzK4ZgE(DDu3R*uL5z{0oZ-Wc5Ve+3O^ut3G|SD~0J?UXV-J}$ICUni_L$mry3VU7VF-~K=F8hvwyzEz@s z0Xom~&)+~p6N@Q-C3uMFJ8*0844@zQUz&5OlE8WB5_2Y^Vswy&U%RAtX#G!lhCK@8i%%ihzgM1;J9m7l}-7#}d-5 zhD^Deaete}95*OM1jrCK-KST^2hdxfa#0zJWR)~X!yGMFi+YSO5Y3~6p(SF?Bw{}% z5stH0dE#iDd;%S-WQ97xOfhk>lF3CJ>N#YB%h1AErfI&AKCxv4}0v!+E>aT-S$7ukMb(((%_;-PSZ^Hj>P5N-?M9zK9d7wED z<~V<(clIZ4_lLvI*T5~W zIL16Rp4ZS1eMZom_?!b>;Mbq^P!o literal 0 HcmV?d00001 diff --git a/main/CSR_Net/__pycache__/map.cpython-37.pyc b/main/CSR_Net/__pycache__/map.cpython-37.pyc new file mode 100644 index 0000000000000000000000000000000000000000..1df41b265cc4baaa98560f047e365f5d07260f29 GIT binary patch literal 1453 zcmZux%Z?m16t!K?uIXtelL-(Ys9^!rB56&si4dTW31XE&N?=nGvMObFQ`6N|72BO; zMjZ($BW1-eL_JG>ffe83Evxj>;Cq3LZIxu`2EK>5g~t| z*c2bgJ*fIFfFOcaB%`innFlSa{8^9%voH%O`IZPz_@_kpk{(4z@hdOu3Hq3%!3su` z!<1QnZpzArPpYPv>H@8s;&nuS21L45GA3x|y(VKSnD9>H%on~0PDvKLCO*kRm^rk; z!|IXLCtsSMMYHL$%+vw-s(t z7JX9}u<6Q1(vkfQxxFEGvgCeaOeJesOl|VyoM+Jk*^L#B+_BM`z}Pj}&yw%vrk$IE zrhaC7aB6u4**?QuwXDX__Bb!9Tx-sMC4b%<{h+1Nqg+hpHAJOlJ$fMZ)HLm=P~|K) z`AwJZO(pecmY4PD;PE4VC=H1AR(q^=ASrwe`fUJ02i|}V5~#?K(*ZjfxEptCvsVS~ z6&BpW0O|o${Tsl9yo9^3Bnx^*UwR9+qzmtiQUW%H?2^sjTk0yLWW&5|Wi8Ym2z3ns zcS^m3fOkv1i*Oy`J%n9=)K~9AW5cqBv(#{Q)3UYng}Q+~9DIm=zsRfVQ(*dI078A1 z&}37?BuwZoC|A7%PUIw6p9Bs|BfPzJK9ii=gmd@`<`wcj=TGN()p;V$MN_~9$Aea% z0NCyB1{JcZz;PZ^O4SG02}5(w1O7&$Lx30}>)&v;G_i?s4`f_5&n+`!8`Sezdu&YI~Gy zG@OXUsa&~|V=w&-d=6iG%2(*6eb1Xfkorb`V?RIJ&%fV${&qAP5)j+I|LyC9kiT(o zAuuM7pt+}@6j9WW8HGPHtmHFJiQ+0bA+zNF=AIGNQ|Sp&spk8=SA3Q#x8jic(jus4HnQ*X}mOkv3V|I9WEPs1{iaC3;V#F)ovt2mjj6O?{Zn z+bZ|4HuHPZXR5F|_hy+DuvWvi-oZrSKDbF!_~lmI>=&I$=h{+dMb)%E+mjB~m)YH2 z>GQpN*=_r zvNWW*DAS0Wn{>!8lkc)EYkMMGxhh;0;Tk)(#(tenBeB0Gs*Ow4DdHo}CJ$mi4p7d} zcR+3SPB$*T&h>G%ZdU8Y_i);s6{dPI)E-s9GewU+O%rxS8$^xxw{t*xqMyqg18bej zbY!eYWRyhog)=AGiJKN%zXIix%6qFzogWGN=91;0d%6vJ4d)HHT?{!1arbc=yl|+0kQL)u ziME_B2YcEb`KlgFZQ5uza7{iCjJUxxtMAqC5)RGNa$GK{8){wk!Y#F>ZmOK00{6?a zeTwZqg5dT6KY;(6HSW>^Lcjy??@$4hUR~oJEh0L8e1R(93Gl%sM}({1JN%JehqPFt zMTZv4eDy%wyi4;A&6n9?(WT#C(Mwjq2{=>Q|Hg~WgXEZ6`&yt?#fNZWgd zcj$F=nPt(FL`ZRwXx?K#&h>Ab3#CY|{g}rvH19zT0ZxVhEVIk(Z6*vz!#3@k=b?I@ z;EAc_wwtUR>tq#@gb^n9Y>o5D28%YC`~VX3#FHB!0vA&P7*62(EaB)0>xH+TkzJ1D zGiq&)+pvKCf5(P|$}?-Z7{>)tu%m8=vDywcf?ahh7|Kk}Ut5_?b(#u7(-a_^S!Ve~ zq-jy*I5V<`RjCDZ$qqQ_Q!t+N!CuxZr|$+eKX&=A&N&Inqt$?zOKV j0Z$3nC`0rkxhb(oILQB0kWA=dJCny8*YPMe)mr!&OGHzJ literal 0 HcmV?d00001 diff --git a/main/CSR_Net/util.py b/main/CSR_Net/util.py new file mode 100644 index 0000000..1479eda --- /dev/null +++ b/main/CSR_Net/util.py @@ -0,0 +1,34 @@ +import tensorflow as tf +from tensorflow.keras import layers + +def SubPixel1D(I, r): + ''' + One-dimensional subpixel upsampling layer + Calls a tensorflow function that directly implements this functionality. + We assume input has dim (batch, width, r) + ''' + X = tf.transpose(I, [2,1,0]) # (r, w, b) + X = tf.batch_to_space(X, [r], [[0,0]]) # (1, r*w, b) + X = tf.transpose(X, [2,1,0]) + return X + + +# ---------------------------------------------------------------------------- +import tensorflow as tf +from tensorflow.keras import layers + + +class MergeTensors(layers.Layer): + """custom layer that handles merging tensors for upscaling purposes""" + def __init__(self, type, name='merger', **kwargs): + super(MergeTensors, self).__init__(name=name, **kwargs) + self.type = type + self.c = layers.Concatenate(axis=-1) + self.a = layers.Add() + + def call(self, inputs): + if self.type == 'concat': + x = self.c(inputs) + if self.type == 'add': + x = self.a(inputs) + return x diff --git a/main/__pycache__/ds.cpython-37.pyc b/main/__pycache__/ds.cpython-37.pyc new file mode 100644 index 0000000000000000000000000000000000000000..da66c84a7484f150cf0d7656b5ca90e1c6ba4dad GIT binary patch literal 3733 zcmZ`+-ESMm5#QbW>R^^ghu0}O^0dlRG+J`F6!jIm+^^@7FlPF2M!*P`7 z-JPlG4pkf#NxIvWQ4w{sB#XyM8t11rB)E0sw2SWax7Qi_ClGvpeRo)&V{8F^6YYJp z{0jihH4B**fLuF(LR)}V1#RsDceDrSY9G+kH9%hnfHg^Lf$asK!T{(QdaT|PXT5*c zTW38#>uc!k&zfbus2?|UT{ z_%NDe81|aRE*@mDsig;#ia9qKX%A<*~v#1yud-P`Zz}PqFTtAXBBx!pK#?IsMj#|V%YKg!F+ObA5h2sRLEEAEA zQ7Z&G1Xc;G0T^p~P#C`RC#J4o_@B48ewoKA--_h!K??EmINkan&i9LHwl!4bw9(c4 zAdA&i73W*zJ6kgELoPoIlQb#9kOavu0x-|z|ML?rzs8$b*dL*l|@mOHr zn5#l_=Gn&}nX(#m{y;yRZ-|Dw43I}ZBF5kZ_y8s+~YYWnT3;Wy9o({7vW%T@C= zunjD;h9I!Y%Ff{%#S7Bb9o;yxVEJyjT3p!kj=Qk^nqI>WSN82ELajX!cdK@rG^vB~ zW4&6gr-5FbpHi;Rb;Wp92dmYsK-Z0mx}w+dt3Rb~LXwRmt^>#>^537yy$MRaqI^yi z-Gqm29KWaoDR#sOCtuSo&^AGI<5AUiv~`Mn&0 zvPh=s50siJ<3PII*mwtNf;2jr2Uib4augv6I6LF%R8_GTnpUJ{m1*NDiqkcCV2O%v zeY4k9mvEG#bgC$ysuqFE0E1Q))p03tn&aeLbwMRnQ=|l%OQQ+$3*8!*s13rgGCUvH zIgmjxQs=f91mmZ;;d^J2Yw85e#;Y$aWFUu1o;}oQ zkC9I>gE3~D!aer1ttrd6zQDbLaVx*!{y0Rys3do~c_zg!x9Q@|))k@D8RQq2=^M z7%siZ1F?wKuo~joOFDl8rw;g@X6+gVM5u1+Eg&f2=x&fvwY|wo2jw93ZoyChLib%E zg|r@9$X8zJqY(6W1a88Q7;Wqw|JvucLYS; z1l>4Og^Fchr4R1n4Nvc*V#2iH+lEKaDQjBuvOA;$sPOs&HpnUDG58GqI()YNWsUM2 z_mC8&&qt%cP$~T~5GV(=p`Z=(+fdDc!c_gD+C8$UP&qTUioo3yY6qR7?RviI5grow z4FOWBVpEc^7{Krc&uh|oPByAF@qff0!JU(63{R%PXWjp+%-MpQJy(&p-@2hlD@FEH zbR#`}12_M1h0VKwsyzxdVQ7LdoJ{2bip|h8!tm}vG(Ma0!cb0!VW?iBo#=Thaxv4$ z3;c&MoW#XwD(Cu-X!6QeK@}880SbgU{-djD~!mkQN4Y6qb7hYwm1^@s6 literal 0 HcmV?d00001 diff --git a/main/__pycache__/prep_vctk.cpython-37.pyc b/main/__pycache__/prep_vctk.cpython-37.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a9672bc8e4132a13d3eae85759b240bfb8d4398a GIT binary patch literal 3588 zcmZ`*UvDHw5$~RVJ3G5x@A>TYzdPRoNyt8&jer1xkc1>0QIJjnjvN|s8OGaVd+goW zxt`uTpVcEFe35vGa>^4TLb4y>Qy@M84?Oh)Lhyk66(Sz^Rj++#Sund@U0q$>UDZ`x zRd3JDMGW8G)}J4|x60VRXmIhlK;A;D_5lPFe8l2e&tq=KPVCy=i#^->v2Xhz4$wQJ z#yE^aPHWsza~#Ey#XaDyxHWFaZO(qggfD`pOa#YHJSX_i*r0KS^Vr=1H?6!-nJg!x zRAm^3A{*vo!W|QAPZj1&W3-cxfDCEP#n%Ax7FzWlK*mVg*byAkE{Udyo^r_B6D`pO z=8HMedCKA-Ylvmh71y5fV;+YXt%wD2-Hw`K5pu1X<_9t>lMmmy_v52i#%W&khoiKr z`a2WZAIdCMd9m9UsY?50US^}b$f_?kHn+ypk;qxt5PB_B-#N+HfjPw)sQ}zq=P}RF@_~Sdc^0_K~lqK0= zIxZpcpktcFbevf6jb9DZQD$6`k4;c1nTt$1kk5pp$>XHN&6sFARl~ic%8xSRRMNOr zIyTUX_h@WCxqcF74Q)6O&PnQ|5_C$H)Z& z*9j~VSOPG7=Z{Qo0Q=(2t)Ew!thQ3IJ1rntl@(j>X4Qe3lv_hdZj|1xre!8?%dDE2 zWve9fJ{qcnO$Zn!NnYeCNl1dK2f)IBU->y5ezfr1D>g5d49TqY_0gUi?|rPv zZ)q1^R{IJDJm#OiUk8V8>)>7X;JbB0H%^>WR)@;n_fC8ro^s)=vbqT~`UipB(B4U4 z=Nf<^=ABc{*~5+dmyAfpCI{JXojNu3ZTUOEIn~*pb66eKEnxGY%wyE9=ZaTp_Pll1 z*^Gzg&)bX!PIc{=Q3FS(o`1#<|EU(Vcg){q&-jDC{nDws^)=lUZjTA?z z4*JqDL7Yw!SAQpJ@y-9i|Cg)$wA(s+vFdzHw1Vij%7H&MyRdG*URG=S z&67UtyP{XH`b!5MB>m4b{?NAjq(}RiBRBPOy;?+id3Hj*I@1*6WgRWsJAtNamb#=@ z@oPP&u0x8oW3D5}$rF77a90&TB$m)$xX5y8>CuykG##t>7z4_M|MWx zQWEwB-x_cm5{^0)tV#+&S=Il3-*`62_@i_?8*L5-2uGPR-lWV51TzG%(m0dKcv!J# z!m1qQDmzprNK0@CRL{8Y;HiD^C$Qzo32 zv*S(Y2NlAd!l6Huaw3fn@pd!Qm=-zG(|8sg8whoKX^F5)r}{gi$wb-+3iU0$W7Ak> zgE8Gr11zcf#@7dZxlU3v_Fl8WKaTREoGP0K{P}-Ot72r&_UO6dYI&f z<3BZCPMd`no5kV9R&nc+H_QfVk{o1N8HeN-2^5Z-);9>XS!4!ok%Fj>*vn|u2!J^r zkL)kxT^{n6c*hNRhj*PtzQ{XHk9XLjyXc0z=Y}4}PKenRzT&L#8*bNGa3bDj5vYL^ zfxG7~@C9#?ulU5b#xZgt7Z5vK&My~k%HYL#FV^>FwLpMpgsch&Ljmz<(6D|aKOpw-FoBiM%~+N4 z9uZ-_J!$>sssX-->0`8t@-f4ucX;G<(b{g?xqL|%Z{o}c-}6jNS(pg+roIJ4Aot>K zkR7$R&T1cJ74?B)no6l)*RjQsi|o{>Lljiuj$=!sT`b1;k2hm)8!=#j8gg4bDN&(H zG6gOuKmjyFVbj>lclUtEyPz9?B2kCLZt>(HUiaiN3L#7zzHNAVMOiaHD{m7zfGVla zut8-jLih~*di-ehMcr|Qdq|4X7or7dR_T|3Ksl-zUG2gA9#r$8Fei02J9eq~I5)OT z!98%~4mwHO4MIs9$b`VJ2#`{iO$x#?0K=bL)ueYtHi{pp*MAT0ibP|0GHrg;|G&zd zFSrZeA#azY)4e;A>?v)~rviDvQLW04EH>*RD&Q#IB#DWVWIPd5l#q#OC&|NUIy#?e zB#D>|lSKAuJ9^`i{LHi~g>rs4&eYyS%r0S%CaJW%1Z-p=Ehq+76az!`*YaAs#_f3#k$WI18mG}Q#stnAortyGpAJF*SS3&)cG*cc3VE@zi3 z&W|oRt6heJ0$HVLfeKC2AEYP*>^|ltZ+Y)y3$!TEM?SSk0e6AE^exCk-0vLD%xc$C z-KWlqmxsKUm-n1|&iT$c)HfTAnugz-qrd;;=U>&dU(m($XM(taBXNYLNll7a>j_T? zqV-ts8J;KlZG@OXwnSDzTs7PjMJ-nwa<;8&4XtwY2z82pPA!l>RyA-n3K;< zugT|mZ(c4-M^+yQZ-MVDdW*6q=N@R@((H-4Y~cFLte;l8>@CmwHQ9`x>8*GxysG*5 zY;VXL>Pywp=+@xFpX!(BzqY-Mb+5+gj_HBa%-`6wYfw9B%}jyXzgg zGyNGLZs16$2{lhh&C{ju47tGE3eS`k>0p(XEY~l<>5`TTGovl!g;!tHTbw)`#=((NA%iZokxlwNzhrO_$M*U4UaOnv*4g=MvYZav7r~mv< z^!u;tN3`9a;=u2h*E^Mx(&((?cnwGLBAQ&}ny=+LT9KP*!2((%x6qoof!1Q$&TY`x zZY)+`SI`u#ie<6IoPGM6DwbaonWsxvk|17EkB9+ozPcdk_IX*>iuvxon6f* zX*f*$VW|9YIOuG73!S}A9QtwCS6!tx)2(8zH;Plg8+6ivYE=p|3FEFJCX_=>oti3| zqCR^RkEsT3{Zvy6IQ{D7wfB=yC2N7)9Q7rru)p?tnA}YV!?lizdO;doPDWV7Wvp_o z7exKFQ5wbV;a=hRe$Ye%i6Qrm?(pUz76qj8>`wfbUC#S>|CVQc4b*R6lZbotTxtj(o17~6Vu@LZ`Y5? zT~M^wr7Y%&n#rfo`@_QM43a`0MlC}%(d(V*Du^>0^!;=&+X_8aEUmIx#5g<`tFs^y z8;vHG#MeZ&c#>O;SnrTJ(T~pLNtn}B5UF;+Ty7i8nRTd>&4B~4-m(kvc2P~j)ZYo> zQCL{LfehoKIvA!=FUms2ujWl|>Add3j0c;^=|wh)4TU}PH^R-R{~TVM@Pb5k;pAFR z4(!mr(2vd*4qo*|5ILCGh7bHbFyw2gIDlyo2iq`+1N%^$aRDPZgYm6ODgJYy)p=^3 zrRD-PBtAu=n8?1c*eDb&^E4-QEU+}0urGfD0}>LZCYmBU|2a-(15fZVlm0C{mREx_VQxfv5Tb47%J6(knJ`=9li495eqg_3g^bBmwwl);Z>NV)bec^OsyNgDT8U z7{!WJj|{IgO|r@v+ic!XWt)3_unl0+62>X8ZHr~G@FeDE!!bRvtZ4rfX(?$|`L9a0 z+Yfr7?-w=ShjziKg0B1iy-^TPdMdsz2OZy6OEd#*l=>1iE;YYJ4cYu-4H7tXEsZbszPvVX0)_&o_TLb9-c5hL9SUAC|wc4Gy==x>IM zE2LmN!TU+D8D3=)p1jL*yVqQ{&ZNzLar(-OrR%@;e0ce#OK$1^uYEnd{K6%EoR)LS zzhCPQ+oUQN`+gGi)&jMew5KSMg{@02?`(4I>n)z))BoQ3#h?8Am;dp~r{7!t*&F|! zJ-EK!p{=GLooOHFU>r+fs=2nWALxggC{fK|R3i*CmthFL#mHtyM4!enixEJP11Oc@ z*oQ8pgEWZ!C=CIyaGOa{ozX>QD*{TZy_UsJ>e%(NJ1*>!BLopu-|xghl0ey5r&yCn z?Y==A5RSpr8C+`6{A_V5<X(Vn z35!qiOza8XK4#Jp+1VMBPXQCsN;-IJegq=dnO|(l5Kra zxd~{Qkgt`!@?O{(DCrJWxDyRViOUuUQ-yLTY1B#FpfBCgP(E2zN*o|yD{3<{NY-2R zQiIyu0bG}1H$doC*uiiJ&(7OWR5z!9e+ncE3v>075}t%xqV zzm0=kv8FpPAy=&EOJY^eo_Tz?jsbHTn9+H*S9fqR)k!JN(GFmL&N0~Iz99{gEhtm6 zltEe2n%qNBzzBTnjS>x8@k$+Yn#)C9ji&cBr-&}813u6fRn{-$4w%B_K zl8LaecjI6K!Ir&Op6h9}xt02Hus0Z`MI{M#!fv!l@n5$kO1E|lgB%)GXi|jx{;}v+ zMUDR)wAsQ;^pkTmupQXgy}AJ}=4cL{f~QxWA?E^+~=Z_+MBkJvLw)XQM^ zl++9J7J9PgME0O3(pOh$-^jU7gulzeSJ&vyE7V-4<_0ydqIrs>tEB)jZ-hY;H9sNi zrBkxTP`L9J&HFVp>(hv6PJM;${WdieXQ|(z=1GCiH>u+$nxZiQtxg!nQ{W<1z+$h^ zAO7!huO4^@lJn35UkkA+bOv@YhS2ap!9u_+jnhsi`<(u5!t>gt#RdD?Vc;`apj)-p_W(FZVFe5 z#FmiI%FVxnDdqyL^T?FdhltnGDjX!;ukrmk8i%WTuA^-{a%A&R|MRbo9hmuC?i|$R zJUQlZb#}eL*R>Ck&wN{ZPy0yAYukeVZ&r=5X)@2-Clz)4%R3Wj^HnXyP+zqn?am} zH}CvPtYx*!-6#%|cDkFk%)(6fhG9|dk9r8-aJw*M)GMk)`!PIxX0PpZ(z}KBomoJZ($ z0YO_$S6di4OU4LSW+B?^xJW4?)%TGAL)AobgA`qoN*TLf$qhpEu}LY~Hk9&_o|{l{ zH?^d(PlyhsmL+!IBFql&CuKbD4g-Y@2Y66oHqziEOx*irIVf?XuA2*eOmUqO`#MJ!(@n2-)u88h?haG+J_MmYKvOuJPxz4p?)QhI6w&i{>3(Im zOy-Nq-h{A99X9E7db8@KG$^?X<41VD45N*jfgQo$f>8HxcNUG38=|0(-OoRONRj4~ zgmEU>!$mxTgW+qVOqCDGZID}n6C(PNS@!2BXr9w1W9`WqPUZ=7V3%ieFg=?C$=_rS zYa0nSN|(+L5OW+<4mC;wAL?V2fzX4(;^%sd)RA-aUUnSy^{fWNBHeI5x|&sPlLcLM zGwYpT7bnNf-FKq?RhP~9c<9GhvpNr@p!s8TzX_p4n!_##*3#{!*?G*HBrYJ4HhbW9 z22^|;h3$45A%?|nfa8GDH<$quL0B#V^unRc)4w0=6m^RGIMh+2UQy{H4M|bsi}5;O zRfv=7Qu7`)x2ZW9a3})5pW#SW&}hQJ-nycWc&|L4hb8hDxK8UGQU?%oa8RKd4k+ir zrGv6075j*`hkJyp9hz+VBkj0>BOxK2NYr&{Ani6oOPbPpfL-@&(6+2FT>(8`Chh2P zQ2Z?myj9sm{I{rfFkSYyoY0<*NvNVrGP>A8B!~<;>Z9n$Bk#bA!KhtSPd8(jSFl3{ zz>)bRokw?+x*K~gVwQm_>%k)#4R}l`>cJ-5!z>r4mu;y_$ontxn4#jX4)IMtlK9Lz zelh`P@>+<5#X}t#hdH*S4nQ_w4W=}5`wZe^fHoj@UzFD?u$H?z{G0uKkyZ{Ij58&& zpzZ1p3`EmnTq#GRt?VOVr5=nQ@^}NFT%kLd)xij79|cgF8zX)^8DZuQ;%@4nlXf{~ zX0`l;v;ec!+&HKMxEpCRuVQVMtQ;bOAJ_8Q!Tg~%uB8jxiw8?NdEhyw&*XE5LOSQP zF@k{YWq|Yx2qKB@G2cKB{(XT)I>@z%L2D?#`Eh3;m`2`;mP!lJW zS<5QzE!BpSWvh?%J*Gs-T4qvQJMj%l;hn876Xb8GK@w1DHju?!Ijc|KAz1BE7>pwn zps|>u`D%iwo5F=}DK${+Ez~EN2YAm}i8C6O&#Q#)R)8Wh>L4GXKnU)vFksJWm$+ph z*`QI*2JoMXd4-beV;|V9ts;xT_A9)5lTRNN*q;ugdVnj;VrM3@!I*HuE;Z z8fv$=#pevrh-ey=63zTqYoW}EhP0ILQbPq_&)lWQ^e$o?nNTtXST1lgR@4>YbH-#8 zUg6L?dD$klUX52o}ir)VP2Uox}ec3SK5Le9={k+)}%SK%+ z;aC+8+66evWq{(cQ9|Pa!WmRo7pFe@xXwfhML%)|XcJUh0zgqk zS*JiI)lpa0spnP+2{5kPtt&fyM9??0j86v=rG$)7$Qk+_qgXGSdgh;~JEPfPAk z@yM@nlpl2r2SHvdq(7A18nf*q zOFJUyKCx2r5WIgkKp=ZTi8RGJ_XT!9Cr6IMqW!~8(X~C<%TqmYPE^bQ2aftaW}+?Y zrs~f?PWB3cLKqDz^%buCB^?2ZL z^q&y?j6G7*WA+%Y4k(#KI!V45%mTM3^%(f?2?Yc%Df*m)6cQ-@5K%cs<5O)eHX#bW zDrXY2nW~ffPn*g4f|+WwnSe9=hNsO`{eqe5vzZ!SFjMWRGod2LYp6p8JPSvW&2{xh zw4?M@%d_tw`$9HWy{W=`qp;sWs&)sZ5&qahAyiP}0|XCsxP$MzM`RYB{T4@EN8~w* z3W7&+;;kl^Lr_~9ZBXD;Qxh2p?77beKy8Vet>?UX z6aa#xJ^haXi6T?*8s%O3e+o$Nipqp23$GT+0mD7ERsKX(EFu!7odAL(<|Inxv^1_- zky|TjK+R;Jy7AzCp{HH-ht%y+vqcSK@I9g;YRD^i^>PX2$0D_dTLr*9!i&jyLncJG zlW4Oa#Oix=3s^l}1~d31sT9s1U`Bktdxdl@xq_w+-{%+$EjrhYg;z^|=;$@qF`Ooj KrKyK3NBj>SUW8i! literal 0 HcmV?d00001 diff --git a/main/ds.py b/main/ds.py new file mode 100644 index 0000000..f38fb50 --- /dev/null +++ b/main/ds.py @@ -0,0 +1,160 @@ +import os, argparse +import numpy as np +import h5py +import random + +import librosa +from scipy import interpolate +from scipy.signal import decimate + +from scipy.signal import butter, lfilter + + +class Prep_VCTK: + """main class for creating data pipelines""" + # we just changed the add_data function into the __init__ method + def __init__(self, type, num_files, dim, file_list, + scale=4, + interpolate=True, + low_pass=False, + batch_size=32, + sr=16000, + sam=0.25): + self.type = type + self.num_files = num_files + self.scale = scale + self.dim = dim + self.stride = dim + self.interpolate = interpolate + self.low_pass = low_pass + self.batch_size = batch_size + self.sr = sr + self.sam = sam + + self.path = '../data/multispeaker' + out = f'{self.path}/vctk-{self.type}.{self.scale}.{self.sr}.{self.dim}.{self.num_files}.{self.sam}.h5' + with h5py.File(out, 'w') as f: # create h5 file for data to be placed in + self.add_data(h5_file=f, inputfiles=file_list, save_examples=False) + + def add_data(self, h5_file, inputfiles, save_examples=False): + # Make a list of all files to be processed + file_list = [] + file_extensions = set(['.wav']) + with open(inputfiles) as f: + for line in f: # for every file in the .txt file + filename = line.strip() # strips any spaces off of the filename + ext = os.path.splitext(filename)[1] + if ext in file_extensions: # if file is wavefile, add to file_list + file_list.append(filename) # add path of filename before adding + file_list = random.sample(file_list, int(self.num_files)) + + # patches to extract and their size + if self.interpolate: # if user wants to replace low-res patches with cubpic splines + d, d_lr = self.dim, self.dim # dimensions for lr and sd are the same + s, s_lr = self.stride, self.stride # extracting low-res stride + else: # apply scaling to lr audio + d, d_lr = self.dim, self.dim / self.scale + s, s_lr = self.stride, self.stride / self.scale + hr_patches, lr_patches = list(), list() + + for j, file_path in enumerate(file_list): #update user on progress (ie. 30/240) + if j % 10 == 0: + print (f'Making {self.type} data...{int(np.ceil(j/self.num_files*100))}% \r', end='') + # load audio file from file_list + x, fs = librosa.load(f'../data/{file_path}', sr=self.sr) + + # crop so that it works with scaling ratio (ie. divisible by 2, 4, 6, etc.) + x_len = len(x) # length of file + x = x[ : x_len - (x_len % self.scale)] + + # generate low-res version + if self.low_pass: + # x_bp = butter_bandpass_filter(x, 0, args.sr / args.scale / 2, fs, order=6) + # x_lr = np.array(x[0::args.scale]) + #x_lr = decimate(x, args.scale, zero_phase=True) + x_lr = decimate(x, self.scale) # downsample signal after applying anti-aliasing filter + else: + x_lr = np.array(x[0::self.scale]) #just sample audio at a lower rate (every 2, 4, 6, etc.) + + if self.interpolate: # zero padd array to have same dim as HD (ie, 4000300020001) + x_lr = Prep_VCTK.upsample(self, x_lr) + assert len(x) % self.scale == 0 + assert len(x_lr) == len(x) + else: + assert len(x) % self.scale == 0 + assert len(x_lr) == len(x) / self.scale + + # generate patches + max_i = len(x) - int(d) + 1 # max iteration?: file length - dimension + 1 + # iterate through the file in strides + for i in range(0, max_i, s): + # keep only a fraction of all the patches (not in use) + u = np.random.uniform() # a single value is returned between 0 and 1 + if u > self.sam: continue # only keeping a random % of the patches if args.sam is specified + + if self.interpolate: + i_lr = i + else: + i_lr = i / self.scale + + hr_patch = np.array( x[i : i+d] ) # current patch = current position + dim + lr_patch = np.array( x_lr[i_lr : i_lr+d_lr] ) + + # print 'a', hr_patch + # print 'b', lr_patch + + assert len(hr_patch) == d + assert len(lr_patch) == d_lr + + # print hr_patch + + hr_patches.append(hr_patch.reshape((d,1))) # create hr patches + lr_patches.append(lr_patch.reshape((d_lr,1))) # create lr patches + + # if j == 1: exit(1) + + # crop # of patches so that it's a multiple of mini-batch size + num_patches = len(hr_patches) + print (f'num_patches = {num_patches}') + num_to_keep = int(np.floor(num_patches / self.batch_size) * self.batch_size) + hr_patches = np.array(hr_patches[:num_to_keep]) + lr_patches = np.array(lr_patches[:num_to_keep]) + + print (hr_patches.shape) + + # create the hdf5 file + data_set = h5_file.create_dataset('data', lr_patches.shape, np.float32) + label_set = h5_file.create_dataset('label', hr_patches.shape, np.float32) + + # fill hdf5 files with patches + data_set[...] = lr_patches + label_set[...] = hr_patches + + def upsample(self, x_lr): #lr = lowres, hr = highres + x_lr = x_lr.flatten() # flatten audio array + x_hr_len = len(x_lr) * self.scale # get (len of audio array * scaling factor) + x_sp = np.zeros(x_hr_len) # create zero-padded array with new length + + i_lr = np.arange(x_hr_len, step=self.scale) # create lr array with step size of scaling factor + i_hr = np.arange(x_hr_len) + + f = interpolate.splrep(i_lr, x_lr) # "Given the set of data points (x[i], y[i]) determine a smooth spline approximation" + + # Given the knots and coefficients of a B-spline representation, evaluate the value of the smoothing polynomial and its derivatives. + x_sp = interpolate.splev(i_hr, f) + + return x_sp + + @staticmethod + def butter_bandpass(lowcut, highcut, fs, order=5): + nyq = 0.5 * fs + low = lowcut / nyq + high = highcut / nyq + b, a = butter(order, [low, high], btype='band') + return b, a + + @staticmethod + def butter_bandpass_filter(data, lowcut, highcut, fs, order=5): + b, a = butter_bandpass(lowcut, highcut, fs, order=order) + y = lfilter(b, a, data) + return y diff --git a/main/main.py b/main/main.py new file mode 100644 index 0000000..3fc3409 --- /dev/null +++ b/main/main.py @@ -0,0 +1,141 @@ +import os +import random +import datetime +import argparse + +import numpy as np +import tensorflow as tf +from tensorflow.keras import Model +from tensorboard.plugins.hparams import api as hp + +import CSR_Net +import util + +# confirm tf is using GPU +# print("Num GPUs Available: ", len(tf.config.experimental.list_physical_device$ +# input('Press enter of gpu settings are good') + +# gundersena@75.86.178.105:~/Desktop/crimata-super-res/train/logs/weights ~/Desktop +# scp rm -r gundersena.75.86.178.105:~/Desktop/crimata-super-res/main +# scp -r ~/Desktop/crimata-super-res/main gundersena@75.86.178.105:~/Desktop/crimata-super-res + + +os.environ['TF_CPP_MIN_LOG_LEVEL'] = '3' + +def make_parser(): + """creates argument parser from train and eval""" + parser = argparse.ArgumentParser() + subparsers = parser.add_subparsers(title='Commands') + + # train + train_parser = subparsers.add_parser('train') + train_parser.set_defaults(func=train) + + train_parser.add_argument('-i','--model-id') + train_parser.add_argument('-c','--from_ckpt') + train_parser.add_argument('-k','--new-data') + train_parser.add_argument('-d','--dim-size',type=int) + train_parser.add_argument('-x','--num-files',type=int) + # train_parser.add_argument('-t','--train-file') + # train_parser.add_argument('-v','--val-file') + train_parser.add_argument('-e','--epochs',type=int) + train_parser.add_argument('-b','--batch-size',type=int) + train_parser.add_argument('-o','--cycle-length',type=int) + train_parser.add_argument('-m','--max-lr',type=float) + train_parser.add_argument('-n','--min-lr',type=float) + + # eval + eval_parser = subparsers.add_parser('eval') + eval_parser.set_defaults(func=eval) + + eval_parser.add_argument('-i','--model-id') + eval_parser.add_argument('-n','--num-examples',type=int) + eval_parser.add_argument('-w','--wavfile-list') + eval_parser.add_argument('-r','--scale',type=int) + eval_parser.add_argument('-s','--sample-rate',type=int) + eval_parser.add_argument('-a','--make-audio') + eval_parser.add_argument('-c','--from-ckpt', default='True') + + return parser + + +def train(args): + """High-level method for training a model""" + # load data + x_train, y_train, n_sam = util.load_data(args, type='train', num_files=args.num_files, full_data=True) + x_val, y_val = util.load_data(args, type='val', num_files=int(np.floor(args.num_files*0.3))) + + # callbacks + checkpointer = tf.keras.callbacks.ModelCheckpoint(filepath=f'logs/weights/weights.{args.model_id}.tf', + monitor='val_loss', save_best_only=True, save_weights_only=True, mode='auto') + + # smart_learn = CSR_Net.util.SGDRScheduler(min_lr=args.min_lr, max_lr=args.max_lr, + # steps_per_epoch=np.ceil(n_sam/args.batch_size), cycle_length=args.cycle_length) + + # lr_finder = CSR_Net.util.LRFinder(min_lr=1e-7, max_lr=3e-2, + # steps_per_epoch=np.ceil(n_sam/args.batch_size), epochs=args.epochs) + + # logdir = f'logs/fit/{args.model_id}-{datetime.datetime.now().strftime("%Y%m%d-%H%M%S")}' + # hparams = {'max_lr':args.max_lr, 'min_lr':args.min_lr, 'cycle_length':args.cycle_length} + # param_logger = hp.KerasCallback(logdir, hparams) + # + # tensorboard_callback = tf.keras.callbacks.TensorBoard(log_dir=logdir, histogram_freq=1, + # write_graph=True, update_freq='epoch') + + # make model + model = make_model(args) + + # compile model + optimizer = tf.keras.optimizers.Adam(learning_rate=args.max_lr) + # optimizer = tf.keras.optimizers.SGD(learning_rate=args.max_lr, momentum=0.8, nesterov=False) + model.compile(optimizer=optimizer, loss='mean_squared_error') + + # final review + util.review_model(args, x_train, y_train) + + # train model + model.fit(x=x_train, y=y_train, batch_size=args.batch_size, epochs=args.epochs, + callbacks=[checkpointer], + validation_data=[x_val, y_val], shuffle=True) + + # plot loss and lr metrics + # lr_finder.plot_lr() + # lr_finder.plot_loss() + + +def eval(args): + """test the model on real audio""" + # make model + model = make_model(args) + + # create list of file names + file_list = [] + with open(args.wavfile_list) as f: + for line in f: + file_list.append(line) # this is gonna get pretty big for a real dataset... + + # eval on random sample of files + file_list = random.sample(file_list, args.num_examples) + for idx, line in enumerate(file_list): + file = line.rstrip('\n') + CSR_Net.util.eval_wav(file, args, model) + + +def make_model(args): + """define a graph and compile model""" + model = CSR_Net.MlRes() + + if args.from_ckpt == 'True': + model.load_weights((f'logs/weights/weights.{args.model_id}.tf')) + + return model + + +def main(): + parser = make_parser() + args = parser.parse_args() + args.func(args) + + +if __name__ == '__main__': + main() diff --git a/main/util.py b/main/util.py new file mode 100644 index 0000000..79a0fe3 --- /dev/null +++ b/main/util.py @@ -0,0 +1,317 @@ +from keras.callbacks import Callback +import keras.backend as K +import numpy as np + + +class SGDRScheduler(Callback): + """custom callback for implementing a SGDR learning rate""" + def __init__(self, min_lr, max_lr, steps_per_epoch, lr_decay=0.9, cycle_length=10, + mult_factor=1.5): + self.min_lr = min_lr + self.max_lr = max_lr + self.lr_decay = lr_decay + + self.batch_since_restart = 0 + self.next_restart = cycle_length + + self.steps_per_epoch = steps_per_epoch + + self.cycle_lenrfrrgth = cycle_length + self.mult_factor = mult_factor + + + def clr(self): + fraction_to_restart = self.batch_since_restart / (self.steps_per_epoch * self.cycle_length) + lr = self.min_lr + 0.5 * (self.max_lr - self.min_lr) * (1 + np.cos(fraction_to_restart * np.pi)) + return lr + + + def on_train_begin(self, logs=None): + K.set_value(self.model.optimizer.lr, self.max_lr) + + + def on_batch_end(self, batch, logs=None): + self.batch_since_restart += 1 + K.set_value(self.model.optimizer.lr, self.clr()) + + + def on_epoch_end(self, epoch, logs=None): + if epoch + 1 == self.next_restart: + self.batch_since_restart = 0 + self.cycle_length = np.ceil(self.cycle_length * self.mult_factor) + self.next_restart += self.cycle_length + self.max_lr *= self.lr_decay + + +# ---------------------------------------------------------------------------- +import matplotlib.pyplot as plt +import keras.backend as K +from keras.callbacks import Callback + + +class LRFinder(Callback): + """ + custom callback for evaluating the optimal lr range for SGDR + Usage: + lr_finder = models.util.LRFinder(min_lr=1e-5, max_lr=3e-2, + steps_per_epoch=np.ceil(n_sam/args.batch_size), epochs=3) + """ + def __init__(self, min_lr=1e-5, max_lr=1e-2, steps_per_epoch=None, epochs=None): + super(LRFinder, self).__init__() + + self.min_lr = min_lr + self.max_lr = max_lr + self.total_iterations = steps_per_epoch * epochs + self.iteration = 0 + self.history = {} + + + def clr(self): + '''Calculate the learning rate.''' + x = self.iteration / self.total_iterations + return self.min_lr + (self.max_lr-self.min_lr) * x + + + def on_train_begin(self, logs=None): + '''Initialize the learning rate to the minimum value at the start of training.''' + logs = logs or {} + K.set_value(self.model.optimizer.lr, self.min_lr) + + + def on_batch_end(self, epoch, logs=None): + '''Record previous batch statistics and update the learning rate.''' + logs = logs or {} + self.iteration += 1 + + self.history.setdefault('lr', []).append(K.get_value(self.model.optimizer.lr)) + self.history.setdefault('iterations', []).append(self.iteration) + + for k, v in logs.items(): + self.history.setdefault(k, []).append(v) + + K.set_value(self.model.optimizer.lr, self.clr()) + + + def plot_lr(self): + '''Helper function to quickly inspect the learning rate schedule.''' + plt.plot(self.history['iterations'], self.history['lr']) + plt.yscale('log') + plt.xlabel('Iteration') + plt.ylabel('Learning rate') + plt.tight_layout() + plt.savefig('plots/lr.png') + plt.clf() + + def plot_loss(self): + '''Helper function to quickly observe the learning rate experiment results.''' + plt.plot(self.history['lr'], self.history['loss']) + plt.xscale('log') + plt.xlabel('Learning rate') + plt.ylabel('Loss') + plt.tight_layout() + plt.savefig('plots/loss.png') + plt.clf() + + +# ---------------------------------------------------------------------------- +import tensorflow as tf +import numpy as np +import h5py +import ds + +def load_data(args, type, num_files, full_data=False): + np.set_printoptions(threshold=100) + path = '../data/multispeaker' + + # load training data + datasets = os.listdir(path) + + for dataset in datasets: + if str(args.dim_size) and str(num_files) in dataset: + if args.new_data == 'False': + make_data = False + break + else: + make_data = True + + if make_data: + ds.Prep_VCTK(type=type, num_files=num_files, dim=args.dim_size, file_list=f'{path}/{type}-files.txt') + + with h5py.File(f'{path}/vctk-{type}.4.16000.{args.dim_size}.{num_files}.0.25.h5', 'r') as hf: + X = np.array(hf.get('data')) + Y = np.array(hf.get('label')) + + n_sam, n_dim, n_chan = Y.shape + r = Y[0].shape[1] / X[0].shape[1] + + if full_data: + return X, Y, n_sam + else: + return X, Y + + +# ---------------------------------------------------------------------------- +import os + + +def review_model(args, x_train, y_train): + """reviews model parameters and raises warnings if something is not recommended""" + # prints preivew of the data + preview_data(x_train, y_train) + + # assert not overwriting weights + if args.from_ckpt == 'False': + files = os.listdir('./logs/weights') + for file in files: + if f'loss.{args.model_id}' in file: + input('Warning: Are you sure you want to write over these weights?') + + +# ---------------------------------------------------------------------------- +import numpy as np + + +def preview_data(X, Y): + print ('Preview X:') + print (f'Shape: {X.shape}') + print (f'Max: {np.amax(X)} | Min: {np.amin(X)}') + print (X[1]) + + print ('Preview Y:') + print (f'Shape of Y: {Y.shape}') + print (f'Max: {np.amax(Y)} | Min: {np.amin(Y)}') + print (Y[1]) + + # data = eval_wav.get_spectrum(X[:100].flatten(), n_fft=2048) + # label = eval_wav.get_spectrum(Y[:100].flatten(), n_fft=2048) + + input('Press enter to continue...') + + +# ---------------------------------------------------------------------------- +import os +import librosa +import numpy as np +from keras.models import Model +from scipy import interpolate +from scipy.signal import decimate +from matplotlib import pyplot as plt + + +class eval_wav: + ''' + Helper function for eval() in main.py + Takes a single wavfile and evaluates it by exporting audio and spectrogram + for hr, lr, and pr + ''' + def __init__(self, file, args, model): + # ../data/VCTK-Corpus--- + x_hr, fs = librosa.load(file, sr=args.sample_rate) + + # ensure that input is a multiple of 2^downsampling layers + ds_layers = 5 + x_hr = eval_wav.clip(x_hr, 2**ds_layers) + assert len(x_hr) % 2**ds_layers == 0 + + # downscale signal + # x_lr = decimate(x_hr, args.scale) + x_lr = np.array(x_hr[0::args.scale]) + # x_lr = downsample_bt(x_hr, args.scale) + assert len(x_hr)/len(x_lr) == args.scale + + + # upsample signal through interpolation + x_ir = eval_wav.upsample(x_lr, args.scale) + assert len(x_ir) == len(x_hr) + + # trim array again to make it a multiple of 800 + x_ir = eval_wav.clip(x_ir, 800) + print(f'Input length: {len(x_ir)}') + + + n_sam = len(x_ir)/800 + x_pr = model.predict(x_ir.reshape(int(n_sam), 800, 1)) + x_pr = x_pr.flatten() + + # save the file + filename = os.path.basename(file) + name = os.path.splitext(filename)[-2] + + if args.make_audio: + audio_data = np.concatenate((x_hr, x_ir, x_pr), axis=0) + audio_outname = f'../samples/audio/{name}' + librosa.output.write_wav(audio_outname + '.hr.wav', audio_data, fs) + + # save the spectrum + spec_outname = f'../samples/spectrograms/{name}' + self.outfile=spec_outname + '.png' + self.S_pr = eval_wav.get_spectrum(x_pr, n_fft=2048) + self.S_hr = eval_wav.get_spectrum(x_hr, n_fft=2048) + self.S_lr = eval_wav.get_spectrum(x_lr, n_fft=2048/args.scale) + self.S_ir = eval_wav.get_spectrum(x_ir, n_fft=2048) + self.save_spectrum() + + @staticmethod + def upsample(x_lr, r): #lr = lowres, hr = highres + + x_lr = x_lr.flatten() # flatten audio array + x_hr_len = len(x_lr) * r # get (len of audio array * scaling factor) + x_sp = np.zeros(x_hr_len) # create zero-padded array with new length + + i_lr = np.arange(x_hr_len, step=r) # create lr array with step size of scaling factor + i_hr = np.arange(x_hr_len) + + f = interpolate.splrep(i_lr, x_lr) # "Given the set of data points (x[i], y[i]) determine a smooth spline approximation" + + # Given the knots and coefficients of a B-spline representation, evaluate the value of the smoothing polynomial and its derivatives. + x_sp = interpolate.splev(i_hr, f) + + return x_sp + + @staticmethod + def clip(array, multiple): + x_len = len(array) + remainder = x_len % multiple + x_len = x_len - remainder + array = array[:x_len] + return array + + @staticmethod + def get_spectrum(data, n_fft=2048): + S = librosa.stft(data, int(n_fft)) + S = np.log1p(np.abs(S)) + p = np.angle(S) + S = np.log1p(np.abs(S)) + return S.T + + def save_spectrum(self, lim=1000): + plt.subplot(2,2,1) + plt.title('Target') + plt.xlabel('Frequency') + plt.ylabel('Time') + plt.imshow(self.S_hr, aspect=10) + plt.xlim([0,lim]) + + plt.subplot(2,2,2) + plt.title('Test') + plt.xlabel('Frequency') + plt.ylabel('Time') + plt.imshow(self.S_lr, aspect=10) + plt.xlim([0,lim]) + + plt.subplot(2,2,3) + plt.title('Interp') + plt.xlabel('Frequency') + plt.ylabel('Time') + plt.imshow(self.S_ir, aspect=10) + plt.xlim([0,lim]) + + plt.subplot(2,2,4) + plt.title('Predict') + plt.xlabel('Frequency') + plt.ylabel('Time') + plt.imshow(self.S_pr, aspect=10) + plt.xlim([0,lim]) + + plt.tight_layout() + plt.savefig(self.outfile) -- 2.43.0