From dc15d9f56b3ba9641e60e73c52db61b3626e3d72 Mon Sep 17 00:00:00 2001 From: AIFSH <1509359472@qq.com> Date: Thu, 11 Apr 2024 01:25:45 +0000 Subject: [PATCH] first commit --- __init__.py | 18 + .../__pycache__/draw_landmark.cpython-310.pyc | Bin 0 -> 6146 bytes ip_lap/__pycache__/face_mask.cpython-310.pyc | Bin 0 -> 2240 bytes ip_lap/__pycache__/inference.cpython-310.pyc | Bin 0 -> 20690 bytes ip_lap/draw_landmark.py | 197 ++++++ ip_lap/face_mask.py | 50 ++ ip_lap/inference.py | 604 ++++++++++++++++++ ip_lap/models/__init__.py | 4 + .../__pycache__/__init__.cpython-310.pyc | Bin 0 -> 347 bytes .../models/__pycache__/audio.cpython-310.pyc | Bin 0 -> 6143 bytes .../landmark_generator.cpython-310.pyc | Bin 0 -> 6980 bytes .../pix2pixHD_disc.cpython-310.pyc | Bin 0 -> 4299 bytes .../video_renderer.cpython-310.pyc | Bin 0 -> 17006 bytes ip_lap/models/audio.py | 237 +++++++ ip_lap/models/landmark_generator.py | 239 +++++++ ip_lap/models/pix2pixHD_disc.py | 137 ++++ ip_lap/models/video_renderer.py | 571 +++++++++++++++++ nodes.py | 143 +++++ note.txt | 10 + requirements.txt | 7 + web/js/previewVideo.js | 155 +++++ web/js/uploadVideo.js | 203 ++++++ 22 files changed, 2575 insertions(+) create mode 100644 __init__.py create mode 100644 ip_lap/__pycache__/draw_landmark.cpython-310.pyc create mode 100644 ip_lap/__pycache__/face_mask.cpython-310.pyc create mode 100644 ip_lap/__pycache__/inference.cpython-310.pyc create mode 100644 ip_lap/draw_landmark.py create mode 100644 ip_lap/face_mask.py create mode 100644 ip_lap/inference.py create mode 100644 ip_lap/models/__init__.py create mode 100644 ip_lap/models/__pycache__/__init__.cpython-310.pyc create mode 100644 ip_lap/models/__pycache__/audio.cpython-310.pyc create mode 100644 ip_lap/models/__pycache__/landmark_generator.cpython-310.pyc create mode 100644 ip_lap/models/__pycache__/pix2pixHD_disc.cpython-310.pyc create mode 100644 ip_lap/models/__pycache__/video_renderer.cpython-310.pyc create mode 100644 ip_lap/models/audio.py create mode 100644 ip_lap/models/landmark_generator.py create mode 100644 ip_lap/models/pix2pixHD_disc.py create mode 100644 ip_lap/models/video_renderer.py create mode 100644 nodes.py create mode 100644 note.txt create mode 100644 requirements.txt create mode 100644 web/js/previewVideo.js create mode 100644 web/js/uploadVideo.js diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..a77e58e --- /dev/null +++ b/__init__.py @@ -0,0 +1,18 @@ +from .nodes import IP_LAP,LoadVideo,PreViewVideo,CombineAudioVideo +WEB_DIRECTORY = "./web" +# A dictionary that contains all nodes you want to export with their names +# NOTE: names should be globally unique +NODE_CLASS_MAPPINGS = { + "IP_LAP": IP_LAP, + "LoadVideo": LoadVideo, + "PreViewVideo": PreViewVideo, + "CombineAudioVideo": CombineAudioVideo +} + +# A dictionary that contains the friendly/humanly readable titles for the nodes +NODE_DISPLAY_NAME_MAPPINGS = { + "IP_LAP": "IP_LAP Node", + "LoadVideo": "Video Loader", + "PreViewVideo": "PreView Video", + "CombineAudioVideo": "Combine Audio Video" +} diff --git a/ip_lap/__pycache__/draw_landmark.cpython-310.pyc b/ip_lap/__pycache__/draw_landmark.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..24d7dc7a0626992925a60f1f8324ef0a644496fc GIT binary patch literal 6146 zcmbuD>2n-M6~O1(dsl}g+mf%0uh^T&atI*_#yMnpZHvf~BiT+eBn+eNUTNgnnRU;s zbub%}$O&gm2m}aMntIYvg#IH}a5cPNCXmbV1vM(G6D*NgBN* zWo#m8qmN{aev&mdlbo@IX6Cn^lt`WwE=#1aqO8cq0O_KGq#Ldtx{dV0wTY&PLX&hG z-BwWsC4OZ|AL+lW8QbZOS&3|ZQHGv7Id9>-i}O~_H*g-{taBdZyqohj&U-j-=e(Em z4$k{H@8rCn?pTrLZ=^fvE~?YrbPwG}_s^!TNX7wrz&J>DksB^cMv)fj0T|6qkne}q zd!Qw>pVh8NEAr}%Fk1aZ`KoN(OmCt$4@kxk%x!@7lig3tbYMlQsALb#%^xCrsY>?I z1XbHJACv|qvY)1)y$048Xjy~Y*p#i?s9jgdk^|)6<%BUzlxeA0T;6?}65Bds*QoA$ zPCc+aS0~I`u-#c5N{&AqV~9)?6OlG<`$3dCZPjW}8>J^}*xGWU#Myeyp;6+T3+195 z=IrtXm->EGDBG;;P?K52uKN?kY?L-l*Q!#}jIyR#^++9R z%_wi07weYOtjRraYV2&u9Gx7WoMHt~$r7knwcQ}+f3aYikqh7VLbZ+c$x$XyE)pR{V{g&|1*co$t^PUdUkq`+ewvKnc`UkcUdY zbX5)&k{pm$lt!W{mqV4y)uucdCcf(pJVtES3aCFEC3p+A6Gm9Z$ZY+PaX!;<`3E5hVFnj8KO0D^@PG`ZtBThLGUb9) z>*F#7WiS$Tf`@918j4D!RV^@B%J$2S=hLFhc7eR@Q#Js_Xwdde7>G^4mFphhBijqr zeiD)$a#mK>mPb2V1fI#KWOh#JsN~{zv`S|H-JJ|;)#mmSwW=Dc%BofRoi7c zRB*A!iYYb-CAJNd?U?MqWDg{fwrJKCBW(%N>!+gsuK}W|eGKhQ$+VV)Ice&EJRqlK zRs>=3U{Q(h0*If*lEQ9AexKkWLOv^a7`a#QVdNuVeCT@%@^gaUfjlDkR^$c2 z??m1!_$K66A%VW{LPo~}-;R7h@Eyqe1>Xr7cQ&*E4N{yFdfts(5PSp~UWTaWJ;=`s z#%C3~4HB5|1j~7=-1>b}G1HtzqKOlGnxl8ad?EZ4Ei5SvmB&wsR=C{ax?pr}d`a**q>GB(55jo6sG->__ku!sD-;2QE51Ybn{qF{#nq~Ht4e-#`cUlm+OUKM-^`31oX z$X^k>i2P;2OUU09ybKxhk-ij2z;+Vg0}WdiLVM>!e? z_z7fO8o(>aIP@a=?Kd)%ehUfIJcTvC6Z~Q1-wXZ-FynI`fF4g`4}Uq?;|{syup_?yUw1am!a2!09WJT4UWZRA^ooVV{2 zR>6tHuL#C)iXRmGJz#qa@P=19-Vm=*xcnhV0%=8Rz^nYY^hnpkN<(R=p}Y#4j;pdP zaZ-^kCPEor>+%G=%vGNTk!lfA)S@&cupdG}rnSBDBiV}Nnd@JPJ$}6k`{SkcZ5Y83 zWFbk)d3deoWld%;L$R12{ZF*BO7RtC*Vqkuy0$@8o4mX4u!ce z9cGV87Y9Rym4g)S1snNLfu2K+LRcW$pw#GsJOO!k*cGNPNo*$U7I}~kwIH*iu=`;b zSs0XDIn)|GL3T9IRhP z7Hm)J)!LHo5{ofwX_%{M-J81)RkgQ)>*LbVwR(+tu#v68F$8wwxTl?=FkXIi0PUX3 z)#;sM#UXag1hsQbUZt(tb_re713cHEF|0tx7_loh6z9B!^@dRCHSYlTZEA)4+QK@` zIJ$05ozZnfFWF#6rdv8%9uAgx(9*q`c?ehCLwQvkIQ0tis?B?WMahOS@cDHp__slC zaOi>75x1!h)@%TG_I1Bro%zrQzx6MX-y-nk^q6rVcI@+lnXH9E?V9K(6+FW+$@n~0K%wtE;Az*ABLK+=k z@f^p)*lM=ab*m_R1`F7&)xlqQmPN0@!htT=E!);iW_T~lT@qvCk5k0Jf_S%GeqdX7 zE(l-^Vb=-P1y9Q5TgEzWRfU84de@CNoS9`hu>xzDA0foitSd_Rb5@N;1rr`4=IE)B ziHXwqbd+sR8a^D5MSf(8TO$oUCB8XygRt1k zeh3-+5hkx-@+u}jh6EJwV?Og1{1(INfESw?ZGsz`aOCVbY)0wUlZ*Wnd*@n3JRTO{ zLAUP_^$fKLVY5etc{VTK!E8EI&gX-T6iMCslyWs4Iuh6ud zOlBe1)-)w8E01BC#!K$e6p-iF6gX|>Ro>7gcS>q=@=5@RVJqF79Fq zCq_m~r(q!G_}H20sJC4#oj7Zj9x6qfHj2ik#-^jq8)eY!*wo|$9ZjdkPM#9ox2!i^ zm+`iJ*U7ow_PwyyoIF1=-Z8?7$>~y*sN)Anqdo2u=f=lHJmY4{=&>AUJ!E)@_154M z%tlc^u4j0JZOCJ@V`S64-`M7P@q_n&Pf=&D}Ao*Wu=CY>t z$VwVRaQ#x0YnznpwR~Sbmrvx?yv9>CpGm3^%yROYra=HFv!Eln2A_YfrBx-VC=kqR enx_3tLtUstRsN}I5Y}tTKeSX{=Jjh`Z~Y6;{5OOE literal 0 HcmV?d00001 diff --git a/ip_lap/__pycache__/face_mask.cpython-310.pyc b/ip_lap/__pycache__/face_mask.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..29e8dd01e7f8d3e939d4ce05dc17fa5122df82ef GIT binary patch literal 2240 zcmZ`)TW=Fb6rP#AcR;!Fqkxby5bnKZ{U)c+NXLH< ze8}StHO^K(gKjM?Vjc~XR768rWYJWn13Yr_^I>uanz{x=kUr%Egg3dv-8XLE;hu20 z4~+_U=<@~-VAbOc#I;OtKOTrjvDz`d6nZN2#{&vn7GDDf9<+PV)XzXPna~NFQ`9W*QdkX9sWaRRC|9`C8K z)}^}D>nUBxcqF=`q8OzjE)&%q6j?75Jb_6mu4e@osp<_u_b82Xp2c!U$p6m4dHRc+ zy)@QB=?}NN8qIF(iexm_YH$4mZg#ul7zbQ#npR@O6KR?@7Tt`f7*c-xP*9(-%g4Ws~b>ugylap9D0D&j(m0XVtSFs>5mP@Vw}auvxM zkpAjVR+$qr`j`KQg|*r=jwHP!aswFr{dI3VyRNgccdy8Xvmdv6gDC_xi*is`Ef%kD zKZzc0J?SMS1hK5clYwWt<;(<8l;nwyBB(l|E&>T?%l=%(D5L&A%n7Iq{aI&k6@7$B z++ExBd`%c&LZiXz4_HD z8}gdM-G*L1Sg9HZXF%53LtF{;@?^DgtKeYmkl6KehjfI66zn@tBF%UeYO@s{efQ#Y zFvZjkF*G51E<}t7lMIU|{9&A^5dOT#MfiNyI=;9CuDNwQLvHrQtuUWv#t5--c9s4)MF7UD3fq(D-782a5v$w1H5rF$;+v(MO+I9 zMXqBAc|Ho`@OhHQa#ovkblBRM;)nxlULxL6%80cUl~u0 z<5Czp1KF(mxnXJc4%D$)*qV@-#HO{GCQ1XEmA4VC6>4#f?0s?~g9mY%Mp1V&EnuW> z!(ed(SUaUw#v2qV6z{rWIe^KYkcE<0k-KL@A(!oV8@*(fY#W2fcfYq#(fNfN=&dZfW}&AGzO`)_dZ1) z)`b8KwTsXWg?JUV)|@t7VQuEpE6$p8Vh2f(znz6q#0pxIBf|H8Sle zdNGaDdPiPDc^Dme2MMCu34)497YII(3ezYcEqZaqu60`SCaf6G)|13I8z1}zxotVZ zYqEIBY^eH|;KfJ=N_M9@NmZ8@ySaV#acaTcvOOW*YEwI7v)2DM=D%uNMFQHU*w)y; D!-;T& literal 0 HcmV?d00001 diff --git a/ip_lap/__pycache__/inference.cpython-310.pyc b/ip_lap/__pycache__/inference.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..5e82e8370334da17c28c3fbb489d495cc043277b GIT binary patch literal 20690 zcmb_^d6*o>bzgT+&$%!md z*inx}Ye{3*HYiAuWLt`)UD~lleNLj-QDVnQoW!wn*^c9X@_qgJi)lNKz0Tz?Oq^J< ziTr+5GqW=bkYe;@cIUlU)pggaSFc{Zo=Pf}(D2FZ9-NqdQ`7#Cf%Z=ffkXJZQ@Wt#Ko+ENzIPGqg#W{5U9jE`kZbJOS((a zBa%)^dQ{RWNw1T1TGHz!-7V=2lJ1f81J2N#HkEOPoe^i$S?8>GHaHJVM&H(~2PZY> zA?IPVxp&T()8El$sLx*07FuKXw}pD`zP3=my;g;+jVCmFz#e>4vj?3`b0KRp!bA2j z!ozch^@y#X*TzPcISb=sdSH6j?9wtuUqc{mi`Ff}bS#Ji#4;pCX77aEUFyP4EW64-tHpV2L=JuMnIic!}T~L5ko>=6!?Jr#aDdSg z5aK37hD73qqKU+HMH7h`MH2~6(L{oiio8TZD4IxIQ#6sNDw;@C6fNtWR5X#es%RoX z21S`fVoK3Oa(oewPHBRy3z|q26-^{ciY5|eMH4maohAKaDxOIEsG^C)k0_c*EGU{t z%qyBm+)^}=_+dp82{I$vLnOGSK@*8NMaz0?iY5}DQZ$k96-^}SiY5}DRy2|LjG~Fe z8;T|pKcr|P@mWO^sd{_Ricc^k=i87-d{U(%5;;W^i7SdG5_v@v30u)b!cjDlctz2& z{L6|a60a+oNRXk?1|nf8nn+wyG?5@fA`Ousn}Q}%<;M_uj3LS<&_seF3N(>;T+u}0 z2}SQAy<5>l;!#BtiCu~&5<5wAjE*z(vU)-!E-0EvoL4lFIHzbL@sgs6#92iXi8G2O z661=N<)2VAk$7IwMB=2PiNp(vCK4|ynn;{dG?6&1Xd+cU$7lyb+tm{yu}#rLVymKw z#F(NtlHR0fBC%P~MB))e6NxRPhfu@A3_YZt5QzsBO(ZgkCK3-Qnn-L=G?7@ZXdlA)*66C&}nqKU*aiY5~K6-^`#E1F0gQZ$h`sAwW_ zK+#0vSw+kA&ncQn98olpII3tOag2bempL%Yax`y{nght=+wn0FI)txF{%|y^qEJJ1 z+jvuh8VVH?*tElT1UPI*p-{zwWK#(`TF85O6B(Ny(jk1=D3^2X>poP|{hDt88^9sp z(0(YaD9IcTLIwL~pyPk`U4#SEb&3<>^T;T^K=9WHzSsocJ-F-Y><+I~-F2i=nwWk0 zs3nGa>aJsX5% z1o-7k$Q797T*+DKl=w0uyRuLyy&N)Rvy=;N48YJ6dIasO4XtWjD`}pC-TaW|S2mE0 z%0c%r_#ABU6Nr$tj)i19#8(gzgscLW);Pe5O)B`$8rQ>oF6k8ma%!nc6`{`L-@L% z1ZeVZJQNSdBk`!c&W_pfH+67v!tRn>oU~Jti_>;DaK!Ghdx4|&dKR~X?E%EZ z?Lm78IAIUlBfwquD7bY)kW`!n76MDjHN7ll#@ncV`3(XoPPq00EzcQjei9(!_{ymn z($%`i$_dKqToYNtHvvOFcq(i95x_8xWH=mwig?jENKo4->o(KIy!cdNstZ(7QmLu5 zZQ?I{8xx-Gwj*|Ql2nhcE%thSODtoJr2D7VdIR2|H{=aZnr~}M`r=yD5_?CNX-0f) z3hAashBV|D_0~@GWrzm$zpy&dsmc8_eugEGZK zbDAG-#60Y6oZ5u%=BY>QUObK2eYe@0E#4SXY1yq(Caf7BGnZ`)mH3$s5SToY1)+R( zMsVK8j3Dk^E##-mj_X>fd_m;FZz5+GW?ZE&dQ26Q!dgmiDbA&QU)kB@mL~VITBg`I z*|iRK+5A-}KV7X9%HGbZceOSMF(S`-s~!|IEDtGh3mOFAIsjA?b2^?>; zw&P{5O4$wKGEd2IuLj9dHQNeVDW0AS<;g4( z7G10Bs0yC1I{6^keC9F~R}HtFn?avDQ^H{zZEKyogEhwwA31jV*!kzPN6w6oA3Jj4 z1A~E9?75RCp1+VicInu2=gwRV`tFSndRuX) zjvbdtgSGcW$=towio4Q8S$%)LQuds(*J_oFH7c(7Hiyle0+R%ZVkOs@?4Z9Xv9q4Y zm0e6RWPUKzKv2Q0Iz;$#MzRTB+lhyu&NOX6f#Q*{WlOFI+s6 zJsKDjIT2<#jcKDPBcWtHIYs!#kX8+z=6 zzR-}7G-_+x1@Fr4Ol_n2`6LMHWt0u>Z$J)v+LXR%NEq@%OS)%5ZkaM1_RS^T9ReNk zP3k(g4PEo1kgfVxwXYeTIma;b$9zaecdHk-P44TyzL>DXOS&By*6iqY=zCt*Vv=bl z^?JB`1ECZ{^@tx?OnY5Z-Jqk8pt0$Y_#A48gHFtaV0`OZ>|rY)S6NTkk1gpU>p@S~8}+Vw((iJAwpm6w;fEJT{3KiEC+rB?k+}91Z`2RZ z>Gc%))$1E8Bj+27>ul^g?(f-a&`ZPaSAFHa@qJp~NLe>L?mULQ7J1U3AMn!{t8TQk z#~$#bqguTeE$)kJ*n|4~-hK`3SmUpmPfR^HsG(imcF!>4rp-ILh?Lj*J!pM54364c zhTlr}mN8)rV-G~1`s~OPTD|Yu`{>J{n*aI4;zK^>8l%x)UkhD9D?VHw@COz*%2C}k zwR!OoNk#1;l>@CC#22R21#?oH)C{9OTpvN3woHwI%Od#p*u%Hcrq6w$KI(1tw)y>Y zx}UN~?9tnyquS-)MO}&dI)B~NcD7$)a2m>4U*F)5F7Ciw?)2AzuXZgN7>@`12c{mq zqq{e~-QZ2o%vI0$nW;Uf^)ZHqwW+y^umrQ0JVEKJ0qkScud&&PULKr><)l|7CyU z)ZY2zH0j?Wy^xeKU2=B*y^QGs+&7<=Y2G8fkd`qinT93oTS(o}ul=$et8en2o=^Il z{0*bpkY@MH84F48nfWdigS8StOc&Cm5wpIPCXF;H71KzQ@-j%Xp1cDp)uNLAhDKbM zq9)AO&HCp0BlRuyF@F$au-Sja-{Ox=nj8nT>XP^P;(mV+V>Ad|a>JXN_pIOV56&4! zQTBq_Hq#F@X8K2epqVz@3M5fRN}|?mH#lu;eOucs>}Z>X1DFLveA(Z+I5HpcH()IH z_G|unj%mL(0m-Oqa<&f2bl>u~eR#S%#JuaVu-~M?h_p<7iwPzlarT;*S!t9N&$0yMJ^8srYvZ zU+~l+e;6z2;oDeSGBnZ%jbP<&Y)J0I{*b-N-i&L2N04qiX6HG33+Kh(&ia@1h3Kv5 zgl>ygQ!wFRN(%(JLGmeyc6F76X<=Op_d3fCf zN!iMy$$Iv*^2ipo)wAb5^{70LePo`tdhUd-z3lHlm1)xBy~<&5o< z<#o>3$<~ZLcFV-9KF&1$V{#7nHs&zCAY(4~E=1tUY#ADO4E9*o*!SVb@R`~1fz4UF5fw~fUY zm$b!GOIrO2?1)c#$NeYNo+iia)q28zf>LB5?C&LR>~~EXv!V7ptlVSR;n!oOobK1W zXCSpU%o(Ftorv9VZ>)VluC1qTg~100Rr}||C>KszHn06AN0L&+e-gUi;ZdHOH=}=&TNlq^jYaI|(C-6F z;rf%e5~%wK>-P659&F{_+sb{PoGh`v*M}x1$3wG^z6I?B`@}Q#{m^r|7SH=3?EDw( zV{dDI)ZgzvgO-~0XQA!%Kw2F@Eid~iwCoj>bI?EFr$!*9p{X764|*59SN(&Wb%YQ4 z2mM2QVs9JK`VodEb(njUcN}u?fV~a7YGf{iF&PH&8&(ZoZ zQj=OJWQ6L+>n9dMjVI67Ph$3u_(w3dmj*TNiI-2*`~2e=-yZL2)Y!)w=Ng(A=JlZ8 z$JwgCFmL)V$Psw4e#$?A^e=Kg{9~m2Q%Vb&4m6H5YQyMH7a_wJ|e0)U8?T}pKCt7P9;l>&dNxAua zYaPjyD{F7wxX;>Cc}{+0p7ymDxwd=WP@K46KyHGIAV&eQf_6iao+gbHD&{oWziIKc zCGB%P+)Wo>_fO9wm%j&`bnosPaeumf_vISDWq`Y1Xsj$_@iOJLT3J{HSobV{>IIZJ zevc$L14(eUZS=~J0%sah;6$4gXwv_Gq`-@fw*FvJAYpHtHts~lV@gIeR-Li<3FyXI z@W)I3*~JmAZ^{b$6eR;R!pODX#wdK!KdZ(>N{9oH5!*1zXOTm#e#}Vy9P`SuPuSar zseiO&kSYy%w^6<65-7Z$H@R#7ur3@X%=<@Tw5t=02a{9|mf(gR!C@lnQc?bR>JUcKUv`CDL3 ze8k^8q*1c^G6gkgoW!6*^9;71oS647HlDT0#c0?Vu~&WU#nd+Im+DsiwfgJ+^9HzNPA{U-mEi>wL?99d%yvUu)Ir_?Izk$$hvjMV@I|eR6S>k|{pai+a8gXX zvXXLzf)x5&2&waW)E@UkE2n8`6+JV4fgUQ(&4TM?E7P^jubqcsFSGmdIR_Tc@?_?E z!FDQ{2>}DVo7uT@C&)@EQ+wFE>SUzZu$iJzbTW1YiOLl(Bg+f)oywvb3PMxW$sklM zPr{;AbtVN3+qK@#0;=xQwZZMMhh2As$Ev+bX3xHfs*74{y`QS@@*J27RZOY6x0afi zD50#(j@itPLZ%kY><~zw38EJZWxI02g;Pz~z3LQ;LB#cJm_H?keO+2{YY(ZO?zw#I zCM;=rS)=r_Y3rrA^G5D^ZGA^YpQ>kexVV1Fdzl@3GM}o~x*APm+QJ*SS(`-X@VRiP zwoBFV*kxEp9b|y4^wS!(vg-wD^<$f%$8?h2}V7x)dNdzp}b;!bgiLMp-i(a8ix9_82Aay zHSV)&58xUobEPs2qd{19IWT5-ShE@;5yQX*#g{->TeAf)Z zH{oV>QzDp05k-K<# zC^?EoZe0K3$`sdW#<-L^vY}(Ty6#brAcl8wx6Mv4x=v?o3G23yyWioQw)dfrSpz*4S85RC zp;c|&*et2`Qe*kAEPW}$)fU4mzKlEV;35&54kdxQ65@uL%jYXXZb{G#xEEEbPMNo~ z_jI4aOsEN&ElF1?&C5&y*%$_`X8hu`kLNb(uXe} zETGlP=N^7ka&Vw|*Su;Ob@Kdv*8kBwj}s_dsR(3NIkCO1U~6LO+xl@zTD^KyYloYy zy>}gTHu&(ztY%e%>CL=~9gZ|h0`G)y$G^;7r-L;(4crT*6nhkelw$pL-P1hX!$!n= znt4bvU*{PP6mA|#%{BB2a40I#qqd!x2|7az;SiR5~1^fjeJT+S#xe&6L&g{nJS&f^Hv(QIV@e_q4$$n?=G z??a@0&xMb%GyMol|LAqaeQEW&dYEyYWAkI*wuZ_4f`8@n+ zGB`ZKK?Ylql+rkKJv}xl{uPo*FAedpLBYwUAn1%1nB#{}9}6P#2tlq%OzpgVJ=Dk?baxnO3TVuDz;nDdwlC*N5pa;3mLd-5#W zcK*WAGcUt~NI((KxoEPiF_NodL$!lY(J8|uP`2G01@Ee$dzlrJE+{zRT4{Ly&;iF2 zl}T6y#Q)+je1{{@CB=g}1_j1sK8WTkMQZvyEsFeA@v|}&PVwGsHlM=|FMe5;eH`oT zG<<;sXE2C6IJNQ(7>DQgT-6Dp&2Gi=*S#Z^B3_M29XWIA%(?7yC(i9Tcj7s`d4%n= z&?%qXV9j{NJ6VEGNmm!gR%0D&3Mlc9P}_<(Xke79L8MqH=ZdgWlrd`JpRhS0o}xlH zHQ^;3A#$@ps8*WC6NLxV2KKNm&8gMUz8TOd18nGyMX5Xq^EBVL6^+`yC&u0pWCxQ;EIzywvy zAX=_uG1s>EO}4<4;xdwRJ-7(LVG2zBB+m)C{IvKjCWtng9t>g{VUCb&!EF-tW-I05 ztoUs<_unyX_y!v|_DLyZt%S5TW!3tS#u!CMu>fUcOxl33oUFAxOk6>olY)kdtg=uA zNmT+4R8FoGgkU`iV({y26axc52?q@;E>rUeVx?R-niZjdNpn{amQSpdG;L*3SeA<+ zM8*$+@vQ`X-||;y%G0jZM;c{hCG%kgI|5dh=5u+#iDOuVJY`wkY_J>()zr3#dE3? z8yIkaGwfOMKR6%Upshr;I0IL$!cC_r{(!N6NbnyC{s_QIuWCr|J&nM-f^2qHF~1eX zQ3GQHPs*0QFUP%KO~I<-23736g|Zl?wGt~w7ptqOgK;9kxYAY8wUW(fnW4KGY7}aX zHY4s;X_T2vIy^SdAj%`LoMqU_U$sVBwl}OJ%HNd*Z$(w>uQrJrRsxI*S4lFX)vtJ} z#UL(Os~M|uHV?;EiiW$&1lY3bSp0uJB5fNs?h)nv3pJV?fkC@!|JW21Q2>aGgx zdwH+Prq&R)@|Y zQw;%TnC5$lXj&gJ5+M^$OoQp*7zuaDgkk9Ey902dWS)e+F|-j*m1#X0>O-0&+$r0? z<{KvB@n`hud(k?)SP+WAikL80N?}VACUdWoU-OOrP!DR0-0eYKNNc2Z<4-pHPWpND z6i@o~jc5;FpkU6lEIX~ojEJ7NtDqU$p!eN1jZTPD1@a}#jc88-Jw{HI7j#?R_FS!U z2JVhvB>GU}2K`CYiaw$)I*|6Gb=y(Sc6|VS>VdPVc{c`E(j;2igAwY-Uk_^TG2nPg zx6%+tJgujVJ|k@m=JMWEXTjNmc+#r+VR zU*o*=y4H)!u63s7rDk7KS6*=5)Ta$G?xpOgbXYTNa~)derR~^aw+|OWy0>1^u7z&k z%?PwOfjW9{%K+EYdJ0aTy|`1*rMtix*k`-zJ#^rl(Zmg}&yLFyyKS_+-i`jO!860w zmO=>iduyeuXM#POGf-l0y|2Ee-tP_A2AoB^=xk75>kaxnQ$s8P4g*O#Ce#OHie5h@ zQw%mz(5$$|PsQ>a3j&p=#5NiTb~9p{`pP9Heh;6{sAgJ;g1&ft=*BA&^N zJ2x_No5@gRe43YK$B#oo%PT^4PH^l8GsnsmQ|F*6XL71_xJ`k+BNNC? z$K5GtnCv+Gsr!qD1%5jQK%OuA?zv&Yi!iuh!X1;kVOo9e9E*teIF7$hAk`kJ1yT1= zs*pTBcAg>sm(8r%@ zj5xw@f+I{l04D%x{G}po$rDIRsl)D%G;-K^2*t|w-_PY;OJYg=5dJ|v~hyU<8I zswALo%;;jBA8Lz(>%AWZ_s0CV;@*0=7Wd6Xjqc=j2r?@HX$fl_|2!DKY!eW_+_9wh zQ`#mj;W|;4d)QCliHVXTlaQEjQBObug*`YZ)w{f?7lXWpyE9z1NaM9H(vXBG@h#K? zXD_(%`$^Q8LOF3?pGvTPiDS4biA{t&eX2`3*(dSpLK<$c@ud))thA}Lyys2ZiQAB> zem5joSM8b8JSnObP5~h$Nt;^7!wfD}VfS}MUg@`qOye_VDw`!H5?4P>pl+R{oES5f zk20F$EPrrdnH!wC!r6=rF!8-X)NQ4xgSJ+Z;6lU&bJ=@*Z(vq%(;!!+m_W2zAXiAq zg?o{dTqTpFq$MW}t-W}J{tP$IR*JKDP3cCiifdRLC2VN+lUOHXJ?h{aZ*|M+Ue#?& zT@$GEP^NKIl;>wFuGDZWuKV#lRt=PG*NPNN*(y$vG_=YCIFAmk0G`l7RzOwj3)HWh|OR<^k4Fl4>1AcO%&Js(11(% zE%l94Af)k@8iRqI7y5)9!iyPU&sYq3#uWSk)GLW_UNWFx7=9SeG`b!2!*r_wZ}KG> zq@DsdgS%oi=QuYe4lz2I;dv*cW&svG&H}bXP%xH(^Kiw4ovvC@z7HMAS{#S{sBt1e zd97HOjc=pv5-jtLAwlMnTn9s@nS+ZB?tMdiAc5SfNsUDnza&rOW?Z+BD?eA95sLXt z7$;!Lac)+5KMI?_tGJI$hHN}kSc#Rk!1SUzDH&Su$3gJMAEtM=_%hqFYHq)XxUHO9 zn&s1G4D;EK*{wa$F@yJVt!#hC8^V-roJDrW{T2W(9H<{k-S8CL8bZ_wRot{GMtxn; zz5%CDsVmmaGG2;8I+>C>7GJ7jDj$-iQPSoGja+a`gR>xCC6D=L!zEp&;$0J80&k@; zKO{|fzjkfgwTMsth>UkA5*QDS%qj0s+Hco=d6l(t3%qiT)wtKQy48`mB|vZ|qED?d zT|bNR<(*#bsg6eIy) zK1SdZ6Q_u+kQDMq2<#{%O{{GAk$MbH(0FYQ95R3lB%EKQ1F*h`SEr|Vy9VbZ20WP8=__{j* zFm{+NjChiIiF}n2vvnTpdkP_*58uR$;H4TFI~E?7mjxYVlE=YS%EyL<`yktXeZYY9 z!jxXI2VMoQuC zmEo4}vWlN;t+*A1^JfrumTjQGm1pu0tQdN&uVYMQTROT<@!A?r9&S8s4o6;}2k$f( zo13J8r6AwAl9VYA&)6IrLxii=5Pno42<2?Mxu#=|4eLB(|;ssbw+%HV3C0D_=?YR z6b}(l5fR+61Z9gLvkH!UwZU!I*zjatoXbt(R*krO@mQ;%&?w|Z3HR5(G84BXq`I2j-Kj<<0QvVl| CB{+`& literal 0 HcmV?d00001 diff --git a/ip_lap/draw_landmark.py b/ip_lap/draw_landmark.py new file mode 100644 index 0000000..5e191a1 --- /dev/null +++ b/ip_lap/draw_landmark.py @@ -0,0 +1,197 @@ +"""MediaPipe solution drawing utils.""" +import math +from typing import List, Mapping, Optional, Tuple, Union +import cv2 +import dataclasses +import numpy as np +import tqdm +from mediapipe.framework.formats import landmark_pb2 +_PRESENCE_THRESHOLD = 0.5 +_VISIBILITY_THRESHOLD = 0.5 +_BGR_CHANNELS = 3 + +WHITE_COLOR = (224, 224, 224) +BLACK_COLOR = (0, 0, 0) +RED_COLOR = (0, 0, 255) +GREEN_COLOR = (0, 128, 0) +BLUE_COLOR = (255, 0, 0) + + +@dataclasses.dataclass +class DrawingSpec: + # Color for drawing the annotation. Default to the white color. + color: Tuple[int, int, int] = WHITE_COLOR + # Thickness for drawing the annotation. Default to 2 pixels. + thickness: int = 2 + # Circle radius. Default to 2 pixels. + circle_radius: int = 2 + +def _normalized_to_pixel_coordinates( + normalized_x: float, normalized_y: float, image_width: int, + image_height: int) -> Union[None, Tuple[int, int]]: + """Converts normalized value pair to pixel coordinates.""" + + # Checks if the float value is between 0 and 1. + def is_valid_normalized_value(value: float) -> bool: + return (value > 0 or math.isclose(0, value)) and (value < 1 or + math.isclose(1, value)) + + if not (is_valid_normalized_value(normalized_x) and + is_valid_normalized_value(normalized_y)): + # TODO: Draw coordinates even if it's outside of the image bounds. + return None + x_px = min(math.floor(normalized_x * image_width), image_width - 1) + y_px = min(math.floor(normalized_y * image_height), image_height - 1) + return x_px, y_px + + +FACEMESH_LIPS = frozenset([(61, 146), (146, 91), (91, 181), (181, 84), (84, 17), + (17, 314), (314, 405), (405, 321), (321, 375), + (375, 291), (61, 185), (185, 40), (40, 39), (39, 37), + (37, 0), (0, 267), + (267, 269), (269, 270), (270, 409), (409, 291), + (78, 95), (95, 88), (88, 178), (178, 87), (87, 14), + (14, 317), (317, 402), (402, 318), (318, 324), + (324, 308), (78, 191), (191, 80), (80, 81), (81, 82), + (82, 13), (13, 312), (312, 311), (311, 310), + (310, 415), (415, 308)]) + +FACEMESH_LEFT_EYE = frozenset([(263, 249), (249, 390), (390, 373), (373, 374), + (374, 380), (380, 381), (381, 382), (382, 362), + (263, 466), (466, 388), (388, 387), (387, 386), + (386, 385), (385, 384), (384, 398), (398, 362)]) + +FACEMESH_LEFT_IRIS = frozenset([(474, 475), (475, 476), (476, 477), + (477, 474)]) + +FACEMESH_LEFT_EYEBROW = frozenset([(276, 283), (283, 282), (282, 295), + (295, 285), (300, 293), (293, 334), + (334, 296), (296, 336)]) + +FACEMESH_RIGHT_EYE = frozenset([(33, 7), (7, 163), (163, 144), (144, 145), + (145, 153), (153, 154), (154, 155), (155, 133), + (33, 246), (246, 161), (161, 160), (160, 159), + (159, 158), (158, 157), (157, 173), (173, 133)]) + +FACEMESH_RIGHT_EYEBROW = frozenset([(46, 53), (53, 52), (52, 65), (65, 55), + (70, 63), (63, 105), (105, 66), (66, 107)]) + +FACEMESH_RIGHT_IRIS = frozenset([(469, 470), (470, 471), (471, 472), + (472, 469)]) + +FACEMESH_FACE_OVAL = frozenset([(389, 356), (356, 454), + (454, 323), (323, 361), (361, 288), (288, 397), + (397, 365), (365, 379), (379, 378), (378, 400), + (400, 377), (377, 152), (152, 148), (148, 176), + (176, 149), (149, 150), (150, 136), (136, 172), + (172, 58), (58, 132), (132, 93), (93, 234), + (234, 127), (127, 162)]) +#(10, 338), (338, 297), (297, 332), (332, 284),(284, 251), (251, 389) (162, 21), (21, 54),(54, 103), (103, 67), (67, 109), (109, 10) + +FACEMESH_NOSE= frozenset([(168, 6),(6,197),(197,195),(195,5),(5,4),\ + (4,45),(45,220),(220,115),(115,48),\ + (4,275),(275,440),(440,344),(344,278),]) +FACEMESH_FULL = frozenset().union(*[ + FACEMESH_LIPS, FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW, FACEMESH_RIGHT_EYE, + FACEMESH_RIGHT_EYEBROW, FACEMESH_FACE_OVAL,FACEMESH_NOSE +]) +connections=FACEMESH_FULL + +def summary_landmark(edge_set): + landmarks=set() + for a,b in edge_set: + landmarks.add(a) + landmarks.add(b) + return landmarks +all_landmark_idx=summary_landmark(FACEMESH_FULL) +pose_landmark_idx=\ +summary_landmark(FACEMESH_NOSE.union(*[FACEMESH_RIGHT_EYEBROW,FACEMESH_RIGHT_EYE,\ + FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW,])).union([162,127,234,93,389,356,454,323]) +content_landmark_idx= all_landmark_idx - pose_landmark_idx + + +def draw_landmarks( + image: np.ndarray, + landmark_list: List, + connections: Optional[List[Tuple[int, int]]] = None, + landmark_drawing_spec: Union[DrawingSpec, + Mapping[int, DrawingSpec]] = DrawingSpec( + color=RED_COLOR), + connection_drawing_spec: Union[DrawingSpec, + Mapping[Tuple[int, int], + DrawingSpec]] = DrawingSpec()): + """Draws the landmarks and the connections on the image. + + Args: + image: A three channel BGR image represented as numpy ndarray. + landmark_list: A normalized landmark list proto message to be annotated on + the image. + connections: A list of landmark index tuples that specifies how landmarks to + be connected in the drawing. + landmark_drawing_spec: Either a DrawingSpec object or a mapping from + hand landmarks to the DrawingSpecs that specifies the landmarks' drawing + settings such as color, line thickness, and circle radius. + If this argument is explicitly set to None, no landmarks will be drawn. + connection_drawing_spec: Either a DrawingSpec object or a mapping from + hand connections to the DrawingSpecs that specifies the + connections' drawing settings such as color and line thickness. + If this argument is explicitly set to None, no landmark connections will + be drawn. + + Raises: + ValueError: If one of the followings: + a) If the input image is not three channel BGR. + b) If any connetions contain invalid landmark index. + """ + if not landmark_list: + return + if image.shape[2] != _BGR_CHANNELS: + raise ValueError('Input image must contain three channel bgr data.') + image_rows, image_cols, _ = image.shape + idx_to_coordinates = {} + for landmark in landmark_list: + # if ((landmark.HasField('visibility') and + # landmark.visibility < _VISIBILITY_THRESHOLD) or + # (landmark.HasField('presence') and + # landmark.presence < _PRESENCE_THRESHOLD)): + # continue + idx=landmark.idx + landmark_px = _normalized_to_pixel_coordinates(landmark.x, landmark.y, + image_cols, image_rows) + if landmark_px: + idx_to_coordinates[idx] = landmark_px + + if connections: + num_landmarks = len(landmark_list) + # Draws the connections if the start and end landmarks are both visible. + for connection in connections: + start_idx = connection[0] + end_idx = connection[1] + # if not (0 <= start_idx < num_landmarks and 0 <= end_idx < num_landmarks): + # raise ValueError(f'Landmark index is out of range. Invalid connection ' + # f'from landmark #{start_idx} to landmark #{end_idx}.') + if start_idx in idx_to_coordinates and end_idx in idx_to_coordinates: + drawing_spec = connection_drawing_spec[connection] if isinstance( + connection_drawing_spec, Mapping) else connection_drawing_spec + # if start_idx in content_landmark and end_idx in content_landmark: + cv2.line(image, idx_to_coordinates[start_idx], + idx_to_coordinates[end_idx], drawing_spec.color, + drawing_spec.thickness) + return image + # Draws landmark points after finishing the connection lines, which is + # aesthetically better. + # if landmark_drawing_spec: + # for idx, landmark_px in idx_to_coordinates.items(): + # drawing_spec = landmark_drawing_spec[idx] if isinstance( + # landmark_drawing_spec, Mapping) else landmark_drawing_spec + # # White circle border + # circle_border_radius = max(drawing_spec.circle_radius + 1, + # int(drawing_spec.circle_radius * 1.2)) + # circle_border_radius=circle_border_radius*0.1 + # cv2.circle(image, landmark_px, circle_border_radius, WHITE_COLOR, + # drawing_spec.thickness) + # Fill color into the circle + + # cv2.circle(image, landmark_px, 1, + # drawing_spec.color, drawing_spec.thickness) + # cv2.putText(image,str(idx),landmark_px,cv2.FONT_HERSHEY_SIMPLEX,0.5,(255,0,0),1,cv2.LINE_AA) diff --git a/ip_lap/face_mask.py b/ip_lap/face_mask.py new file mode 100644 index 0000000..1b1b4de --- /dev/null +++ b/ip_lap/face_mask.py @@ -0,0 +1,50 @@ +import cv2 +import numpy as np +from typing import Any +import mediapipe as mp +from basicsr.utils.download_util import load_file_from_url + +class FaceMask: + def __init__(self) -> None: + BaseOptions = mp.tasks.BaseOptions + FaceLandmarker = mp.tasks.vision.FaceLandmarker + FaceLandmarkerOptions = mp.tasks.vision.FaceLandmarkerOptions + VisionRunningMode = mp.tasks.vision.RunningMode + + face_landmarks_detector_path = load_file_from_url(url="https://storage.googleapis.com/mediapipe-models/face_landmarker/face_landmarker/float16/latest/face_landmarker.task", + model_dir="weights", + file_name="face_landmarker.task") + options = FaceLandmarkerOptions( + base_options=BaseOptions(model_asset_path=face_landmarks_detector_path), + running_mode=VisionRunningMode.IMAGE) + self.face_landmarks_detector = FaceLandmarker.create_from_options(options) + + def __call__(self,image,*args: Any, **kwds: Any) -> Any: + """ + Calculate face mask from image. This is done by + + Args: + image: numpy array of an image + Returns: + A uint8 numpy array with the same height and width of the input image, containing a binary mask of the face in the image + """ + # initialize mask + mask = np.zeros((image.shape[0], image.shape[1]), dtype=np.uint8) + + # detect face landmarks + mp_image = mp.Image(image_format=mp.ImageFormat.SRGB, data=image) + detection = self.face_landmarks_detector.detect(mp_image) + + if len(detection.face_landmarks) == 0: + # no face detected - set mask to all of the image + mask[:] = 1 + return mask + + # extract landmarks coordinates + face_coords = np.array([[lm.x * image.shape[1], lm.y * image.shape[0]] for lm in detection.face_landmarks[0]]) + + # calculate convex hull from face coordinates + convex_hull = cv2.convexHull(face_coords.astype(np.float32)) + + # apply convex hull to mask + return cv2.fillPoly(mask, pts=[convex_hull.squeeze().astype(np.int32)], color=1) diff --git a/ip_lap/inference.py b/ip_lap/inference.py new file mode 100644 index 0000000..88f4cc4 --- /dev/null +++ b/ip_lap/inference.py @@ -0,0 +1,604 @@ +import os,cv2,torch,subprocess,platform +import mediapipe as mp +import numpy as np +from tqdm import tqdm +from .draw_landmark import draw_landmarks +import face_alignment +from .face_mask import FaceMask +from cuda_malloc import cuda_malloc_supported +from .models import Landmark_generator as Landmark_transformer,Renderer,audio + +NAME = "IP_LAP" + +# the following is the index sequence for fical landmarks detected by mediapipe +ori_sequence_idx = [162, 127, 234, 93, 132, 58, 172, 136, 150, 149, 176, 148, 152, 377, 400, 378, 379, 365, 397, 288, + 361, 323, 454, 356, 389, # + 70, 63, 105, 66, 107, 55, 65, 52, 53, 46, # + 336, 296, 334, 293, 300, 276, 283, 282, 295, 285, # + 168, 6, 197, 195, 5, # + 48, 115, 220, 45, 4, 275, 440, 344, 278, # + 33, 246, 161, 160, 159, 158, 157, 173, 133, 155, 154, 153, 145, 144, 163, 7, # + 362, 398, 384, 385, 386, 387, 388, 466, 263, 249, 390, 373, 374, 380, 381, 382, # + 61, 185, 40, 39, 37, 0, 267, 269, 270, 409, 291, 375, 321, 405, 314, 17, 84, 181, 91, 146, # + 78, 191, 80, 81, 82, 13, 312, 311, 310, 415, 308, 324, 318, 402, 317, 14, 87, 178, 88, 95] + +# the following is the connections of landmarks for drawing sketch image +FACEMESH_LIPS = frozenset([(61, 146), (146, 91), (91, 181), (181, 84), (84, 17), + (17, 314), (314, 405), (405, 321), (321, 375), + (375, 291), (61, 185), (185, 40), (40, 39), (39, 37), + (37, 0), (0, 267), + (267, 269), (269, 270), (270, 409), (409, 291), + (78, 95), (95, 88), (88, 178), (178, 87), (87, 14), + (14, 317), (317, 402), (402, 318), (318, 324), + (324, 308), (78, 191), (191, 80), (80, 81), (81, 82), + (82, 13), (13, 312), (312, 311), (311, 310), + (310, 415), (415, 308)]) +FACEMESH_LEFT_EYE = frozenset([(263, 249), (249, 390), (390, 373), (373, 374), + (374, 380), (380, 381), (381, 382), (382, 362), + (263, 466), (466, 388), (388, 387), (387, 386), + (386, 385), (385, 384), (384, 398), (398, 362)]) +FACEMESH_LEFT_EYEBROW = frozenset([(276, 283), (283, 282), (282, 295), + (295, 285), (300, 293), (293, 334), + (334, 296), (296, 336)]) +FACEMESH_RIGHT_EYE = frozenset([(33, 7), (7, 163), (163, 144), (144, 145), + (145, 153), (153, 154), (154, 155), (155, 133), + (33, 246), (246, 161), (161, 160), (160, 159), + (159, 158), (158, 157), (157, 173), (173, 133)]) +FACEMESH_RIGHT_EYEBROW = frozenset([(46, 53), (53, 52), (52, 65), (65, 55), + (70, 63), (63, 105), (105, 66), (66, 107)]) +FACEMESH_FACE_OVAL = frozenset([(389, 356), (356, 454), + (454, 323), (323, 361), (361, 288), (288, 397), + (397, 365), (365, 379), (379, 378), (378, 400), + (400, 377), (377, 152), (152, 148), (148, 176), + (176, 149), (149, 150), (150, 136), (136, 172), + (172, 58), (58, 132), (132, 93), (93, 234), + (234, 127), (127, 162)]) +FACEMESH_NOSE = frozenset([(168, 6), (6, 197), (197, 195), (195, 5), (5, 4), + (4, 45), (45, 220), (220, 115), (115, 48), + (4, 275), (275, 440), (440, 344), (344, 278), ]) +FACEMESH_CONNECTION = frozenset().union(*[ + FACEMESH_LIPS, FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW, FACEMESH_RIGHT_EYE, + FACEMESH_RIGHT_EYEBROW, FACEMESH_FACE_OVAL, FACEMESH_NOSE +]) + +full_face_landmark_sequence = [*list(range(0, 4)), *list(range(21, 25)), *list(range(25, 91)), #upper-half face + *list(range(4, 21)), # jaw + *list(range(91, 131))] # mouth + +class LandmarkDict(dict):# Makes a dictionary that behave like an object to represent each landmark + def __init__(self, idx, x, y): + self['idx'] = idx + self['x'] = x + self['y'] = y + def __getattr__(self, name): + try: + return self[name] + except: + raise AttributeError(name) + def __setattr__(self, name, value): + self[name] = value + +class IP_LAP_infer: + + def __init__(self,T=5,Nl=15,ref_img_N=25, + img_size=128,mel_step_size=16, + face_det_batch_size=4, + checkpoints_path=""): + self.T = T + self.Nl = Nl + self.ref_img_N = ref_img_N + self.img_size = img_size + self.mel_step_size = mel_step_size + self.face_det_batch_size = face_det_batch_size + self.pads = [100,100,100,100] + self.device = "cuda" if cuda_malloc_supported() else "cpu" + + self.mp_face_mesh = mp.solutions.face_mesh + self.drawing_spec = mp.solutions.drawing_utils.DrawingSpec(thickness=1, circle_radius=1) + self.lip_index = [0, 17] + self.all_landmarks_idx = self.summarize_landmark(FACEMESH_CONNECTION) + self.pose_landmark_idx = \ + self.summarize_landmark(FACEMESH_NOSE.union(*[FACEMESH_RIGHT_EYEBROW, FACEMESH_RIGHT_EYE, + FACEMESH_LEFT_EYE, FACEMESH_LEFT_EYEBROW, ])).union( + [162, 127, 234, 93, 389, 356, 454, 323]) + # pose landmarks are landmarks of the upper-half face(eyes,nose,cheek) that represents the pose information + + self.content_landmark_idx = self.all_landmarks_idx - self.pose_landmark_idx + # content_landmark include landmarks of lip and jaw which are inferred from audio + + + landmark_gen_checkpoint_path = os.path.join(checkpoints_path, "landmarkgenerator_checkpoint.pth") + renderer_checkpoint_path = os.path.join(checkpoints_path, "renderer_checkpoint.pth") + self.landmark_generator_model = self.load_model( + model=Landmark_transformer(T=self.T, d_model=512, nlayers=4, nhead=4, dim_feedforward=1024, dropout=0.1), + path=landmark_gen_checkpoint_path) + self.renderer = self.load_model(model=Renderer(), path=renderer_checkpoint_path) + + self.fa = face_alignment.FaceAlignment(face_alignment.LandmarksType.TWO_D, flip_input=False, device=self.device) + + self.face_mask = FaceMask() + + def __call__(self,video_file, audio_file, outfile): + temp_dir = os.path.join(os.path.dirname(outfile), NAME) + if not os.path.exists(temp_dir): os.makedirs(temp_dir, exist_ok=True) + ##(1) Reading input video frames ### + print(f'[Step 1]Reading video frames ... from {video_file}', NAME) + if not os.path.isfile(video_file): + raise ValueError('the input video file does not exist') + elif video_file.split('.')[1] in ['jpg', 'png', 'jpeg']: #if input a single image for testing + ori_background_frames = [cv2.imread(video_file)] + else: + video_stream = cv2.VideoCapture(video_file) + fps = video_stream.get(cv2.CAP_PROP_FPS) + if fps != 25: + print(" input video fps:", fps,',converting to 25fps...') + tmp_file = '{}/temp_25fps.mp4'.format(temp_dir) + if os.path.exists(tmp_file): os.remove(tmp_file) + print(tmp_file) + command = 'ffmpeg -y -i ' + video_file + f' -r 25 {tmp_file}' + subprocess.call(command, shell=platform.system() != 'Windows',stdout=subprocess.PIPE, stderr=subprocess.STDOUT) + video_file = '{}/temp_25fps.mp4'.format(temp_dir) + video_stream.release() + video_stream = cv2.VideoCapture(video_file) + fps = video_stream.get(cv2.CAP_PROP_FPS) + assert fps == 25 + + ori_background_frames = [] #input videos frames (includes background as well as face) + frame_idx = 0 + while 1: + still_reading, frame = video_stream.read() + if not still_reading: + video_stream.release() + break + ori_background_frames.append(frame) + frame_idx = frame_idx + 1 + input_vid_len = len(ori_background_frames) + + ##(2) Extracting audio#### + print(f'[Step 2]Extracting audio ... from {audio_file}', NAME) + if not audio_file.endswith('.wav'): + command = 'ffmpeg -y -i {} -strict -2 {}'.format(audio_file, '{}/temp.wav'.format(temp_dir)) + subprocess.call(command, shell=platform.system() != 'Windows', stdout=subprocess.PIPE, stderr=subprocess.STDOUT) + audio_file = '{}/temp.wav'.format(temp_dir) + wav = audio.load_wav(audio_file, 16000) + mel = audio.melspectrogram(wav) # (H,W) extract mel-spectrum + ##read audio mel into list### + mel_chunks = [] # each mel chunk correspond to 5 video frames, used to generate one video frame + mel_idx_multiplier = 80. / fps + mel_chunk_idx = 0 + while 1: + start_idx = int(mel_chunk_idx * mel_idx_multiplier) + if start_idx + self.mel_step_size > len(mel[0]): + break + mel_chunks.append(mel[:, start_idx: start_idx + self.mel_step_size]) # mel for generate one video frame + mel_chunk_idx += 1 + + print('[Step 3]detect facial using face detection tool', NAME) + ori_face_frames, ori_face_coords = self.face_detect(ori_background_frames) + # print(len(ori_face_frames)) + import gc; gc.collect(); torch.cuda.empty_cache() + + ##(3) detect facial landmarks using mediapipe tool + print('[Step 4]detect facial landmarks using mediapipe tool', NAME) + boxes = [] #bounding boxes of human face + lip_dists = [] #lip dists + #we define the lip dist(openness): distance between the midpoints of the upper lip and lower lip + face_crop_results = [] + all_pose_landmarks, all_content_landmarks = [], [] #content landmarks include lip and jaw landmarks + with self.mp_face_mesh.FaceMesh(static_image_mode=True, max_num_faces=1, refine_landmarks=True, + min_detection_confidence=0) as face_mesh: + # (1) get bounding boxes and lip dist + for frame_idx, full_frame in tqdm(enumerate(ori_face_frames),total=input_vid_len, + desc="get bounding boxes and lip dist"): + h, w = full_frame.shape[0], full_frame.shape[1] + results = face_mesh.process(cv2.cvtColor(full_frame, cv2.COLOR_BGR2RGB)) + if not results.multi_face_landmarks: + raise NotImplementedError # not detect face + face_landmarks = results.multi_face_landmarks[0] + + ## calculate the lip dist + dx = face_landmarks.landmark[self.lip_index[0]].x - face_landmarks.landmark[self.lip_index[1]].x + dy = face_landmarks.landmark[self.lip_index[0]].y - face_landmarks.landmark[self.lip_index[1]].y + dist = np.linalg.norm((dx, dy)) + lip_dists.append((frame_idx, dist)) + + # (1)get the marginal landmarks to crop face + x_min,x_max,y_min,y_max = 999,-999,999,-999 + for idx, landmark in enumerate(face_landmarks.landmark): + if idx in self.all_landmarks_idx: + if landmark.x < x_min: + x_min = landmark.x + if landmark.x > x_max: + x_max = landmark.x + if landmark.y < y_min: + y_min = landmark.y + if landmark.y > y_max: + y_max = landmark.y + ##########plus some pixel to the marginal region########## + #note:the landmarks coordinates returned by mediapipe range 0~1 + plus_pixel = 25 + x_min = max(x_min - plus_pixel / w, 0) + x_max = min(x_max + plus_pixel / w, 1) + + y_min = max(y_min - plus_pixel / h, 0) + y_max = min(y_max + plus_pixel / h, 1) + y1, y2, x1, x2 = int(y_min * h), int(y_max * h), int(x_min * w), int(x_max * w) + boxes.append([y1, y2, x1, x2]) + boxes = np.array(boxes) + + # (2)croppd face + face_crop_results = [[image[y1:y2, x1:x2], (y1, y2, x1, x2)] \ + for image, (y1, y2, x1, x2) in zip(ori_face_frames, boxes)] + + # (3)detect facial landmarks + for frame_idx, full_frame in tqdm(enumerate(ori_face_frames),total=input_vid_len, + desc="detect facial landmarks"): + h, w = full_frame.shape[0], full_frame.shape[1] + results = face_mesh.process(cv2.cvtColor(full_frame, cv2.COLOR_BGR2RGB)) + if not results.multi_face_landmarks: + raise ValueError("not detect face in some frame!") # not detect + face_landmarks = results.multi_face_landmarks[0] + + + + pose_landmarks, content_landmarks = [], [] + for idx, landmark in enumerate(face_landmarks.landmark): + if idx in self.pose_landmark_idx: + pose_landmarks.append((idx, w * landmark.x, h * landmark.y)) + if idx in self.content_landmark_idx: + content_landmarks.append((idx, w * landmark.x, h * landmark.y)) + + # normalize landmarks to 0~1 + y_min, y_max, x_min, x_max = face_crop_results[frame_idx][1] #bounding boxes + pose_landmarks = [ \ + [idx, (x - x_min) / (x_max - x_min), (y - y_min) / (y_max - y_min)] for idx, x, y in pose_landmarks] + content_landmarks = [ \ + [idx, (x - x_min) / (x_max - x_min), (y - y_min) / (y_max - y_min)] for idx, x, y in content_landmarks] + all_pose_landmarks.append(pose_landmarks) + all_content_landmarks.append(content_landmarks) + + all_pose_landmarks = self.get_smoothened_landmarks(all_pose_landmarks, windows_T=1) + all_content_landmarks=self.get_smoothened_landmarks(all_content_landmarks,windows_T=1) + + ##randomly select N_l reference landmarks for landmark transformer## + print("randomly select N_l reference landmarks for landmark transformer", NAME) + dists_sorted = sorted(lip_dists, key=lambda x: x[1]) + lip_dist_idx = np.asarray([idx for idx, dist in dists_sorted]) #the frame idxs sorted by lip openness + + Nl_idxs = [lip_dist_idx[int(i)] for i in torch.linspace(0, input_vid_len - 1, steps=self.Nl)] + Nl_pose_landmarks, Nl_content_landmarks = [], [] #Nl_pose + Nl_content=Nl reference landmarks + for reference_idx in Nl_idxs: + frame_pose_landmarks = all_pose_landmarks[reference_idx] + frame_content_landmarks = all_content_landmarks[reference_idx] + Nl_pose_landmarks.append(frame_pose_landmarks) + Nl_content_landmarks.append(frame_content_landmarks) + + Nl_pose = torch.zeros((self.Nl, 2, 74)) # 74 landmark + Nl_content = torch.zeros((self.Nl, 2, 57)) # 57 landmark + for idx in range(self.Nl): + #arrange the landmark in a certain order, since the landmark index returned by mediapipe is is chaotic + Nl_pose_landmarks[idx] = sorted(Nl_pose_landmarks[idx], + key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0])) + Nl_content_landmarks[idx] = sorted(Nl_content_landmarks[idx], + key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0])) + + Nl_pose[idx, 0, :] = torch.FloatTensor( + [Nl_pose_landmarks[idx][i][1] for i in range(len(Nl_pose_landmarks[idx]))]) # x + Nl_pose[idx, 1, :] = torch.FloatTensor( + [Nl_pose_landmarks[idx][i][2] for i in range(len(Nl_pose_landmarks[idx]))]) # y + Nl_content[idx, 0, :] = torch.FloatTensor( + [Nl_content_landmarks[idx][i][1] for i in range(len(Nl_content_landmarks[idx]))]) # x + Nl_content[idx, 1, :] = torch.FloatTensor( + [Nl_content_landmarks[idx][i][2] for i in range(len(Nl_content_landmarks[idx]))]) # y + Nl_content = Nl_content.unsqueeze(0) # (1,Nl, 2, 57) + Nl_pose = Nl_pose.unsqueeze(0) # (1,Nl,2,74) + + + ##select reference images and draw sketches for rendering according to lip openness## + print("select reference images and draw sketches for rendering according to lip openness", NAME) + ref_img_idx = [int(lip_dist_idx[int(i)]) for i in torch.linspace(0, input_vid_len - 1, steps=self.ref_img_N)] + ref_imgs = [face_crop_results[idx][0] for idx in ref_img_idx] + ## (N,H,W,3) + ref_img_pose_landmarks, ref_img_content_landmarks = [], [] + for idx in ref_img_idx: + ref_img_pose_landmarks.append(all_pose_landmarks[idx]) + ref_img_content_landmarks.append(all_content_landmarks[idx]) + + ref_img_pose = torch.zeros((self.ref_img_N, 2, 74)) # 74 landmark + ref_img_content = torch.zeros((self.ref_img_N, 2, 57)) # 57 landmark + + for idx in range(self.ref_img_N): + ref_img_pose_landmarks[idx] = sorted(ref_img_pose_landmarks[idx], + key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0])) + ref_img_content_landmarks[idx] = sorted(ref_img_content_landmarks[idx], + key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0])) + ref_img_pose[idx, 0, :] = torch.FloatTensor( + [ref_img_pose_landmarks[idx][i][1] for i in range(len(ref_img_pose_landmarks[idx]))]) # x + ref_img_pose[idx, 1, :] = torch.FloatTensor( + [ref_img_pose_landmarks[idx][i][2] for i in range(len(ref_img_pose_landmarks[idx]))]) # y + + ref_img_content[idx, 0, :] = torch.FloatTensor( + [ref_img_content_landmarks[idx][i][1] for i in range(len(ref_img_content_landmarks[idx]))]) # x + ref_img_content[idx, 1, :] = torch.FloatTensor( + [ref_img_content_landmarks[idx][i][2] for i in range(len(ref_img_content_landmarks[idx]))]) # y + + ref_img_full_face_landmarks = torch.cat([ref_img_pose, ref_img_content], dim=2).cpu().numpy() # (N,2,131) + ref_img_sketches = [] + for frame_idx in range(ref_img_full_face_landmarks.shape[0]): # N + full_landmarks = ref_img_full_face_landmarks[frame_idx] # (2,131) + h, w = ref_imgs[frame_idx].shape[0], ref_imgs[frame_idx].shape[1] + drawn_sketech = np.zeros((int(h * self.img_size / min(h, w)), int(w * self.img_size / min(h, w)), 3)) + mediapipe_format_landmarks = [LandmarkDict(ori_sequence_idx[full_face_landmark_sequence[idx]], full_landmarks[0, idx], + full_landmarks[1, idx]) for idx in range(full_landmarks.shape[1])] + drawn_sketech = draw_landmarks(drawn_sketech, mediapipe_format_landmarks, connections=FACEMESH_CONNECTION, + connection_drawing_spec=self.drawing_spec) + drawn_sketech = cv2.resize(drawn_sketech, (self.img_size, self.img_size)) # (128, 128, 3) + ref_img_sketches.append(drawn_sketech) + ref_img_sketches = torch.FloatTensor(np.asarray(ref_img_sketches) / 255.0).cuda().unsqueeze(0).permute(0, 1, 4, 2, 3) + # (1,N, 3, 128, 128) + ref_imgs = [cv2.resize(face.copy(), (self.img_size, self.img_size)) for face in ref_imgs] + ref_imgs = torch.FloatTensor(np.asarray(ref_imgs) / 255.0).unsqueeze(0).permute(0, 1, 4, 2, 3).cuda() + # (1,N,3,H,W) + + ##prepare output video strame## + frame_h, frame_w = ori_background_frames[0].shape[:-1] + ''' + out_stream = cv2.VideoWriter('{}/result.avi'.format(temp_dir), cv2.VideoWriter_fourcc(*'DIVX'), fps, + (frame_w, frame_h)) # +frame_h*3 + ''' + out_stream = cv2.VideoWriter(outfile, cv2.VideoWriter_fourcc(*'mp4v'), fps, + (frame_w, frame_h)) # +frame_h*3 + + ##generate final face image and output video## + input_mel_chunks_len = len(mel_chunks) + input_frame_sequence = torch.arange(input_vid_len).tolist() + #the input template video may be shorter than audio + #in this case we repeat the input template video as following + num_of_repeat=input_mel_chunks_len//input_vid_len+1 + input_frame_sequence = input_frame_sequence + list(reversed(input_frame_sequence)) + input_frame_sequence=input_frame_sequence*((num_of_repeat+1)//2) + file_num = 0 + for batch_idx, batch_start_idx in tqdm(enumerate(range(0, input_mel_chunks_len-2, 1)), + total=len(range(0, input_mel_chunks_len-2, 1)), desc="[IP_LAP] [Step 5]Lipsync..."): + T_input_frame, T_ori_face_coordinates = [], [] + #note: input_frame include background as well as face + T_mel_batch, T_crop_face,T_pose_landmarks = [], [],[] + + A_input_frame, A_ori_face_coordinates = [], [] + + # (1) for each batch of T frame, generate corresponding landmarks using landmark generator + for mel_chunk_idx in range(batch_start_idx, batch_start_idx + self.T): # for each T frame + # 1 input audio + T_mel_batch.append(mel_chunks[max(0, mel_chunk_idx - 2)]) + + # 2.input face + input_frame_idx = int(input_frame_sequence[mel_chunk_idx]) + face, coords = face_crop_results[input_frame_idx] + T_crop_face.append(face) + T_ori_face_coordinates.append((face, coords)) ##input face + # 3.pose landmarks + T_pose_landmarks.append(all_pose_landmarks[input_frame_idx]) + # 3.face background + T_input_frame.append(ori_face_frames[input_frame_idx].copy()) + # 4.frame background + A_ori_face_coordinates.append(ori_face_coords[input_frame_idx]) + A_input_frame.append(ori_background_frames[input_frame_idx].copy()) + + T_mels = torch.FloatTensor(np.asarray(T_mel_batch)).unsqueeze(1).unsqueeze(0) # 1,T,1,h,w + #prepare pose landmarks + T_pose = torch.zeros((self.T, 2, 74)) # 74 landmark + for idx in range(self.T): + T_pose_landmarks[idx] = sorted(T_pose_landmarks[idx], + key=lambda land_tuple: ori_sequence_idx.index(land_tuple[0])) + T_pose[idx, 0, :] = torch.FloatTensor( + [T_pose_landmarks[idx][i][1] for i in range(len(T_pose_landmarks[idx]))]) # x + T_pose[idx, 1, :] = torch.FloatTensor( + [T_pose_landmarks[idx][i][2] for i in range(len(T_pose_landmarks[idx]))]) # y + T_pose = T_pose.unsqueeze(0) # (1,T, 2,74) + + #landmark generator inference + Nl_pose, Nl_content = Nl_pose.cuda(), Nl_content.cuda() # (Nl,2,74) (Nl,2,57) + T_mels, T_pose = T_mels.cuda(), T_pose.cuda() + with torch.no_grad(): # require (1,T,1,hv,wv)(1,T,2,74)(1,T,2,57) + predict_content = self.landmark_generator_model(T_mels, T_pose, Nl_pose, Nl_content) # (1*T,2,57) + T_pose = torch.cat([T_pose[i] for i in range(T_pose.size(0))], dim=0) # (1*T,2,74) + T_predict_full_landmarks = torch.cat([T_pose, predict_content], dim=2).cpu().numpy() # (1*T,2,131) + + #1.draw target sketch + T_target_sketches = [] + for frame_idx in range(self.T): + full_landmarks = T_predict_full_landmarks[frame_idx] # (2,131) + h, w = T_crop_face[frame_idx].shape[0], T_crop_face[frame_idx].shape[1] + drawn_sketech = np.zeros((int(h * self.img_size / min(h, w)), int(w * self.img_size / min(h, w)), 3)) + mediapipe_format_landmarks = [LandmarkDict(ori_sequence_idx[full_face_landmark_sequence[idx]] + , full_landmarks[0, idx], full_landmarks[1, idx]) for idx in + range(full_landmarks.shape[1])] + drawn_sketech = draw_landmarks(drawn_sketech, mediapipe_format_landmarks, connections=FACEMESH_CONNECTION, + connection_drawing_spec=self.drawing_spec) + drawn_sketech = cv2.resize(drawn_sketech, (self.img_size, self.img_size)) # (128, 128, 3) + if frame_idx == 2: + show_sketch = cv2.resize(drawn_sketech, (frame_w, frame_h)).astype(np.uint8) + T_target_sketches.append(torch.FloatTensor(drawn_sketech) / 255) + T_target_sketches = torch.stack(T_target_sketches, dim=0).permute(0, 3, 1, 2) # (T,3,128, 128) + target_sketches = T_target_sketches.unsqueeze(0).cuda() # (1,T,3,128, 128) + + # 2.lower-half masked face + ori_face_img = torch.FloatTensor(cv2.resize(T_crop_face[2], (self.img_size, self.img_size)) / 255).permute(2, 0, 1).unsqueeze( + 0).unsqueeze(0).cuda() #(1,1,3,H, W) + + # 3. render the full face + # require (1,1,3,H,W) (1,T,3,H,W) (1,N,3,H,W) (1,N,3,H,W) (1,1,1,h,w) + # return (1,3,H,W) + with torch.no_grad(): + generated_face, _, _, _ = self.renderer(ori_face_img, target_sketches, ref_imgs, ref_img_sketches, + T_mels[:, 2].unsqueeze(0)) # T=1 + gen_face = (generated_face.squeeze(0).permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8) # (H,W,3) + + # 4. paste each generated face + y1, y2, x1, x2 = T_ori_face_coordinates[2][1] # coordinates of face bounding box + original_background = T_input_frame[2].copy() + T_input_frame[2][y1:y2, x1:x2] = cv2.resize(gen_face,(x2 - x1, y2 - y1)) #resize and paste generated face + # 5. post-process + full_face = self.merge_face_contour_only(original_background, T_input_frame[2], T_ori_face_coordinates[2][1],self.fa) #(H,W,3) + # 6.output + # full = np.concatenate([show_sketch, full], axis=1) + # print(f"full_face.shape{full_face.shape}") + ori_x1, ori_y1, ori_x2, ori_y2 = A_ori_face_coordinates[2] + # print(ori_face_coords[file_num]) + full_frame = A_input_frame[2] + #full_mask = np.zeros_like(full_frame) + # print(full_frame.shape) + if ori_x1 != -1: + p = cv2.resize(full_face.astype(np.uint8), (ori_x2 - ori_x1, ori_y2 - ori_y1)) + # print(p.shape) + full_frame[ori_y1:ori_y2, ori_x1:ori_x2] = p + # height, width = full_frame.shape[:2] + # img = self.Laplacian_Pyramid_Blending_with_mask(full_frame, ori_background_frames[file_num], full_mask[:, :, 0], 6) + # pp = np.uint8(cv2.resize(np.clip(img, 0 ,255), (width, height))) + mask = self.face_mask(p) + full_frame[ori_y1:ori_y2, ori_x1:ori_x2] = full_frame[ori_y1:ori_y2, ori_x1:ori_x2] * (1 - mask[..., None]) + p * mask[..., None] + full = full_frame.copy() + + out_stream.write(full) + + try: + # cv2.imwrite(temp_frame_paths[batch_idx+2],full) + file_num += 1 + except: + pass + + if batch_idx == 0: + out_stream.write(full) + out_stream.write(full) + # cv2.imwrite(temp_frame_paths[batch_idx],full) + # cv2.imwrite(temp_frame_paths[batch_idx+1],full) + + out_stream.release() + # command = 'ffmpeg -y -i {} -i {} -strict -2 -q:v 1 {}'.format(voice_file, '{}/result.avi'.format(temp_dir), outfile) + # subprocess.call(command, shell=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT) + print(f"succeed output results to:{outfile}", NAME) + + + def face_detect(self, images): + + batch_size = self.face_det_batch_size + + while 1: + predictions = [] + try: + for i in tqdm(range(0, len(images), batch_size)): + imgs = np.array(images[i:i + batch_size]) + imgs_numpy = imgs.transpose(0, 3, 1, 2) + image_batch = torch.from_numpy(imgs_numpy.copy()) + _, _, bboxes =self.fa.get_landmarks_from_batch(image_batch,return_bboxes=True) + predictions.extend(bboxes) + except RuntimeError: + if batch_size == 1: + raise RuntimeError('Image too big to run face detection on GPU. Please use the --resize_factor argument') + batch_size //= 2 + print('Recovering from OOM error; New batch size: {}'.format(batch_size)) + continue + break + + results = [] + pady1, pady2, padx1, padx2 = self.pads + for rect, image in zip(predictions, images): + + if rect is None: + # cv2.imwrite('temp/faulty_frame.jpg', image) # check this frame where the face was not detected. + # results.append([-1,-1,-1,-1]) + raise ValueError('Face not detected! Ensure the video contains a face in all the frames.') + else: + rect = rect[0] + rect = np.clip(rect, 0, None) + x1_0, y1_0, x2_0, y2_0 = map(int, rect[:-1]) + + y1 = max(0, y1_0 - pady1) + y2 = min(image.shape[0], y2_0 + pady2) + x1 = max(0, x1_0 - padx1) + x2 = min(image.shape[1], x2_0 + padx2) + + results.append([x1, y1, x2, y2]) + + boxes = np.array(results) + + faces = [image[y1: y2, x1:x2] for image, (x1, y1, x2, y2) in zip(images, boxes)] + return faces, boxes + + + def merge_face_contour_only(self,src_frame, generated_frame, face_region_coord, fa): #function used in post-process + """Merge the face from generated_frame into src_frame + """ + input_img = src_frame + y1, y2, x1, x2 = 0, 0, 0, 0 + if face_region_coord is not None: + y1, y2, x1, x2 = face_region_coord + input_img = src_frame[y1:y2, x1:x2] + ### 1) Detect the facial landmarks + try: + preds = fa.get_landmarks(input_img)[0] # 68x2 + except: + preds = np.int64(-1 * np.ones((68,2))) + if face_region_coord is not None: + preds += np.array([x1, y1]) + lm_pts = preds.astype(int) + contour_idx = list(range(0, 17)) + list(range(17, 27))[::-1] + contour_pts = lm_pts[contour_idx] + ### 2) Make the landmark region mark image + mask_img = np.zeros((src_frame.shape[0], src_frame.shape[1], 1), np.uint8) + cv2.fillConvexPoly(mask_img, contour_pts, 255) + ### 3) Do swap + img = self.swap_masked_region(src_frame, generated_frame, mask=mask_img) + return img + + def swap_masked_region(self,target_img, src_img, mask): #function used in post-process + """From src_img crop masked region to replace corresponding masked region + in target_img + """ # swap_masked_region(src_frame, generated_frame, mask=mask_img) + mask_img = cv2.GaussianBlur(mask, (21, 21), 11) + mask1 = mask_img / 255 + mask1 = np.tile(np.expand_dims(mask1, axis=2), (1, 1, 3)) + img = src_img * mask1 + target_img * (1 - mask1) + return img.astype(np.uint8) + + # smooth landmarks + def get_smoothened_landmarks(self,all_landmarks, windows_T=1): + for i in range(len(all_landmarks)): # frame i + if i + windows_T > len(all_landmarks): + window = all_landmarks[len(all_landmarks) - windows_T:] + else: + window = all_landmarks[i: i + windows_T] + ##### + for j in range(len(all_landmarks[i])): # landmark j + all_landmarks[i][j][1] = np.mean([frame_landmarks[j][1] for frame_landmarks in window]) # x + all_landmarks[i][j][2] = np.mean([frame_landmarks[j][2] for frame_landmarks in window]) # y + return all_landmarks + + def load_model(self, model, path): + print("Load checkpoint from: {}".format(path)) + checkpoint = self._load(path) + s = checkpoint["state_dict"] + new_s = {} + for k, v in s.items(): + if k[:6] == 'module': + new_k=k.replace('module.', '', 1) + else: + new_k =k + new_s[new_k] = v + model.load_state_dict(new_s) + model = model.to(self.device) + return model.eval() + + def _load(self,checkpoint_path): + if self.device == 'cuda': + checkpoint = torch.load(checkpoint_path) + else: + checkpoint = torch.load(checkpoint_path, map_location=lambda storage, loc: storage) + return checkpoint + + def summarize_landmark(self, edge_set): # summarize all ficial landmarks used to construct edge + landmarks = set() + for a, b in edge_set: + landmarks.add(a) + landmarks.add(b) + return landmarks \ No newline at end of file diff --git a/ip_lap/models/__init__.py b/ip_lap/models/__init__.py new file mode 100644 index 0000000..6e4582c --- /dev/null +++ b/ip_lap/models/__init__.py @@ -0,0 +1,4 @@ +from . import audio +from .landmark_generator import Landmark_generator +from .video_renderer import Renderer +from .pix2pixHD_disc import define_D diff --git a/ip_lap/models/__pycache__/__init__.cpython-310.pyc b/ip_lap/models/__pycache__/__init__.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..3d3157593c59c7c9e3c7b9e5a05f43537844c4ad GIT binary patch literal 347 zcmZXP&q~8U5XN`&rwNjZ6nuyr>VkL_QLJF0NNMqs%MiL7UEDumHxct7K7_9nkHv$p z;K@l)5FD6qe)A0=S8FC~P-V?D-jrtm(#Qtjr0)9k9L-jVi{X<#Maf7;GkQe7 E0Tx|aw*UYD literal 0 HcmV?d00001 diff --git a/ip_lap/models/__pycache__/audio.cpython-310.pyc b/ip_lap/models/__pycache__/audio.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..a46f8ebc766eff32faebba872fadb3fcb18d1551 GIT binary patch literal 6143 zcma)AU2Ggz6`p(N?(BHIcARA6IL_ZTb&_t|#0gE)zqC%1v}xMJP3!cBZOe4LcXoH| z{p+1^9J_0%>L7S2LZUAqkZgD%F9;+ac;JOnN`aPtsnUQ1g$FP#0#OhW0?c>r?5=;3 z0JFO1o_p?{JLjBx{_fU$dQt|Sw?5fl+_Kv+eoKSZPXdGEXu%uIFr*<_#mL7!YxO*D zg3c?Vn#d;@&6yRen#?CP&MJ1^K5a-L6PFE{nCJN}uS*tvLs~NV876Jn1(=fEfN9wS z*ef>x_Q{Qan>4;zW`J*z46t9X->Tt&hTG&I=-V~EL+%8=OYR2D${PUp$QuENwUB90oij4+DNk-UfI?!`t|A}nU~K1-zJ}x&!Hcb zACV6F?ecl)qTeCMWdZ$8DW!*gmtD6K@0W)_(s!k zj3!%Pi^eJAxvppV9G^3rX49B3n*5mYLQ)!I#uYZ2!1N?dd+Q{ELrrLpxtP7If44usXl@X30$IP=7r;lfl9)~imfF1?^lJaFc` zbN2r8!+yi5xQ*c|<|@IkJ0<=4!NyEvJC0xTL&xbNj%@&o_3<qg0S~s96i7|& zhj0_^BwCOKXfl%NB2(QVldQ@5jd}L6rZw3%$e81!vqSq^d+yKH$0xi(m@T_OwpP!& zVW|A^snE;r4RTyr*e~k6zfF8Vsk#E0*i+=nTK*XH(9&+lDS07n;y4>X2xx;8OS2T4 z-Mk`3XH9`PZ$E^0BQr&!OXsm5N-8fXxE1e0z9+`~TFI%pmr8HHMbG;qQF1EqoXT|Y zV3epi#bW4Zp_<4p*BegY&w5L5(l++AU(>^=J8-KFaIM_X^Iy6TbG=2#?m0o(FNRJv zh%&MF!6;SLt@yKU=+|q}ZnsjYPdkN*-)ImI*l+EE-e9I$^4F?JM3zuj0sY|MSbs%TZc+#SUqS zjw%vGEL*yxg}RsgukIvK{c1^z>Sy8Sky&g6OWSBoYTOMA2;N*_?esgqiB%iE#je4ORROq-2ZSvluB%B;1xqv&Z7!6(!epcF z*FvZz^b{oa8(R$oSPy82md^z4MbD`>Lca>xl@q`LynHH{sTFFl!s)Q3CA_(*OWa#5X)jvy8UIz-U4DCe?i$SHlhO|g^u6~V`^K)1zxTBfglM8%Qz2Ic5fb_& zT0o}|8WU^+?$<;Jd4)@^1BW?!g+)A2iXv^!jFJ_9T-5_N5{0_*B2lTklEP=h4a@2% z)~s`o$fgwzg8#Gh9rDQxQ}>{|Dh^N-)Xdxoq=`x(MAi@4f`JCM7TF@3=i80Dj5#AT z7X)G+s0)dZO<0&SsTx3}Ge`a9j~Pbk#mrYuzdL(*B$td#*l)yZjmU%%D+)JJVp<{8 zN0u9eGYt=MzZM=odc}-PFskmxifhZFPaR91$BKZIw~GyObqdH;(op2m()0sqsTL8r ziPZK=n^!)nZpSPln>Kx`J{2~m!hDxzj&!C@V0yhq``TQ1+SHI%u}EvSpQ#6dU!Rv< zm~C>^6B-LTPjnk4@KXX`i{=<~j%>(Bm*}Xka7Aa8Pgan$;RxW7q>0$1x)J?4p41R% z6KEe$@E5F~3HJYTQR;BA!i73b-#7!%`mTByc)riCU36B$3>WRP8HrOvDV@3BiR z5%F?E4PqNbI*v>jltK)nJ`I_QM~WR5!$H$EL_a zL-Ap~-#vDg)tL&r=>zMh<+wUCS@)w0?xS-aHEFtuLRZbd10ouuk9(4t-^=945w zvc8tc>T!B`jKI2M?4|U+BH$;mq$S`$TfnD5E1<1;$V;^04FG{j;V36~XwFkFsP=tD z=un3Gx#4+*a*pSdCQ7N|zYVb;Ycc{F(z)+_1%nUC^OWzq=w5oai}PxFVUZz35C|Mbi~0 zN@q5}^*OJJJmfRvJ6hdmfl%Int$_C???;xI|tC05VXB4iivSgW$VQ(iUKdIOR<24s|+NS#iLW{5=B?_5CM`3=V`UHlql+B zwu-upi_C^ABX(|`>`{tpm-riB!21DhVw-6;z=S!l496;iutpI{5GZA7MbL+1$C}4= z0q$Z#A+cfvg_U^<$_R82G(oB%rDUKSC@2BtSxu;>X>B!%2q!sn_QtjnDqbn~47a(G4z2Fx1bSmo0nANV< zXWBUV4DK7M0@@mp>D)UaQ;LF6a5jlMkv?ObqZX@08)%6$s!u3pl|_l5>^3~)Am)-qZ@7r*3&=O{vreoQh`I>bNfxBygg7B45=7Mjz?eu<06?SImC&(E(Rt?*aUn&qAaB+5eL~%Yfx4My zTTy^;oh{=!C>fU=xOY=P*3oSXtUt?k z031FPB~f_#)u}2Bo{T5z>Rg7_pA4m<@1g0UauD^CPEKpMr?A{90QxX!6>%t%!L#`z$?IM(}A$;!&bepGFnOVk=%zqPE=>ay{igB>o>_v9{Yjrd|YoeHLEF zEFHKy49L}Ps>w-h6A1(gnz05n>^8O=8Nv`a2zF~o^(MVRl&e>x?gu@tNf(lf5OI>J zsm)dj;(LD%;xi&<<4|AvU=A#)F$ZV%-tDk-k7+b?1k{I1ch)>21% zB+XW>i*9JQExoQ!5td`!qLyvx5vwg)M{liKLz;Ua-uJ)l6`#-5@2e)*gd3bgr?Lb( z8dGyL>?jOsxLZi)-!Sg4fEM z>_Tq0jvDz5tpI$m9X%3Ak&T&_`3D6bILCXDt?`3?U0(&&Gqe{)<5&*5n(7;!zDmU< z<5{9nMXt+LMJ1gsTH>mba+%IOI?$^+4JoSX0Hpm`ds19Qb~s$3 z{sVE8e_!0m-xPQ8?~1$m8{!`R9dV4mF7D&6i4pNF`=l*wDi7X6`r1%ot>Tik+PEQ*iMMBU4{+VN(>Z##xZf?RH_`s5Qm?IFg4zul}6f~ zUH6P^ON@$8vav#ysDc}ucB{DP$b~A-963>NYO1(Eiqe7NU`{zCe6MFlyOLH(;LNDz zZTIWfKl9$}_orU9>S%cW^6s(rZ$71If1}R$W1@2vZ}bZQu5s4Un&~g=>P?+dTko*0 z(KNbd(^P%dDKrbOXx!lD1C5)$zHDt8%_1-bZUJLu7#o-(w}G)UOo^Ac^FV7l7%B4# zMk?7z8JH@s0aME`70%w!8dJMCm$uN*W&M`$!l>O7U0=K$wtCzb^2A}>4R6gCEt)j` z&}Y)~_&k8GH5u2MI=3EaZN_zOJSa8sF}#4InU6G6YZkC|QO?}#ML`_&La+03_qNY@ z5H9X&x~8=ZwOR@+{Lo9d=(hl|wy6!+P-9xG4~#@l*iavsLv6$0EHT^yuol-7J;e@+ z+!(SAGs7}&4h^D??LmnbhAd&HwX!y=4IHdMbIXaD6lOJy7#!kxsN)MDhKEL5zf3z; z4viCCHl7&RPse^nq7ACtN{pcrGJWsDb-{4_T?O!!+9F;}+xq=gsrUUxO%|fnKDchX zZV(2s>&imh6Rjm#SocLQl7)7s=f$!j{5z|G@ZCk>acOxF+oCT^t6_9!)%VwZX?DGM zNgAErqBQ*beQ88ND2*0Ym=gYC5XHW5Z?Cr7zGzs|jQmbp7P;Gn#5=Os_3pbJKNKaL zOxpcSK+@*HT^AF(J4}{bx7G2Y$aQzL_gugn9=6uk1J6?ZYZ57RTypVT2vEnZJp)Wjq>&5<>C}ZOy-lz`XuxZ0$$My4U zPM>3rKF!vjnxHfnbB?u$t@%T3Qk-jeqo)88ZJF*HiBwD2fN?ft5B0WwQiFK(1wq?3 z^z_rsC|LK=c@j9u?u%#en}RvQ6me1Rd4iunC!&p>V(X_Toh+lnN^}$jb}SNc=yiRF z*l~&Ts~u`rT=&kZ*GWedBg9$SteP>r<##%+J9d<(X~{VP1pwM0B4BGZHa)FyEp+!i zlW%HK-BdC-bW@d?8M~>rw668SyU+9M()^}>mwj|jda)*%x2d6wpqcigTM8ul+}|}nrtAw z*YG|*0{|-nRffF*8<&}=kTv|GM!f-9K$!R)wyTyRWlRIaJYk@V?@52QItR!Zn8TfRWJ3Bn?^I7(Q0Dy&QV zqlG2DM(bZ8@B)FW1V${Iqb@Q1tJHWAAZ0d|&X_GoBJn)dY%s~X2bd_!=?vvT)pyoFSz$Fc%PP2Rv&?*7ECAKQyYIM`ONu5VM@ z9BUV13(}x5L{i&`FI*d#QetsS(K>M%w;#q87=(q}g^gldNs5WFNhY?cq)}Q>9;m^@ zNo>4J->k#~R+ckM<_7#eakhVpkhjThy6r$jaicb7W!Wv|UZj&p2=3jnwwKDbv=`M$ z8&1ya1o4_Iz9f2mvZqC!HUu$3=m2tN0uvi{YOCHPDmt|+h8+q*5fFt-K0=@=9(3Kd z@ALesd=FI3S|$WU97H_nZodY)Xd1xLEoeuD)uAVKXvzB7qlG%U>$%2Y0j-Lcyr5FK zzoMZGvkE-Xg5QIQE|5{?EH;)+go$)AeTzLvSUPI_gs=Z6%++osyR2miWE%d~<34E781wfj;(2v9?sZUas zMi9nuva8%fnCiuX*pzgtajmsN8K89&>!`iGb5c$Zz9r)F{C#_ zI=NQ5cblWO*DR|Q#>b7{l-<+NM6xRXPR-5<|l%D0(wvII(pNbsa{?6(v@fM)kXAad_QJ( zY{o37-uU<=4DmAGW1O?;9+NP{$$Z8*7t}sd>$8|x5u}O!(-rbDRxFVSYmBbK2?S;BW>U(R3=u< zdQYle4OV|Do>l>&F?fPk$YZ&66-t~;^X1(+*aG1JGz{@2YJZu)R{(A`%mWw|YVTgf z81iA7qnya@+i2~goU?}|^OfeyKx!&VKfD$tTgu~zZ_s?2gSdwBuwC-$XcGlC+DpxL z$$IWm($knpW0?@`G|@hoG$-XU3+^*M`)63qy{g{H+?|0TzXevvp$VSfb6bB$+*!ePn=t<&1yEVg3!Ka9RgaHQq6{ z41}1KvCJ_w7L>B?OK2%+p!j75?_;W{EF$iH%i3UDD2klaey-oq-x!OOc66yv#Xl)jcUA&AMlGll@ppj2K@_gcL ztTz_5{TN+3$1NA8CK{bXJjB4Ci1sG{|HE=*nhq@Y=X|-{5AZ0X`QI)))fYYwTCs~v zJVt5lzz+Y+cW6wffjtHKVi7u5=fK2TRLcDCzR|(m0mQB1-BytqNQwGEYXB& z8m#A0eYJzq^j=2Mtde;ib&qAlsUgqH*3Fp=O~g4;u?4=0jIP)_pZF-9B;tOHg&fG;wNLXhQKHg=XcD=T)mOk|q&RL{8AJw3D4 zH8#8SNP#i(2INQxUf^l>3Hb$H`3wAo;t3&y{KhkaEZ?b~z1d!+M}4NMPF+r&%Xd!q z>-CC-=WqW!)HyV3S^uEU?#Dsr3f@eC2urZUY8C&i$6K5gZM$WI#}lXLw%o~>*D48H zIFGHCFAj>Ts6A$_@(oMWMdPt08WCTrjO|uc%!&hI3cMP4(_#j^nF+5h*lnwM=s8Zd z7Mo1_-w$OR-bS zZ(C-DJDG82dxlrcK;OfPk|=F3y3=j*HvL8IyxUJ7KHtRC#d+rG{%rP7H~zl*yUV(I zEmZCMi+$OLjOq=}z1#0~Rz6-h*B)l7-wV>d zh_bE0GYhwZn^$k0iw8jx4$k$!OR{r=_@@`}zCRy`IBTCBtms-HM;63stbA;rH9&mU zV6%9q*bH;NkOwjJ+$M4bMEh}?sW5Fv*PoL)?zBswTTdAc+Q%5^TWBaNw*`C3MkU2^ z7$N6*DR*)=_r|4>pZn`v*`)H^hxs{KzD(K()XE@gI=az~6t0GWBwUH4KJwu(Q4l^% zqIvwuxEH6P>Pzj0gF&)lWNns7#=#H^+uBY=r&URdB_ESoitMOM*jtw!OgTUA{a9Y=E*J@-MZqqyn7yZ z(#tNO5m_xJtQMSB@CeQdSG7Ct6RX37BizU3mMc6_LO^*>;PI^zwBzf8i#JINdu}cd zVK<3zwyy871R0X7L59mdg-a@nEHLKle8Ymiz1#}y+y-?7S;tl0dMPgne$*27Yt{xE zvyrd-rE+02VduVZp1R6kuL#Jt!MnW6M&+E3Dyk|k%YF-)DR4uOvd)5;E z82r3!_Q4qTpX>u4sEpaJ-Rv%V+ZvUZFMP@RQX$8F>nCSbAQ6JOnA>H=-nlD(P5kxx zhP6HB$v>*9YUJfHZj2jDkE-fGZi%v*%By2`#DeV8c{#5hvp%=RHnWz0OYqROM)f`G zeB8WV%)>qE%YQcO>YoFmf*D%TtSmk!?2;alBOpy%S06+&jS{%xYNWkP$yh|XGzf)= z({A?%H$T>{4AX9;Z3L4p-JV2+okZwmkqVWP+Juc5tUh+Rw}e?Ff=|G{gRBCm6`+Ve7ya zC%{2b*n7q#;mS_9$usEEp4kWq@t<#4AGS?3NbW1Z-Omg)N!PKw1k8$~p3M4yPJ%3aKRo@e#{7&+N!n z2r0ppgMgq>K5}xt?!sNUu#aIDf!%T)Jx5r-A3vi^egGa;0&WZ{WaY z$>>qOg|<o5%S3)Uhrft>lCmF@DG_kU+3EfzjB+o@`hafz-*jEd^K@#;4 zR*`s5%D!xJc@$?BQAD|AR}dUS_cq>1o`z$){(nL3+NLDlOG$^6NKZyy@rtAaHFc0; zUQRp07iFZMl1V?5oGbD;w%DC+$o^YT6%DE`O}hCD8vCak6aZF$bb}Ir0QViFnUYB| zwkZrsV-IPiB(GrxN{2~g9YdP&O=OjGpCU`xfaWqK5&R>?2xqG|<_4zUT%I*6T#TC) zd3mX7IJ2s@R7ZMc0x6=)>84K|$kdp?w249BD^(fVD@q3wKiV~RH`!KFWL|ID@-+Ts z6GTt9;R6at+)jof3V?Xc1ColMLUROQa*l`rNrJ{AqzS_1J47g)CAo7EuH-kmeDz`X z7W@g`QW5<}X%>>dG^-}g_eb72+Qq;xi{;m7Beo<>lo^*Yp+*w@1aC(E)!B^+@skwIGl8nCf?19VgVbs9oTu#-3tq_|z@~km@9nfg6+c3Rz%v6cn z@K}T4`8GVFW*h@-;ez(!E>LuKf#P>?j(iWKS-@YRzUvs!9uN>Uqx>OF*&9vzK!z-f zN-~&W+A`*3H4xIn6gyg=;G6s1|Ao~rWerMTq?+tkAZvWd`i!A$fC5MnYlRZt{){+W zo_u0GVNZAm9xwP~AAABa_&BKwrI?e}m9vJdq1-hl+@sc-brj><8OB@VDBqya4Ne(m z%_?T$Yv>v{KE`qCay92rqqq1HmQpBFoo{N>E^K!9(1;0QS2YH|2) z_%dy2%a5?y#DpY#0)`nkCWpqC3_p8b{s_%yc43D(vr@J)GgcD`BRlFIMLms{Z95|4NdZWht&ExhAxfA1`I&Agz zFS!&XCX2`gA{0a~CQ~CUohLsavg6k#!i_Y?!3gjghxYMN2T3;l>8Y|oBh(c1&9DRK zMV>R7L-sR0ly@$^t=jjsZ~o4vX(7iVR<;_&SU6Put_+1FdznwE!HPUjTa!PTFCp!H sXetDg#!Z?k0C7j^9WU@N7S+ThhSeE~=vz=9sgU`0r94xwymh$tfAXWt761SM literal 0 HcmV?d00001 diff --git a/ip_lap/models/__pycache__/video_renderer.cpython-310.pyc b/ip_lap/models/__pycache__/video_renderer.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..aff043a7d18d8c7273d82e794e179c69fc5bc817 GIT binary patch literal 17006 zcmcJ0dyr(;dDne)_w749Jw3g>v!mU$)RCptc92(|(OQmWMQdqSE6Z7rg?5t^ZHZQI z-#ars+mF$?clOcg0VN~FYhWD=vGahc%or-Ln zNGVjM0%-aDzH@IsW@l~w;db5A=Y7vT=R4o``@Zwt{_L!0;QF2ae024{HVxyOOiW%5 z5}&{wo-_@^H+-{Wbj`YDGH!M3x-I`k*Qq;6UZ$RT+3;=OdB^abz}d*|*!3K8GJY00 z*}&e=c`ovDzKc9}l$S?d-uIB_jq*I?&G-f66_Va&{8_*Fj!`f8bAAc`v;Mqa#(&X2 z;#crLhZaZu1+-Wg^-@CKF@F(xi=(`G(>HzV9T)Q{Ggo&$XE;VZgMKry_k!PiW#v-yR-jsGdv=q-ox{z%a2-Kp>=?V| zo?#l1wP6pX&pX!8XrR|^^tAdl>vlF9;d)aAew=4~t=a81W0%pZLDa;I zIAPFPjoo%{b1Q1}TJ_m(-{0yqIq68ox!!u)_k&*SZiPXEqq=Lx`9`DFX@+5=ao6~x z?|W+f)`z0*=2PeU-PK#)|H@OXtuX3$8@;|CglXo7UU{`~>4jIHYHv0=&CRFqjGz-f zb-nEe{e}v9exL$%X7iTHZOVGlAPG;|AruhEd9FkWah-r8z*bcrmgdYsRhO$etIbUKa3#3&X}syGxyF3BYfvl1zor8&uK zAvtxiBCqfif^_@)Ibi_Y_<1k@VU!aLz>PEV7W>@i>Lh{|ACq2;AK(I!IW%@aGQtVS zZ!_4_CU^WWsvyg)8+*Iv16$1~Rg)eAZe zE_$2|BZch|=Qf+Z-|nrg{q#@q`mJ-dS;hHLe~bb5N}P%Msdl&DL*h-UR;~F~bRB6*?t7 zKaD#qA}CngDh1gvjy0IoFO=*O^!Fs{2)XHi6g$!1nD{wAHQWjD7p;;j>-<3Ce{YKT zT_JwRA|d{)pM#7YOTRn>#j7t|cuAz+J|vql9#@NGpI(fAhiw23iFqFpe@+X*A-G|( zDO}kUIcyY&z={0)1M)o)hSrKYgNN5N}}xaJQPquAc9&>mQg*+jE}R%j230Gzy2SH^+mLi)h&Whi_kU-;EJe+3Jxg*Vn zq%uvjB7sP;sl2g9@&@b#rH~W&Wh%xokrS5<^}T3Za}+yLB>wx~?h)pmLl8ThirAHX zPZmF#hzUv?QSSGlMacOC-CFoT(hkZK@0BumM{yR$bn+eOmkY-GJ#tbPyS$WFOh1r* z)hAF|%f}gc_JN1pN0j7yHtukUATl1nn+at4EYjq532{e<0K3{Qm$u8WEkRnD;=Un??*-VVah23)3X? zoMf6gVJ;V^NHGPe#APAJMrC2c9Ol}_(EY3C7fdP*yHp1tb5hQ@mq8lJ!k#G#LKdmq z9ymk89g^5Y8PIdQ_7c)-^>T6x#rMv8#fx?E`$( z@>-19t#&VVT4V_a2<0_4a++!_jI-AO|JK68WcD@Gk~J!WW=f_DAnrMyIhZ>yNMlSk z7AmF}$%glN+#xy9hy?=wY?!KmX)!-)2mrS@DJiqF)vKtjI5hRg5hTJ`J4obft(fxkzSlv-tzv}5y1i)uRGi3wzX|L+Ek)X-&6qgl3w`g>LddQ@ZsXj2A4>)ORGgc6@cnKROta)Ymu-~$y^S3CV1 z)zx;;@xyA=ulfO&xho7YJmXbyy#QK%R$@#^_0v(%3;U|}Z1v)(b+vyr0`d*~>eXA- zXg#QMhC9;kwIJG3y@^gH9#S3mbp?$kOQ0nCRg9*J?X=GEA2#pZf!-l*r9 zBRD8_r1s3jfDSJGiDU{J(kSXTChGn$Mj@D>nENDlQ3cywwUIYP4IVPcXN4Mm5s_(X z5Ly71fmUr$mzK0WN&_7N3p&2VwXc9K7P0ckqlP@02Pyy!Aa4&m(S82nE#)|pQfyAQ#E^Q zL2N-!`~Y&}tO_=RW~3TScw4;?=A&Q`G~#S+ycJuw zaNoRH&vt`wJ@5|@ntBt}ehjxp9l4TKGAkAs4DOSNgS1#Vb%8BBJdv z*PH|QD24MVMwgVSh`tf`)E8UR#}Am26>OR|Cy0 z0$(~I$C~JPg4s;uhE?G-7{_>|9Ao(L>h)%4E2y@6Rax}v=?^VESzW40*Y`0sfC8BP zCLY1J64nJarbno$OQh0Ri>0KB-K*`67#r1}KoYQ3D6Wb40l}nRXZ0BpbVG}sTCQH0 zoF28#8fK%OS#R_^em!%8e-l(UE%f6Y#|ZBFvuLJOG#9!Cev1;SE##AUr!K-+Ua8Je zx&Vm~T*T&H38EW)b#0&2w@*|jibyMMz($PE;I?sxB}~WR>PJd~Fb9uKL*P!3WjQe) zi(b5&7v$$@!x2hd-+{3n_9|&LqjhBDrmYIhNm|WDMc*Z=+yRem%uS_AQg7bRi*8;< znK!C=Bzc;LC&@NRo>qyR-J_$jWDBt@eP&WdPC9BQWeYHN!*1F*Cc68~%6=7-v(Ig; z8IjNn_4zAzt++r9W))m#3z$X;-g%KVNlN>lKu}xMxR(&)P<8483qFHDE4fRJt38c- zJqooYF=%Nkz_PBfj2f{NE^BKSwF0fUTw*tcSKBv(&gJ#3)zwZA&uuaVYq=z82b)yO z&L)<$v#P+Rxy$XfZXcpq+u1$>CB4kfKh0o;!AlGl7>uzhA^lfz3nGRpW|v_zTC^Az z9kDVXZ&73~Q0)I1Jamvk$6$MsHjkl8A~d1R_d%OrayVXs{B%=;7cAzrD=ofdla@6UQ%Z@+^)5l$bm#&u4ue`p{Hw)@^7>hxBGiZx(vcC76=SX%wFume2!pOUTV6HRdIE9=TZ5j^wUJ_7 z=2DA9d=)xIZB|3gDcqK7F-Jfr?@_(LxrXtEw{70W+REyhT-`5PZ$Su(9x8N;jcc9x zaBQ}5CRMu9@KyUd*lr`d7C;XjQ+Dh<*J+1QtKZ#x{j(l=_RPsT)h*Nw*%mJvcSe^EX-lO9-HrZ8Oeb?f2Ke zdHYxX@IO6YJ367x2`0;Ez?L6DGu30OKh5AzF}T3sM;QDtgSQz_a~pFv^$wHoGWbyj zw;2%5i*T$R)5d{04tb@qs(6Nj8Hvl28OgicL92R33UDwyMFREW4b|M-4E)CWQX|~z zPNbJ5J(rft+@Qjdu)i90n>V3(^!i;YIeitISCTdHZ`S8UH4(!RhQwpZ*jiu0y^Pc8 zk>c8T=4^dtf9;u{ zwN6T{lLF|DX)&U1Vg2<9fM7X5okQ%1!qLFp#ySGQw2*F^OQhE}AQ|+zx8Nn(pDx0t z1=Ce(>ki`rcJC;4@L>Wke~5jvKlD9I7w7b$>XO=F zSup{ZDYi`NE~nKH(XG#p&kEu^aV<^>a1#7EbaL0!7*n>sa01!?d>0S9Go)r-1V<`? z%ap;7rbBAr+21=}A=o^wA*T)4iGwZP1^YPK4Nlz862_2|3z}Ypn_LOu4gVUE2v*&( z?@eR@XwvMGff*?(x8KBh|t`VeV6L7{{?>3A`yEA?UA@?cNp*(OYdn6DZMw`UwUP zGx+uZ;jf_z=u8`XRFe(WRSk|mc-|gt z*ymQyvwKpmYd}&zshHQoJ#~?ad8LxkeJycGnp%GJBnJUP^Uol0&;+mp`fflb%0J|W z59cA@g2NE0pUS!`d^NKYXRfcUEqx3!7^gLzt;Q0xO&wbv#hxC;&W>UqX)I~3XY+7C zY!-1=+%`gZ{XoAK<`=^F$Ni~(nSro6;(>T(0$9mB#4{5s85duLPh7AW;ppg6KSYNx zIwxuMH`&u6eCIQqU2;X>8QZlSD{ES2;r+pUI)R6Ck2$m$agX6G#{hDzOw&E z2x}HZHT(&~=oW-^hvar#WA!8pClh1xyC^CkjMBfxu6S4?mg_O!gVm?g?*-}~vFJAt z;H)fY_S8=zTAR~S^PjNnA2awV20x9U?rD)P9DB^Y_E?l#iwGiudI6DcGrT4#^-Cy= zi)#TMhSM2eS2;2sY6C^`KK05(Casl{EAsvx z9{f(6dyZ-0us)5IEyv8>8u1 zd4YCwDA&II6&ikskYTie*~WogUYBIXW?c%FXjPKY8;9aiYglxCujj_#{qlpp3GA zSjevi)z69%{EejF`^?~@LBZ%vg_G8GX^ehH1EhNmM=s}#`XES~IR%H>Yu<1FmCf(< zzHx5tx32%%&m4d1Kb%|p^7%j8`PTpX_vZu@*X*?q)cuct<1=4Bw{~m))RO!2Upu$< z^{>fw4*S9D1veV{&E${<4r-5J4eF)zxu7-~occ{<$2O)Z&T2JR@468Oh)+Tz@7&S{ z4-gY_HaHsjb&`RHE(+dp>o{aAj$*!!WDEDyWhW^i%8`GiS+__25^miA8D!uukr5Yb zLZ)n7ZoE1Y99UQy)|)Tmp65#v!eE@`yRQ5mey=pvh5p z)FMYtFK$_`6Yr5H%5LJeJ~U~i-?75EPhlhPm_LF|3=l$xb*Qxvw{TO3P=nBdoi~G% ze4Mj3@^5M7&t4gPcyh&hL7)%Yv70dT=+(PRhXptTZ~E8{qA;?(`j-r(K}H%J22H+5I0*(8XL}aLp>sXK$+)^x8Ny>U1AMRzND$YIzi;=J~#vK)Hoji(20~$ zqj!#N$b5P6#4A_(czg}RQhyFXoRNu*z481W;0@x2{g*g3CDeByU}%e?ZOVF^F#C{c zNpJ(s_nzQ$D56x256d+ge!m4bZh@GOGyl682qTg}5Ti)Ae~yCB@v)>2*MQk9U-A@l zr$i9xm~c{N`#D6YR&H1jHP#(-13onbV}N>41U2XqkcCfJ`1#6Tp2f;&k1Jz`cJTQ@ z>19Y3!PNKKB~Z1P%xI5M{XcvpA)gq6UF^m1Y+phpJ$RZyU&^1=N#@+@7w38 zJ!oP^9KEk&WA%G+R-9Er^(!cg-IR7@PGLRd?idr!#F}t9;%q}E`4icck;Qs>z71Q@FS!=B*i@7@a#Kgou;%o$PD=1Ilx(1 zr?UEWbkO`A?7Ok(5S9;a+U931BdL1`E=YO&X|mIl@6&~YY>K=Uq9YYN&sLw{i$C*3 zb*N(>d*XifjDD){;7NJ%M7Of7#z*gyvPugm$^%{+%YxNInu;a8C?S5&T{%{zu(~$ zy|P+;^UW)2D|qwGC#y*zr)m;IQG_iopS_8DJQ`h>4yfo&i?0_tYkiDyy$d0g_&!N3 z$k1`SR1l1y2x{ul92g9dOFN(xP2Ma;KZ3fZnf&~ePND*8*OXJ~; zRpshw_{^n3SeB#R=@2mnbfX}|ugdzh*KktNNxXu3V@C}xcC_9O@%SFFK^x96)$t(( zz%1rH@kx`7D*R$4wmM4}l>*TVS7?!erWfO>SPzTXpp##s4Xw!DvDEC)T7&x?#Eh{8 zQwr=1CK#W!W5SWn?#)Kmuwz#f7GFL6cv!1;@e!J=Bfmo;Et6KT?n%{j*Rk#yS((Af6%AIakAwWcG5Qg< zMytQp!xH-HTX-mzZ8^5PVK8|3{AgPqR(nu7z{j2)-i=yuz1%OP4^wuo|ahi*uj8 z{F3Cs;Y&U$nwi+Z@sV_sP3+1zv#Id#IE#G*kIKY{SPK#nDZj_3|2~8F82kZ)|H(jx zb%=@RHz{Mj2Q_HRqjxSG!SG*(!~rh5=iGB)?<-guow>kgCEQ4(4r;?gVsLCcnS7ZM zP7|xZ)c0G&G~7P|ABK1^;Fq}S(jl7wjr(E8wD9SW_>{g&o`G*%Xdd|an0aE_29$462Tme34dE$rGi2%)?Pt+rSRgep&UmU`2@37Tr0;=l98~cIl+A zZR*o0qAt-WIga~m0?0c}b81s1y`N@Q; z9B}pVd7lc!7`CEq`IG@@CsqHD^@$P_2KOJ&UBOu!i{irMsg`2Jc!A7GbY*SzO^>|F z>UUY}P=JYIFB>(61N`gA*BkaH7~_VG?Lm7}V?}Q4IP2pJQ8;$Uj^-}XK7HKpaUp#i z)Uj@@r2PlLsW#)RUI-Z7U?)@ml+FGbTQ~$ou{UAl%gB4gMVbAPH&{KjW{yDYzh(2R zW^eW6%fqCwUQA4%Ng6DY*0}=+SM;3!M%J@|uDltt7c!fyjqhFnb+hj~cJiNbz=&W{ zySHGuZUIb(M-DDDBLd>=`><=9p2Ls73OFgl*AMrsw~G#Yfe}u6;Xkw|2c{HjurMZh zVO9ynPdU1hMU4z?pG1Ntv_1ONYgbo%TgR`&AXwomT`_RvMSp0~f}mCA5VvvQCEwNI zi(pc{I1VCL(NIYpQ;Lrw|_;9G4Q$vG)(pH zI3UiJET`hLAFN{jI!eFF6{i%SifPSp1&i-{kCk2#4Wo7}c143a)9Wd`gQ5O6gJlM^ zh^cQPh_evXTb)2YjnyC7zQD3iF?gE68w}bE-ej=F;5vh^F!*H#G^mLtCQ@7;J~5Gt zEEPS>eGc*+4|^Wt*4E#GK^f?pWvB98mBq@@a>XsVIk)H@tIStktGwi%F3-D) + + frame_shift_ms=None, # Can replace hop_size parameter. (Recommended: 12.5) + + # Mel and Linear spectrograms normalization/scaling and clipping + signal_normalization=True, + # Whether to normalize mel spectrograms to some predefined range (following below parameters) + allow_clipping_in_normalization=True, # Only relevant if mel_normalization = True + symmetric_mels=True, + # Whether to scale the data to be symmetric around 0. (Also multiplies the output range by 2, + # faster and cleaner convergence) + max_abs_value=4., + # max absolute value of data. If symmetric, data will be [-max, max] else [0, max] (Must not + # be too big to avoid gradient explosion, + # not too small for fast convergence) + # Contribution by @begeekmyfriend + # Spectrogram Pre-Emphasis (Lfilter: Reduce spectrogram noise and helps model certitude + # levels. Also allows for better G&L phase reconstruction) + preemphasize=True, # whether to apply filter + preemphasis=0.97, # filter coefficient. + + # Limits + min_level_db=-100, + ref_level_db=20, + fmin=55, + # Set this to 55 if your speaker is male! if female, 95 should help taking off noise. (To + # test depending on dataset. Pitch info: male~[65, 260], female~[100, 525]) + fmax=7600, # To be increased/reduced depending on data. + + ###################### Our training parameters ################################# + img_size=288, + fps=25, + + batch_size=8, + initial_learning_rate=1e-4, + nepochs=200000000000000000, + ### ctrl + c, stop whenever eval loss is consistently greater than train loss for ~10 epochs + num_workers=4, + checkpoint_interval=6000, + eval_interval=6000, + save_optimizer_state=True, + + syncnet_wt=0.0, # is initially zero, will be set automatically to 0.03 later. Leads to faster convergence. + syncnet_batch_size=128, + syncnet_lr=1e-4, + syncnet_eval_interval=4500, + syncnet_checkpoint_interval=4500, + + disc_wt=0.07, + disc_initial_learning_rate=1e-4, +) + + +def load_wav(path, sr): + return librosa.core.load(path, sr=sr)[0] + + +def save_wav(wav, path, sr): + wav *= 32767 / max(0.01, np.max(np.abs(wav))) + # proposed by @dsmiller + wavfile.write(path, sr, wav.astype(np.int16)) + + +def save_wavenet_wav(wav, path, sr): + librosa.output.write_wav(path, wav, sr=sr) + + +def preemphasis(wav, k, preemphasize=True): + if preemphasize: + return signal.lfilter([1, -k], [1], wav) + return wav + + +def inv_preemphasis(wav, k, inv_preemphasize=True): + if inv_preemphasize: + return signal.lfilter([1], [1, -k], wav) + return wav + + +def get_hop_size(): + hop_size = hp.hop_size + if hop_size is None: + assert hp.frame_shift_ms is not None + hop_size = int(hp.frame_shift_ms / 1000 * hp.sample_rate) + return hop_size + + +def linearspectrogram(wav): + D = _stft(preemphasis(wav, hp.preemphasis, hp.preemphasize)) + S = _amp_to_db(np.abs(D)) - hp.ref_level_db + + if hp.signal_normalization: + return _normalize(S) + return S + + +def melspectrogram(wav): + D = _stft(preemphasis(wav, hp.preemphasis, hp.preemphasize)) + S = _amp_to_db(_linear_to_mel(np.abs(D))) - hp.ref_level_db + + if hp.signal_normalization: + return _normalize(S) + return S + + +def _lws_processor(): + return lws.lws(hp.n_fft, get_hop_size(), fftsize=hp.win_size, mode="speech") + + +def _stft(y): + if hp.use_lws: + return _lws_processor(hp).stft(y).T + else: + return librosa.stft(y=y, n_fft=hp.n_fft, hop_length=get_hop_size(), win_length=hp.win_size) + + +########################################################## +# Those are only correct when using lws!!! (This was messing with Wavenet quality for a long time!) +def num_frames(length, fsize, fshift): + """Compute number of time frames of spectrogram + """ + pad = (fsize - fshift) + if length % fshift == 0: + M = (length + pad * 2 - fsize) // fshift + 1 + else: + M = (length + pad * 2 - fsize) // fshift + 2 + return M + + +def pad_lr(x, fsize, fshift): + """Compute left and right padding + """ + M = num_frames(len(x), fsize, fshift) + pad = (fsize - fshift) + T = len(x) + 2 * pad + r = (M - 1) * fshift + fsize - T + return pad, pad + r + + +########################################################## +# Librosa correct padding +def librosa_pad_lr(x, fsize, fshift): + return 0, (x.shape[0] // fshift + 1) * fshift - x.shape[0] + + +# Conversions +_mel_basis = None + + +def _linear_to_mel(spectogram): + global _mel_basis + if _mel_basis is None: + _mel_basis = _build_mel_basis() + return np.dot(_mel_basis, spectogram) + + +def _build_mel_basis(): + assert hp.fmax <= hp.sample_rate // 2 + return librosa.filters.mel(sr=hp.sample_rate, n_fft=hp.n_fft, n_mels=hp.num_mels, + fmin=hp.fmin, fmax=hp.fmax) + + +def _amp_to_db(x): + min_level = np.exp(hp.min_level_db / 20 * np.log(10)) + return 20 * np.log10(np.maximum(min_level, x)) + + +def _db_to_amp(x): + return np.power(10.0, (x) * 0.05) + + +def _normalize(S): + if hp.allow_clipping_in_normalization: + if hp.symmetric_mels: + return np.clip((2 * hp.max_abs_value) * ((S - hp.min_level_db) / (-hp.min_level_db)) - hp.max_abs_value, + -hp.max_abs_value, hp.max_abs_value) + else: + return np.clip(hp.max_abs_value * ((S - hp.min_level_db) / (-hp.min_level_db)), 0, hp.max_abs_value) + + assert S.max() <= 0 and S.min() - hp.min_level_db >= 0 + if hp.symmetric_mels: + return (2 * hp.max_abs_value) * ((S - hp.min_level_db) / (-hp.min_level_db)) - hp.max_abs_value + else: + return hp.max_abs_value * ((S - hp.min_level_db) / (-hp.min_level_db)) + + +def _denormalize(D): + if hp.allow_clipping_in_normalization: + if hp.symmetric_mels: + return (((np.clip(D, -hp.max_abs_value, + hp.max_abs_value) + hp.max_abs_value) * -hp.min_level_db / (2 * hp.max_abs_value)) + + hp.min_level_db) + else: + return ((np.clip(D, 0, hp.max_abs_value) * -hp.min_level_db / hp.max_abs_value) + hp.min_level_db) + + if hp.symmetric_mels: + return (((D + hp.max_abs_value) * -hp.min_level_db / (2 * hp.max_abs_value)) + hp.min_level_db) + else: + return ((D * -hp.min_level_db / hp.max_abs_value) + hp.min_level_db) diff --git a/ip_lap/models/landmark_generator.py b/ip_lap/models/landmark_generator.py new file mode 100644 index 0000000..7d41bae --- /dev/null +++ b/ip_lap/models/landmark_generator.py @@ -0,0 +1,239 @@ +import torch +import torch.nn as nn +from torch.nn import TransformerEncoder, TransformerEncoderLayer +import math + +class PositionalEmbedding(nn.Module): + def __init__(self, d_model=512, max_len=512): + super().__init__() + + # Compute the positional encodings once in log space. + pe = torch.zeros(max_len, d_model).float() + pe.require_grad = False + + position = torch.arange(0, max_len).float().unsqueeze(1) + div_term = (torch.arange(0, d_model, 2).float() * -(math.log(10000.0) / d_model)).exp() + + pe[:, 0::2] = torch.sin(position * div_term) + pe[:, 1::2] = torch.cos(position * div_term) + + pe = pe.unsqueeze(0) + self.register_buffer('pe', pe) + + def forward(self, x): + return self.pe[:, :x.size(1)] + +class Conv1d(nn.Module): + def __init__(self, cin, cout, kernel_size, stride, padding, residual=False,act='ReLU', *args, **kwargs): + super().__init__(*args, **kwargs) + self.conv_block = nn.Sequential( + nn.Conv1d(cin, cout, kernel_size, stride, padding), + nn.BatchNorm1d(cout) + ) + if act=='ReLU': + self.act = nn.ReLU() + elif act=='Tanh': + self.act =nn.Tanh() + self.residual = residual + + def forward(self, x): + out = self.conv_block(x) + if self.residual: + out += x + return self.act(out) + +class Conv2d(nn.Module): + def __init__(self, cin, cout, kernel_size, stride, padding, residual=False, act='ReLU',*args, **kwargs): + super().__init__(*args, **kwargs) + self.conv_block = nn.Sequential( + nn.Conv2d(cin, cout, kernel_size, stride, padding), + nn.BatchNorm2d(cout) + ) + if act == 'ReLU': + self.act = nn.ReLU() + elif act == 'Tanh': + self.act = nn.Tanh() + + self.residual = residual + + def forward(self, x): + out = self.conv_block(x) + if self.residual: + out += x + return self.act(out) + + + +def weight_init(m): + if isinstance(m, nn.Linear): + nn.init.xavier_normal_(m.weight) + nn.init.constant_(m.bias, 0) + elif isinstance(m, nn.BatchNorm1d): + nn.init.constant_(m.weight, 1) + nn.init.constant_(m.bias, 0) + + +class Fusion_transformer_encoder(nn.Module): + def __init__(self,T, d_model, nlayers, nhead, dim_feedforward, # 1024 128 + dropout=0.1): + super().__init__() + self.T=T + self.position_v = PositionalEmbedding(d_model=512) #for visual landmarks + self.position_a = PositionalEmbedding(d_model=512) #for audio embedding + self.modality = nn.Embedding(4, 512, padding_idx=0) # 1 for pose, 2 for audio, 3 for reference landmarks + self.dropout = nn.Dropout(p=dropout) + encoder_layers = TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout, batch_first=True) + self.transformer_encoder = TransformerEncoder(encoder_layers, nlayers) + + def forward(self,ref_embedding,mel_embedding,pose_embedding):#(B,Nl,512) (B,T,512) (B,T,512) + + # (1). positional(temporal) encoding + position_v_encoding = self.position_v(pose_embedding) # (1, T, 512) + position_a_encoding = self.position_a(mel_embedding) + + #(2) modality encoding + modality_v = self.modality(1 * torch.ones((pose_embedding.size(0), self.T), dtype=torch.int).cuda()) + modality_a = self.modality(2 * torch.ones((mel_embedding.size(0), self.T), dtype=torch.int).cuda()) + + pose_tokens = pose_embedding + position_v_encoding + modality_v #(B , T, 512 ) + audio_tokens = mel_embedding + position_a_encoding + modality_a #(B , T, 512 ) + ref_tokens = ref_embedding + self.modality( + 3 * torch.ones((ref_embedding.size(0), ref_embedding.size(1)), dtype=torch.int).cuda()) + + #(3) concat tokens + input_tokens = torch.cat((ref_tokens, audio_tokens, pose_tokens), dim=1) # (B, 1+T+T, 512 ) + input_tokens = self.dropout(input_tokens) + + #(4) input to transformer + output = self.transformer_encoder(input_tokens) + return output + + +class Landmark_generator(nn.Module): + def __init__(self,T,d_model,nlayers,nhead,dim_feedforward,dropout=0.1): + super(Landmark_generator, self).__init__() + self.mel_encoder=nn.Sequential( # (B*T,1,hv,wv) + Conv2d(1, 32, kernel_size=3, stride=1, padding=1), + Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True), + Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1), + Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True), + Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(64, 128, kernel_size=3, stride=3, padding=1), + Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True), + Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1), + Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(256, 512, kernel_size=3, stride=1, padding=0), + Conv2d(512, 512, kernel_size=1, stride=1, padding=0,act='Tanh'), + ) + + self.ref_encoder=nn.Sequential( # (B*Nl,2,131) + Conv1d(2, 4, 3, 1, 1), #131 + + Conv1d(4, 8, 3, 2,1), #66 + Conv1d(8, 8, 3, 1, 1,residual=True), + Conv1d(8, 8, 3, 1, 1,residual=True), + + Conv1d(8, 16, 3, 2, 1), # 33 + Conv1d(16, 16, 3, 1, 1, residual=True), + Conv1d(16, 16, 3, 1, 1, residual=True), + + Conv1d(16, 32, 3, 2,1),# 17 + Conv1d(32, 32, 3, 1, 1,residual=True), + Conv1d(32, 32, 3, 1, 1,residual=True), + + Conv1d(32, 64, 3, 2,1), # 9 + Conv1d(64, 64, 3, 1, 1,residual=True), + Conv1d(64, 64, 3, 1, 1,residual=True), + + Conv1d(64, 128, 3, 2,1), # 5 + Conv1d(128, 128, 3, 1, 1,residual=True), + Conv1d(128, 128, 3, 1, 1,residual=True), + + Conv1d(128, 256, 3, 2,1), #3 + Conv1d(256, 256, 3, 1, 1,residual=True), + + Conv1d(256, 512, 3, 1,0), #1 + Conv1d(512, 512, 1, 1,0,act='Tanh'), #1 + ) + self.pose_encoder=nn.Sequential( # (B*T,2,74) + Conv1d(2, 4, 3, 1, 1), + + Conv1d(4, 8, 3, 1, 1), #74 + Conv1d(8, 8, 3, 1, 1,residual=True), + Conv1d(8, 8, 3, 1, 1, residual=True), + + Conv1d(8, 16, 3, 2, 1), # 37 + Conv1d(16, 16, 3, 1, 1, residual=True), + Conv1d(16, 16, 3, 1, 1, residual=True), + + Conv1d(16, 32, 3, 2, 1), # 19 + Conv1d(32, 32, 3, 1, 1, residual=True), + Conv1d(32, 32, 3, 1, 1, residual=True), + + Conv1d(32, 64, 3, 2, 1), #10 + Conv1d(64, 64, 3, 1, 1, residual=True), + Conv1d(64, 64, 3, 1, 1, residual=True), + + Conv1d(64, 128, 3, 2, 1), # 5 + Conv1d(128, 128, 3, 1, 1, residual=True), + Conv1d(128, 128, 3, 1, 1, residual=True), + + Conv1d(128, 256, 3, 2, 1), # 3 + Conv1d(256, 256, 3, 1, 1, residual=True), + Conv1d(256, 256, 3, 1, 1, residual=True), + + Conv1d(256, 512, 3, 1, 0), # 1 + Conv1d(512, 512, 1, 1, 0, residual=True,act='Tanh'), + ) + + self.fusion_transformer = Fusion_transformer_encoder(T,d_model,nlayers,nhead,dim_feedforward,dropout) + + self.mouse_keypoint_map = nn.Linear(d_model, 40 * 2) + self.jaw_keypoint_map = nn.Linear(d_model, 17 * 2) + + self.apply(weight_init) + self.Norm=nn.LayerNorm(512) + + def forward(self, T_mels, T_pose, Nl_pose, Nl_content): + # (B,T,1,hv,wv) (B,T,2,74) (B,N_l,2,74) (B,N_l,2,57) + B,T,N_l= T_mels.size(0),T_mels.size(1),Nl_content.size(1) + + #1. obtain full reference landmarks + Nl_ref = torch.cat([Nl_pose, Nl_content], dim=3) #(B,Nl,2,131=74+57) + Nl_ref = torch.cat([Nl_ref[i] for i in range(Nl_ref.size(0))], dim=0) # (B*Nl,2,131) + + T_mels=torch.cat([T_mels[i] for i in range(T_mels.size(0))],dim=0) #(B*T,1,hv,wv) + T_pose = torch.cat([T_pose[i] for i in range(T_pose.size(0))],dim=0) # (B*T,2,74) + + # 2. get embedding + mel_embedding=self.mel_encoder(T_mels).squeeze(-1).squeeze(-1)#(B*T,512) + pose_embedding=self.pose_encoder(T_pose).squeeze(-1) # (B*T,512) + ref_embedding = self.ref_encoder(Nl_ref).squeeze(-1) # (B*Nl,512) + #normalization + mel_embedding = self.Norm(mel_embedding) # (B*T,512) + pose_embedding =self.Norm(pose_embedding) # (B*T,512) + ref_embedding = self.Norm(ref_embedding) # (B*Nl,512) + + mel_embedding = torch.stack(torch.split(mel_embedding,T),dim=0) #(B,T,512) + pose_embedding = torch.stack(torch.split(pose_embedding, T), dim=0) # (B,T,512) + ref_embedding=torch.stack(torch.split(ref_embedding,N_l,dim=0),dim=0) #(B,N_l,512) + + #3. fuse embedding + output_tokens=self.fusion_transformer(ref_embedding,mel_embedding,pose_embedding) + + #4.output landmark + lip_embedding=output_tokens[:,N_l:N_l+T,:] #(B,T,dim) + jaw_embedding=output_tokens[:,N_l+T:,:] #(B,T,dim) + output_mouse_landmark=self.mouse_keypoint_map(lip_embedding) ##(B,T,40*2) + output_jaw_landmark=self.jaw_keypoint_map(jaw_embedding) ##(B,T,17*2) + + predict_content=torch.reshape(torch.cat([output_jaw_landmark,output_mouse_landmark],dim=2),(B,T,-1,2)) #(B,T,57,2) + predict_content=torch.cat([predict_content[i] for i in range(predict_content.size(0))],dim=0).permute(0,2,1)#(B*T,2,57) + return predict_content #(B*T,2,57) + diff --git a/ip_lap/models/pix2pixHD_disc.py b/ip_lap/models/pix2pixHD_disc.py new file mode 100644 index 0000000..6a13ea5 --- /dev/null +++ b/ip_lap/models/pix2pixHD_disc.py @@ -0,0 +1,137 @@ +import torch +import torch.nn as nn +import functools +from torch.autograd import Variable +import numpy as np + + +def weights_init(m): + classname = m.__class__.__name__ + if classname.find('Conv') != -1: + m.weight.data.normal_(0.0, 0.02) + elif classname.find('BatchNorm2d') != -1: + m.weight.data.normal_(1.0, 0.02) + m.bias.data.fill_(0) + + +def define_D(input_nc=3, ndf=64, n_layers_D=3, norm='instance', use_sigmoid=False, num_D=2, getIntermFeat=True): + #('--ndf', type=int, default=64, help='# of discrim filters in first conv layer') + #('--input_nc', type=int, default=3, help='# of input image channels') + #('--n_layers_D', type=int, default=3, help='only used if which_model_netD==n_layers') + # ('--num_D', type=int, default=2, help='number of discriminators to use') + + norm_layer = get_norm_layer(norm_type=norm) + netD = MultiscaleDiscriminator(input_nc, ndf, n_layers_D, norm_layer, use_sigmoid, num_D, getIntermFeat) + #print(netD) + netD.apply(weights_init) + return netD + + +class NLayerDiscriminator(nn.Module): + def __init__(self, input_nc, ndf=64, n_layers=3, norm_layer=nn.BatchNorm2d, use_sigmoid=False, getIntermFeat=False): + super(NLayerDiscriminator, self).__init__() + self.getIntermFeat = getIntermFeat + self.n_layers = n_layers + + kw = 4 + padw = int(np.ceil((kw-1.0)/2)) + sequence = [[nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw), nn.LeakyReLU(0.2, True)]] + + nf = ndf + for n in range(1, n_layers): + nf_prev = nf + nf = min(nf * 2, 512) + sequence += [[ + nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=2, padding=padw), + norm_layer(nf), nn.LeakyReLU(0.2, True) + ]] + + nf_prev = nf + nf = min(nf * 2, 512) + sequence += [[ + nn.Conv2d(nf_prev, nf, kernel_size=kw, stride=1, padding=padw), + norm_layer(nf), + nn.LeakyReLU(0.2, True) + ]] + + sequence += [[nn.Conv2d(nf, 1, kernel_size=kw, stride=1, padding=padw)]] + + if use_sigmoid: + sequence += [[nn.Sigmoid()]] + + if getIntermFeat: + for n in range(len(sequence)): + setattr(self, 'model'+str(n), nn.Sequential(*sequence[n])) + else: + sequence_stream = [] + for n in range(len(sequence)): + sequence_stream += sequence[n] + self.model = nn.Sequential(*sequence_stream) + + def forward(self, input): + + if self.getIntermFeat: + res = [input] + for n in range(self.n_layers+2): + model = getattr(self, 'model'+str(n)) + res.append(model(res[-1])) + return res[1:] + else: + return self.model(input) + + +def get_norm_layer(norm_type='instance'): + if norm_type == 'batch': + norm_layer = functools.partial(nn.BatchNorm2d, affine=True) + elif norm_type == 'instance': + norm_layer = functools.partial(nn.InstanceNorm2d, affine=False) + else: + raise NotImplementedError('normalization layer [%s] is not found' % norm_type) + return norm_layer + + + + +class MultiscaleDiscriminator(nn.Module): + def __init__(self, input_nc, ndf=64, n_layers=3, norm_layer=nn.BatchNorm2d, + use_sigmoid=False, num_D=3, getIntermFeat=False): + super(MultiscaleDiscriminator, self).__init__() + self.num_D = num_D + self.n_layers = n_layers + self.getIntermFeat = getIntermFeat + + for i in range(num_D): + netD = NLayerDiscriminator(input_nc, ndf, n_layers, norm_layer, use_sigmoid, getIntermFeat) + if getIntermFeat: + for j in range(n_layers + 2): + setattr(self, 'scale' + str(i) + '_layer' + str(j), getattr(netD, 'model' + str(j))) + else: + setattr(self, 'layer' + str(i), netD.model) + + self.downsample = nn.AvgPool2d(3, stride=2, padding=[1, 1], count_include_pad=False) + + def singleD_forward(self, model, input): + if self.getIntermFeat: + result = [input] + for i in range(len(model)): + result.append(model[i](result[-1])) + return result[1:] + else: + return [model(input)] + + def forward(self, input): #: (B,T,C,H,W) + # input = torch.cat([input[i,:] for i in range(input.size(0))], dim=0)# : (B*T,C,H,W) + num_D = self.num_D + result = [] + input_downsampled = input + for i in range(num_D): + if self.getIntermFeat: + model = [getattr(self, 'scale' + str(num_D - 1 - i) + '_layer' + str(j)) for j in + range(self.n_layers + 2)] + else: + model = getattr(self, 'layer' + str(num_D - 1 - i)) + result.append(self.singleD_forward(model, input_downsampled)) + if i != (num_D - 1): + input_downsampled = self.downsample(input_downsampled) + return result + diff --git a/ip_lap/models/video_renderer.py b/ip_lap/models/video_renderer.py new file mode 100644 index 0000000..9fb2443 --- /dev/null +++ b/ip_lap/models/video_renderer.py @@ -0,0 +1,571 @@ +from torch.nn import functional as F +import torch +import torch.nn as nn +import torchvision + + + +class AdaINLayer(nn.Module): + def __init__(self, input_nc, modulation_nc): + super().__init__() + + self.InstanceNorm2d = nn.InstanceNorm2d(input_nc, affine=False) + + nhidden = 128 + use_bias=True + + self.mlp_shared = nn.Sequential( + nn.Linear(modulation_nc, nhidden, bias=use_bias), + nn.ReLU() + ) + self.mlp_gamma = nn.Linear(nhidden, input_nc, bias=use_bias) + self.mlp_beta = nn.Linear(nhidden, input_nc, bias=use_bias) + + def forward(self, input, modulation_input): + + # Part 1. generate parameter-free normalized activations + normalized = self.InstanceNorm2d(input) + + # Part 2. produce scaling and bias conditioned on feature + modulation_input = modulation_input.view(modulation_input.size(0), -1) + actv = self.mlp_shared(modulation_input) + gamma = self.mlp_gamma(actv) + beta = self.mlp_beta(actv) + + # apply scale and bias + gamma = gamma.view(*gamma.size()[:2], 1,1) + beta = beta.view(*beta.size()[:2], 1,1) + out = normalized * (1 + gamma) + beta + return out + +class AdaIN(torch.nn.Module): + + def __init__(self, input_channel, modulation_channel,kernel_size=3, stride=1, padding=1): + super(AdaIN, self).__init__() + self.conv_1 = torch.nn.Conv2d(input_channel, input_channel, kernel_size=kernel_size, stride=stride, padding=padding) + self.conv_2 = torch.nn.Conv2d(input_channel, input_channel, kernel_size=kernel_size, stride=stride, padding=padding) + self.leaky_relu = torch.nn.LeakyReLU(0.2) + self.adain_layer_1 = AdaINLayer(input_channel, modulation_channel) + self.adain_layer_2 = AdaINLayer(input_channel, modulation_channel) + + def forward(self, x, modulation): + + x = self.adain_layer_1(x, modulation) + x = self.leaky_relu(x) + x = self.conv_1(x) + x = self.adain_layer_2(x, modulation) + x = self.leaky_relu(x) + x = self.conv_2(x) + + return x + + + + +class SPADELayer(torch.nn.Module): + def __init__(self, input_channel, modulation_channel, hidden_size=256, kernel_size=3, stride=1, padding=1): + super(SPADELayer, self).__init__() + self.instance_norm = torch.nn.InstanceNorm2d(input_channel) + + self.conv1 = torch.nn.Conv2d(modulation_channel, hidden_size, kernel_size=kernel_size, stride=stride, + padding=padding) + self.gamma = torch.nn.Conv2d(hidden_size, input_channel, kernel_size=kernel_size, stride=stride, padding=padding) + self.beta = torch.nn.Conv2d(hidden_size, input_channel, kernel_size=kernel_size, stride=stride, padding=padding) + + def forward(self, input, modulation): + norm = self.instance_norm(input) + + conv_out = self.conv1(modulation) + + gamma = self.gamma(conv_out) + beta = self.beta(conv_out) + + return norm + norm * gamma + beta + + +class SPADE(torch.nn.Module): + def __init__(self, num_channel, num_channel_modulation, hidden_size=256, kernel_size=3, stride=1, padding=1): + super(SPADE, self).__init__() + self.conv_1 = torch.nn.Conv2d(num_channel, num_channel, kernel_size=kernel_size, stride=stride, padding=padding) + self.conv_2 = torch.nn.Conv2d(num_channel, num_channel, kernel_size=kernel_size, stride=stride, padding=padding) + self.leaky_relu = torch.nn.LeakyReLU(0.2) + self.spade_layer_1 = SPADELayer(num_channel, num_channel_modulation, hidden_size, kernel_size=kernel_size, + stride=stride, padding=padding) + self.spade_layer_2 = SPADELayer(num_channel, num_channel_modulation, hidden_size, kernel_size=kernel_size, + stride=stride, padding=padding) + + def forward(self, input, modulations): + input = self.spade_layer_1(input, modulations) + input = self.leaky_relu(input) + input = self.conv_1(input) + input = self.spade_layer_2(input, modulations) + input = self.leaky_relu(input) + input = self.conv_2(input) + return input + +class Conv2d(nn.Module): + def __init__(self, cin, cout, kernel_size, stride, padding, residual=False, *args, **kwargs): + super().__init__(*args, **kwargs) + self.conv_block = nn.Sequential( + nn.Conv2d(cin, cout, kernel_size, stride, padding), + nn.BatchNorm2d(cout) + ) + self.act = nn.ReLU() + self.residual = residual + + def forward(self, x): + out = self.conv_block(x) + if self.residual: + out += x + return self.act(out) + +def downsample(x, size): + if len(x.size()) == 5: + size = (x.size(2), size[0], size[1]) + return torch.nn.functional.interpolate(x, size=size, mode='nearest') + return torch.nn.functional.interpolate(x, size=size, mode='nearest') + + +def convert_flow_to_deformation(flow): + r"""convert flow fields to deformations. + Args: + flow (tensor): Flow field obtained by the model + Returns: + deformation (tensor): The deformation used for warpping + """ + b, c, h, w = flow.shape + flow_norm = 2 * torch.cat([flow[:, :1, ...] / (w - 1), flow[:, 1:, ...] / (h - 1)], 1) + grid = make_coordinate_grid(flow) + deformation = grid + flow_norm.permute(0, 2, 3, 1) + return deformation + + +def make_coordinate_grid(flow): + r"""obtain coordinate grid with the same size as the flow filed. + Args: + flow (tensor): Flow field obtained by the model + Returns: + grid (tensor): The grid with the same size as the input flow + """ + b, c, h, w = flow.shape + + x = torch.arange(w).to(flow) + y = torch.arange(h).to(flow) + + x = (2 * (x / (w - 1)) - 1) + y = (2 * (y / (h - 1)) - 1) + + yy = y.view(-1, 1).repeat(1, w) + xx = x.view(1, -1).repeat(h, 1) + + meshed = torch.cat([xx.unsqueeze_(2), yy.unsqueeze_(2)], 2) + meshed = meshed.expand(b, -1, -1, -1) + return meshed + + +def warping(source_image, deformation): + r"""warp the input image according to the deformation + Args: + source_image (tensor): source images to be warpped + deformation (tensor): deformations used to warp the images; value in range (-1, 1) + Returns: + output (tensor): the warpped images + """ + _, h_old, w_old, _ = deformation.shape + _, _, h, w = source_image.shape + if h_old != h or w_old != w: + deformation = deformation.permute(0, 3, 1, 2) + deformation = torch.nn.functional.interpolate(deformation, size=(h, w), mode='bilinear') + deformation = deformation.permute(0, 2, 3, 1) + return torch.nn.functional.grid_sample(source_image, deformation) + + +class DenseFlowNetwork(torch.nn.Module): + def __init__(self, num_channel=6, num_channel_modulation=3*5, hidden_size=256): + super(DenseFlowNetwork, self).__init__() + + # Convolutional Layers + self.conv1 = torch.nn.Conv2d(num_channel, 32, kernel_size=7, stride=1, padding=3) + self.conv1_bn = torch.nn.BatchNorm2d(num_features=32, affine=True) + self.conv1_relu = torch.nn.ReLU() + + self.conv2 = torch.nn.Conv2d(32, 256, kernel_size=3, stride=2, padding=1) + self.conv2_bn = torch.nn.BatchNorm2d(num_features=256, affine=True) + self.conv2_relu = torch.nn.ReLU() + + + # SPADE Blocks + self.spade_layer_1 = SPADE(256, num_channel_modulation, hidden_size) + self.spade_layer_2 = SPADE(256, num_channel_modulation, hidden_size) + self.pixel_shuffle_1 = torch.nn.PixelShuffle(2) + self.spade_layer_4 = SPADE(64, num_channel_modulation, hidden_size) + + # Final Convolutional Layer + self.conv_4 = torch.nn.Conv2d(64, 2, kernel_size=7, stride=1, padding=3) + self.conv_5= nn.Sequential(torch.nn.Conv2d(64, 32, kernel_size=7, stride=1, padding=3), + torch.nn.ReLU(), + torch.nn.Conv2d(32, 1, kernel_size=7, stride=1, padding=3), + torch.nn.Sigmoid(), + )#predict weight + + def forward(self, ref_N_frame_img, ref_N_frame_sketch, T_driving_sketch): #to output: (B*T,3,H,W) + # (B, N, 3, H, W)(B, N, 3, H, W) (B, 5, 3, H, W) # + ref_N = ref_N_frame_img.size(1) + + driving_sketch=torch.cat([T_driving_sketch[:,i] for i in range(T_driving_sketch.size(1))], dim=1) #(B, 3*5, H, W) + + wrapped_h1_sum, wrapped_h2_sum, wrapped_ref_sum=0.,0.,0. + softmax_denominator=0. + T = 1 # during rendering, generate T=1 image at a time + for ref_idx in range(ref_N): # each ref img provide information for each B*T frame + ref_img= ref_N_frame_img[:, ref_idx] #(B, 3, H, W) + ref_img = ref_img.unsqueeze(1).expand(-1, T, -1, -1, -1) # (B,T, 3, H, W) + ref_img = torch.cat([ref_img[i] for i in range(ref_img.size(0))], dim=0) # (B*T, 3, H, W) + + ref_sketch = ref_N_frame_sketch[:, ref_idx] #(B, 3, H, W) + ref_sketch = ref_sketch.unsqueeze(1).expand(-1, T, -1, -1, -1) # (B,T, 3, H, W) + ref_sketch = torch.cat([ref_sketch[i] for i in range(ref_sketch.size(0))], dim=0) # (B*T, 3, H, W) + + #predict flow and weight + flow_module_input = torch.cat((ref_img, ref_sketch), dim=1) #(B*T, 3+3, H, W) + # Convolutional Layers + h1 = self.conv1_relu(self.conv1_bn(self.conv1(flow_module_input))) #(32,128,128) + h2 = self.conv2_relu(self.conv2_bn(self.conv2(h1))) #(256,64,64) + # SPADE Blocks + downsample_64 = downsample(driving_sketch, (64, 64)) # driving_sketch:(B*T, 3, H, W) + + spade_layer = self.spade_layer_1(h2, downsample_64) #(256,64,64) + spade_layer = self.spade_layer_2(spade_layer, downsample_64) #(256,64,64) + + spade_layer = self.pixel_shuffle_1(spade_layer) #(64,128,128) + + spade_layer = self.spade_layer_4(spade_layer, driving_sketch) #(64,128,128) + + # Final Convolutional Layer + output_flow = self.conv_4(spade_layer) # (B*T,2,128,128) + output_weight=self.conv_5(spade_layer) # (B*T,1,128,128) + + deformation=convert_flow_to_deformation(output_flow) + wrapped_h1 = warping(h1, deformation) #(32,128,128) + wrapped_h2 = warping(h2, deformation) #(256,64,64) + wrapped_ref = warping(ref_img, deformation) #(3,128,128) + + softmax_denominator+=output_weight + wrapped_h1_sum+=wrapped_h1*output_weight + wrapped_h2_sum+=wrapped_h2*downsample(output_weight, (64,64)) + wrapped_ref_sum+=wrapped_ref*output_weight + #return weighted warped feataure and images + softmax_denominator+=0.00001 + wrapped_h1_sum=wrapped_h1_sum/softmax_denominator + wrapped_h2_sum = wrapped_h2_sum / downsample(softmax_denominator, (64,64)) + wrapped_ref_sum = wrapped_ref_sum / softmax_denominator + return wrapped_h1_sum, wrapped_h2_sum, wrapped_ref_sum + + +class TranslationNetwork(torch.nn.Module): + def __init__(self): + super(TranslationNetwork, self).__init__() + self.audio_encoder = nn.Sequential( + Conv2d(1, 32, kernel_size=3, stride=1, padding=1), + Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True), + Conv2d(32, 32, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(32, 64, kernel_size=3, stride=(3, 1), padding=1), + Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True), + Conv2d(64, 64, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(64, 128, kernel_size=3, stride=3, padding=1), + Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True), + Conv2d(128, 128, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(128, 256, kernel_size=3, stride=(3, 2), padding=1), + Conv2d(256, 256, kernel_size=3, stride=1, padding=1, residual=True), + + Conv2d(256, 512, kernel_size=3, stride=1, padding=0), + Conv2d(512, 512, kernel_size=1, stride=1, padding=0), ) + + # Encoder + self.conv1 = torch.nn.Conv2d(in_channels=3+3*5, out_channels=32, kernel_size=7, stride=1, padding=3, bias=False) + self.conv1_bn = torch.nn.BatchNorm2d(num_features=32, affine=True) + self.conv1_relu = torch.nn.ReLU() + + self.conv2 = torch.nn.Conv2d(in_channels=32, out_channels=256, kernel_size=3, stride=2, padding=1, bias=False) + self.conv2_bn = torch.nn.BatchNorm2d(num_features=256, affine=True) + self.conv2_relu = torch.nn.ReLU() + + # Decoder + self.spade_1 = SPADE(num_channel=256, num_channel_modulation=256) + self.adain_1 = AdaIN(256,512) + self.pixel_suffle_1 = nn.PixelShuffle(upscale_factor=2) + + self.spade_2 = SPADE(num_channel=64, num_channel_modulation=32) + self.adain_2 = AdaIN(input_channel=64,modulation_channel=512) + + self.spade_4 = SPADE(num_channel=64, num_channel_modulation=3) + + # Final layer + self.leaky_relu = torch.nn.LeakyReLU() + self.conv_last = torch.nn.Conv2d(in_channels=64, out_channels=3, kernel_size=7, stride=1, padding=3, bias=False) + self.Sigmoid=torch.nn.Sigmoid() + def forward(self, translation_input, wrapped_ref, wrapped_h1, wrapped_h2, T_mels): + # (B,3+3,H,W) (B,3,128,128) (B,32,128,128) (B,256,64,64) (B,T,1,h,w) #T=1 + # Encoder + T_mels=torch.cat([T_mels[i] for i in range(T_mels.size(0))],dim=0)# B*T,1,h,w + x = self.conv1_relu(self.conv1_bn(self.conv1(translation_input))) #32,128,128 + x = self.conv2_relu(self.conv2_bn(self.conv2(x))) #256,64,64 + + audio_feature = self.audio_encoder(T_mels).squeeze(-1).permute(0,2,1) #(B*T,1,512) + + # Decoder + x = self.spade_1(x, wrapped_h2) # (C=256,64,64) + x = self.adain_1(x, audio_feature) # (C=256,64,64) + x = self.pixel_suffle_1(x) # (C=64,128,128) + + x = self.spade_2(x, wrapped_h1) # (64,128,128) + x = self.adain_2(x, audio_feature) # (64,128,128) + x = self.spade_4(x, wrapped_ref) # (64,128,128) + + # output layer + x = self.leaky_relu(x) + x = self.conv_last(x) + x = self.Sigmoid(x) + return x + +class Renderer(torch.nn.Module): + def __init__(self): + super(Renderer, self).__init__() + + # 1.flow Network + self.flow_module = DenseFlowNetwork() + #2. translation Network + self.translation = TranslationNetwork() + #3.return loss + self.perceptual = PerceptualLoss(network='vgg19', + layers=['relu_1_1', 'relu_2_1', 'relu_3_1', 'relu_4_1', 'relu_5_1'], + num_scales=2) + + def forward(self, face_frame_img, target_sketches, ref_N_frame_img, ref_N_frame_sketch, audio_mels): #T=1 + # (B,1,3,H,W) (B,5,3,H,W) (B,N,3,H,W) (B,N,3,H,W) (B,T,1,hv,wv)T=1 + # (1)warping reference images and their feature + wrapped_h1, wrapped_h2, wrapped_ref = self.flow_module(ref_N_frame_img, ref_N_frame_sketch, target_sketches) + #(B,C,H,W) + + # (2)translation module + target_sketches = torch.cat([target_sketches[:, i] for i in range(target_sketches.size(1))], dim=1) + # (B,3*T,H,W) + gt_face = torch.cat([face_frame_img[i] for i in range(face_frame_img.size(0))], dim=0) + # (B,3,H,W) + gt_mask_face = gt_face.clone() + gt_mask_face[:, :, gt_mask_face.size(2) // 2:, :] = 0 # (B,3,H,W) + # + translation_input=torch.cat([gt_mask_face, target_sketches], dim=1) # (B*T,3+3,H,W) + generated_face = self.translation(translation_input, wrapped_ref, wrapped_h1, wrapped_h2, audio_mels) #translation_input + + perceptual_gen_loss = self.perceptual(generated_face, gt_face, use_style_loss=True, + weight_style_to_perceptual=250).mean() + perceptual_warp_loss = self.perceptual(wrapped_ref, gt_face, use_style_loss=False, + weight_style_to_perceptual=0.).mean() + return generated_face, wrapped_ref, torch.unsqueeze(perceptual_warp_loss, 0), torch.unsqueeze( + perceptual_gen_loss, 0) + # (B,3,H,W) and losses + +#the following is the code for Perceptual(VGG) loss + +def apply_imagenet_normalization(input): + r"""Normalize using ImageNet mean and std. + + Args: + input (4D tensor NxCxHxW): The input images, assuming to be [-1, 1]. + + Returns: + Normalized inputs using the ImageNet normalization. + """ + # normalize the input back to [0, 1] + normalized_input = (input + 1) / 2 + # normalize the input using the ImageNet mean and std + mean = normalized_input.new_tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) + std = normalized_input.new_tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) + output = (normalized_input - mean) / std + return output + +class _PerceptualNetwork(nn.Module): + r"""The network that extracts features to compute the perceptual loss. + + Args: + network (nn.Sequential) : The network that extracts features. + layer_name_mapping (dict) : The dictionary that + maps a layer's index to its name. + layers (list of str): The list of layer names that we are using. + """ + + def __init__(self, network, layer_name_mapping, layers): + super().__init__() + assert isinstance(network, nn.Sequential), \ + 'The network needs to be of type "nn.Sequential".' + self.network = network + self.layer_name_mapping = layer_name_mapping + self.layers = layers + for param in self.parameters(): + param.requires_grad = False + + def forward(self, x): + r"""Extract perceptual features.""" + output = {} + for i, layer in enumerate(self.network): + x = layer(x) + layer_name = self.layer_name_mapping.get(i, None) + if layer_name in self.layers: + # If the current layer is used by the perceptual loss. + output[layer_name] = x + return output + +def _vgg19(layers): + r"""Get vgg19 layers""" + network = torchvision.models.vgg19(pretrained=True).features + layer_name_mapping = {1: 'relu_1_1', + 3: 'relu_1_2', + 6: 'relu_2_1', + 8: 'relu_2_2', + 11: 'relu_3_1', + 13: 'relu_3_2', + 15: 'relu_3_3', + 17: 'relu_3_4', + 20: 'relu_4_1', + 22: 'relu_4_2', + 24: 'relu_4_3', + 26: 'relu_4_4', + 29: 'relu_5_1'} + return _PerceptualNetwork(network, layer_name_mapping, layers) + +class PerceptualLoss(nn.Module): + r"""Perceptual loss initialization. + + Args: + network (str) : The name of the loss network: 'vgg16' | 'vgg19'. + layers (str or list of str) : The layers used to compute the loss. + weights (float or list of float : The loss weights of each layer. + criterion (str): The type of distance function: 'l1' | 'l2'. + resize (bool) : If ``True``, resize the input images to 224x224. + resize_mode (str): Algorithm used for resizing. + instance_normalized (bool): If ``True``, applies instance normalization + to the feature maps before computing the distance. + num_scales (int): The loss will be evaluated at original size and + this many times downsampled sizes. + """ + + def __init__(self, network='vgg19', layers='relu_4_1', weights=None, + criterion='l1', resize=False, resize_mode='bilinear', + instance_normalized=False, num_scales=1,): + super().__init__() + if isinstance(layers, str): + layers = [layers] + if weights is None: + weights = [1.] * len(layers) + elif isinstance(layers, float) or isinstance(layers, int): + weights = [weights] + + assert len(layers) == len(weights), \ + 'The number of layers (%s) must be equal to ' \ + 'the number of weights (%s).' % (len(layers), len(weights)) + if network == 'vgg19': + self.model = _vgg19(layers) + else: + raise ValueError('Network %s is not recognized' % network) + + self.num_scales = num_scales + self.layers = layers + self.weights = weights + if criterion == 'l1': + self.criterion = nn.L1Loss() + elif criterion == 'l2' or criterion == 'mse': + self.criterion = nn.MSELoss() + else: + raise ValueError('Criterion %s is not recognized' % criterion) + self.resize = resize + self.resize_mode = resize_mode + self.instance_normalized = instance_normalized + + + print('Perceptual loss:') + print('\tMode: {}'.format(network)) + + def forward(self, inp, target, mask=None,use_style_loss=False,weight_style_to_perceptual=0.): + r"""Perceptual loss forward. + + Args: + inp (4D tensor) : Input tensor. + target (4D tensor) : Ground truth tensor, same shape as the input. + + Returns: + (scalar tensor) : The perceptual loss. + """ + # Perceptual loss should operate in eval mode by default. + self.model.eval() + inp, target = \ + apply_imagenet_normalization(inp), \ + apply_imagenet_normalization(target) + if self.resize: + inp = F.interpolate( + inp, mode=self.resize_mode, size=(256, 256), + align_corners=False) + target = F.interpolate( + target, mode=self.resize_mode, size=(256, 256), + align_corners=False) + + # Evaluate perceptual loss at each scale. + loss = 0 + style_loss=0 + for scale in range(self.num_scales): + input_features, target_features = \ + self.model(inp), self.model(target) + for layer, weight in zip(self.layers, self.weights): + # Example per-layer VGG19 loss values after applying + # [0.03125, 0.0625, 0.125, 0.25, 1.0] weighting. + # relu_1_1, 0.014698 + # relu_2_1, 0.085817 + # relu_3_1, 0.349977 + # relu_4_1, 0.544188 + # relu_5_1, 0.906261 + input_feature = input_features[layer] + target_feature = target_features[layer].detach() + if self.instance_normalized: + input_feature = F.instance_norm(input_feature) + target_feature = F.instance_norm(target_feature) + + if mask is not None: + mask_ = F.interpolate(mask, input_feature.shape[2:], + mode='bilinear', + align_corners=False) + input_feature = input_feature * mask_ + target_feature = target_feature * mask_ + # print('mask',mask_.shape) + + + loss += weight * self.criterion(input_feature, + target_feature) + if use_style_loss and scale==0: + style_loss += self.criterion(self.compute_gram(input_feature), + self.compute_gram(target_feature)) + + # Downsample the input and target. + if scale != self.num_scales - 1: + inp = F.interpolate( + inp, mode=self.resize_mode, scale_factor=0.5, + align_corners=False, recompute_scale_factor=True) + target = F.interpolate( + target, mode=self.resize_mode, scale_factor=0.5, + align_corners=False, recompute_scale_factor=True) + + if use_style_loss: + return loss + style_loss*weight_style_to_perceptual + else: + return loss + + + def compute_gram(self, x): + b, ch, h, w = x.size() + f = x.view(b, ch, w * h) + f_T = f.transpose(1, 2) + G = f.bmm(f_T) / (h * w * ch) + return G + diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..fe81d67 --- /dev/null +++ b/nodes.py @@ -0,0 +1,143 @@ +import os +import platform +import subprocess +import folder_paths +from pydub import AudioSegment +from moviepy.editor import VideoFileClip,AudioFileClip + +parent_directory = os.path.dirname(os.path.abspath(__file__)) + +from .ip_lap.inference import IP_LAP_infer + +input_path = folder_paths.get_input_directory() +out_path = folder_paths.get_output_directory() + +class CombineAudioVideo: + @classmethod + def INPUT_TYPES(s): + return {"required": + {"vocal_AUDIO": ("AUDIO",), + "bgm_AUDIO": ("AUDIO",), + "video": ("VIDEO",) + } + } + + CATEGORY = "AIFSH_IP_LAP" + DESCRIPTION = "hello world!" + + RETURN_TYPES = ("VIDEO",) + + OUTPUT_NODE = False + + FUNCTION = "combine" + + def combine(self, vocal_AUDIO,bgm_AUDIO,video): + vocal = AudioSegment.from_file(vocal_AUDIO) + bgm = AudioSegment.from_file(bgm_AUDIO) + audio = vocal.overlay(bgm) + audio_file = os.path.join(out_path,"ip_lap_voice.wav") + audio.export(audio_file, format="wav") + cm_video_file = os.path.join(out_path,"voice_"+os.path.basename(video)) + video_clip = VideoFileClip(video) + audio_clip = AudioFileClip(audio_file) + new_video_clip = video_clip.set_audio(audio_clip) + new_video_clip.write_videofile(cm_video_file) + return (cm_video_file,) + + +class PreViewVideo: + @classmethod + def INPUT_TYPES(s): + return {"required":{ + "video":("VIDEO",), + }} + + CATEGORY = "AIFSH_IP_LAP" + DESCRIPTION = "hello world!" + + RETURN_TYPES = () + + OUTPUT_NODE = True + + FUNCTION = "load_video" + + def load_video(self, video): + video_name = os.path.basename(video) + video_path_name = os.path.basename(os.path.dirname(video)) + return {"ui":{"video":[video_name,video_path_name]}} + +class LoadVideo: + @classmethod + def INPUT_TYPES(s): + files = [f for f in os.listdir(input_path) if os.path.isfile(os.path.join(input_path, f)) and f.split('.')[-1] in ["mp4", "webm","mkv","avi"]] + return {"required":{ + "video":(files,), + }} + + CATEGORY = "AIFSH_IP_LAP" + DESCRIPTION = "hello world!" + + RETURN_TYPES = ("VIDEO","AUDIO") + + OUTPUT_NODE = False + + FUNCTION = "load_video" + + def load_video(self, video): + video_path = os.path.join(input_path,video) + video_clip = VideoFileClip(video_path) + audio_path = os.path.join(input_path,video+".wav") + video_clip.audio.write_audiofile(audio_path) + return (video_path,audio_path,) + +class IP_LAP: + + @classmethod + def INPUT_TYPES(s): + return {"required": + { + "audio": ("AUDIO",), + "video": ("VIDEO",), + "T":("INT",{ + "default": 5, + }), + "Nl":("INT",{ + "default": 15, + }), + "ref_img_N":("INT",{ + "default": 25, + }), + "img_size":("INT",{ + "default": 128, + }), + "mel_step_size":("INT",{ + "default": 16, + }), + "face_det_batch_size":("INT",{ + "default": 4, + }), + "checkpoints_path":("STRING",{ + "default": os.path.join(parent_directory,"weights") + }) + } + } + + CATEGORY = "AIFSH_IP_LAP" + DESCRIPTION = "hello world!" + + RETURN_TYPES = ("VIDEO",) + + OUTPUT_NODE = False + + FUNCTION = "process" + + def process(self, audio, video, T=5,Nl=15,ref_img_N=25,img_size=128, + mel_step_size=16,face_det_batch_size=4,checkpoints_path=""): + ip_lap = IP_LAP_infer(T,Nl,ref_img_N,img_size,mel_step_size,face_det_batch_size,checkpoints_path) + video_name = os.path.basename(video) + out_video_file = os.path.join(out_path, f"ip_lap_{video_name}") + ip_lap(video,audio,out_video_file) + # res_video_file = os.path.join(out_path, f"result_ip_lap_{video_name}") + # command = f'ffmpeg -y -i {out_video_file} -i {audio} -map 0:0 -map 1:0 -c:a libmp3lame -q:a 1 -q:v 1 -shortest {res_video_file}' + # subprocess.call(command, shell=platform.system() != 'Windows') + return (out_video_file,) \ No newline at end of file diff --git a/note.txt b/note.txt new file mode 100644 index 0000000..738f613 --- /dev/null +++ b/note.txt @@ -0,0 +1,10 @@ +1. ModuleNotFoundError: No module named 'torchvision.transforms.functional_tensor' + +from +from torchvision.transforms.functional_tensor import rgb_to_grayscale +to +from torchvision.transforms.functional import rgb_to_grayscale + +2.ImportError: libGL.so.1: cannot open shared object file: No such file or directory +apt update +apt install ffmpeg -y \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..ace7f0f --- /dev/null +++ b/requirements.txt @@ -0,0 +1,7 @@ +face_alignment +mediapipe +basicsr +lws +librosa +moviepy +pydub \ No newline at end of file diff --git a/web/js/previewVideo.js b/web/js/previewVideo.js new file mode 100644 index 0000000..5c8f1ca --- /dev/null +++ b/web/js/previewVideo.js @@ -0,0 +1,155 @@ +import { app } from "../../../scripts/app.js"; +import { api } from '../../../scripts/api.js' + +function fitHeight(node) { + node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]]) + node?.graph?.setDirtyCanvas(true); +} +function chainCallback(object, property, callback) { + if (object == undefined) { + //This should not happen. + console.error("Tried to add callback to non-existant object") + return; + } + if (property in object) { + const callback_orig = object[property] + object[property] = function () { + const r = callback_orig.apply(this, arguments); + callback.apply(this, arguments); + return r + }; + } else { + object[property] = callback; + } +} + +function addPreviewOptions(nodeType) { + chainCallback(nodeType.prototype, "getExtraMenuOptions", function(_, options) { + // The intended way of appending options is returning a list of extra options, + // but this isn't used in widgetInputs.js and would require + // less generalization of chainCallback + let optNew = [] + try { + const previewWidget = this.widgets.find((w) => w.name === "videopreview"); + + let url = null + if (previewWidget.videoEl?.hidden == false && previewWidget.videoEl.src) { + //Use full quality video + //url = api.apiURL('/view?' + new URLSearchParams(previewWidget.value.params)); + url = previewWidget.videoEl.src + } + if (url) { + optNew.push( + { + content: "Open preview", + callback: () => { + window.open(url, "_blank") + }, + }, + { + content: "Save preview", + callback: () => { + const a = document.createElement("a"); + a.href = url; + a.setAttribute("download", new URLSearchParams(previewWidget.value.params).get("filename")); + document.body.append(a); + a.click(); + requestAnimationFrame(() => a.remove()); + }, + } + ); + } + if(options.length > 0 && options[0] != null && optNew.length > 0) { + optNew.push(null); + } + options.unshift(...optNew); + + } catch (error) { + console.log(error); + } + + }); +} +function previewVideo(node,file,type){ + var element = document.createElement("div"); + const previewNode = node; + var previewWidget = node.addDOMWidget("videopreview", "preview", element, { + serialize: false, + hideOnZoom: false, + getValue() { + return element.value; + }, + setValue(v) { + element.value = v; + }, + }); + previewWidget.computeSize = function(width) { + if (this.aspectRatio && !this.parentEl.hidden) { + let height = (previewNode.size[0]-20)/ this.aspectRatio + 10; + if (!(height > 0)) { + height = 0; + } + this.computedHeight = height + 10; + return [width, height]; + } + return [width, -4];//no loaded src, widget should not display + } + // element.style['pointer-events'] = "none" + previewWidget.value = {hidden: false, paused: false, params: {}} + previewWidget.parentEl = document.createElement("div"); + previewWidget.parentEl.className = "video_preview"; + previewWidget.parentEl.style['width'] = "100%" + element.appendChild(previewWidget.parentEl); + previewWidget.videoEl = document.createElement("video"); + previewWidget.videoEl.controls = true; + previewWidget.videoEl.loop = false; + previewWidget.videoEl.muted = false; + previewWidget.videoEl.style['width'] = "100%" + previewWidget.videoEl.addEventListener("loadedmetadata", () => { + + previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight; + fitHeight(this); + }); + previewWidget.videoEl.addEventListener("error", () => { + //TODO: consider a way to properly notify the user why a preview isn't shown. + previewWidget.parentEl.hidden = true; + fitHeight(this); + }); + + let params = { + "filename": file, + "type": type, + } + + previewWidget.parentEl.hidden = previewWidget.value.hidden; + previewWidget.videoEl.autoplay = !previewWidget.value.paused && !previewWidget.value.hidden; + let target_width = 256 + if (element.style?.width) { + //overscale to allow scrolling. Endpoint won't return higher than native + target_width = element.style.width.slice(0,-2)*2; + } + if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") { + params.force_size = target_width+"x?" + } else { + let size = params.force_size.split("x") + let ar = parseInt(size[0])/parseInt(size[1]) + params.force_size = target_width+"x"+(target_width/ar) + } + + previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params)); + + previewWidget.videoEl.hidden = false; + previewWidget.parentEl.appendChild(previewWidget.videoEl) +} + +app.registerExtension({ + name: "IP_LAP.VideoPreviewer", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeData?.name == "PreViewVideo") { + nodeType.prototype.onExecuted = function (data) { + previewVideo(this, data.video[0], data.video[1]); + } + addPreviewOptions(nodeType) + } + } +}); \ No newline at end of file diff --git a/web/js/uploadVideo.js b/web/js/uploadVideo.js new file mode 100644 index 0000000..1c92ce4 --- /dev/null +++ b/web/js/uploadVideo.js @@ -0,0 +1,203 @@ +import { app } from "../../../scripts/app.js"; +import { api } from '../../../scripts/api.js' +import { ComfyWidgets } from "../../../scripts/widgets.js" + +function fitHeight(node) { + node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]]) + node?.graph?.setDirtyCanvas(true); +} + +function previewVideo(node,file){ + while (node.widgets.length > 2){ + node.widgets.pop() + } + try { + var el = document.getElementById("uploadVideo"); + el.remove(); + } catch (error) { + console.log(error); + } + var element = document.createElement("div"); + element.id = "uploadVideo"; + const previewNode = node; + var previewWidget = node.addDOMWidget("videopreview", "preview", element, { + serialize: false, + hideOnZoom: false, + getValue() { + return element.value; + }, + setValue(v) { + element.value = v; + }, + }); + previewWidget.computeSize = function(width) { + if (this.aspectRatio && !this.parentEl.hidden) { + let height = (previewNode.size[0]-20)/ this.aspectRatio + 10; + if (!(height > 0)) { + height = 0; + } + this.computedHeight = height + 10; + return [width, height]; + } + return [width, -4];//no loaded src, widget should not display + } + // element.style['pointer-events'] = "none" + previewWidget.value = {hidden: false, paused: false, params: {}} + previewWidget.parentEl = document.createElement("div"); + previewWidget.parentEl.className = "video_preview"; + previewWidget.parentEl.style['width'] = "100%" + element.appendChild(previewWidget.parentEl); + previewWidget.videoEl = document.createElement("video"); + previewWidget.videoEl.controls = true; + previewWidget.videoEl.loop = false; + previewWidget.videoEl.muted = false; + previewWidget.videoEl.style['width'] = "100%" + previewWidget.videoEl.addEventListener("loadedmetadata", () => { + + previewWidget.aspectRatio = previewWidget.videoEl.videoWidth / previewWidget.videoEl.videoHeight; + fitHeight(this); + }); + previewWidget.videoEl.addEventListener("error", () => { + //TODO: consider a way to properly notify the user why a preview isn't shown. + previewWidget.parentEl.hidden = true; + fitHeight(this); + }); + + let params = { + "filename": file, + "type": "input", + } + + previewWidget.parentEl.hidden = previewWidget.value.hidden; + previewWidget.videoEl.autoplay = !previewWidget.value.paused && !previewWidget.value.hidden; + let target_width = 256 + if (element.style?.width) { + //overscale to allow scrolling. Endpoint won't return higher than native + target_width = element.style.width.slice(0,-2)*2; + } + if (!params.force_size || params.force_size.includes("?") || params.force_size == "Disabled") { + params.force_size = target_width+"x?" + } else { + let size = params.force_size.split("x") + let ar = parseInt(size[0])/parseInt(size[1]) + params.force_size = target_width+"x"+(target_width/ar) + } + + previewWidget.videoEl.src = api.apiURL('/view?' + new URLSearchParams(params)); + + previewWidget.videoEl.hidden = false; + previewWidget.parentEl.appendChild(previewWidget.videoEl) +} + +function videoUpload(node, inputName, inputData, app) { + const videoWidget = node.widgets.find((w) => w.name === "video"); + let uploadWidget; + /* + A method that returns the required style for the html + */ + var default_value = videoWidget.value; + Object.defineProperty(videoWidget, "value", { + set : function(value) { + this._real_value = value; + }, + + get : function() { + let value = ""; + if (this._real_value) { + value = this._real_value; + } else { + return default_value; + } + + if (value.filename) { + let real_value = value; + value = ""; + if (real_value.subfolder) { + value = real_value.subfolder + "/"; + } + + value += real_value.filename; + + if(real_value.type && real_value.type !== "input") + value += ` [${real_value.type}]`; + } + return value; + } + }); + async function uploadFile(file, updateNode, pasted = false) { + try { + // Wrap file in formdata so it includes filename + const body = new FormData(); + body.append("image", file); + if (pasted) body.append("subfolder", "pasted"); + const resp = await api.fetchApi("/upload/image", { + method: "POST", + body, + }); + + if (resp.status === 200) { + const data = await resp.json(); + // Add the file to the dropdown list and update the widget value + let path = data.name; + if (data.subfolder) path = data.subfolder + "/" + path; + + if (!videoWidget.options.values.includes(path)) { + videoWidget.options.values.push(path); + } + + if (updateNode) { + videoWidget.value = path; + previewVideo(node,path) + + } + } else { + alert(resp.status + " - " + resp.statusText); + } + } catch (error) { + alert(error); + } + } + + const fileInput = document.createElement("input"); + Object.assign(fileInput, { + type: "file", + accept: "video/webm,video/mp4,video/mkv,video/avi", + style: "display: none", + onchange: async () => { + if (fileInput.files.length) { + await uploadFile(fileInput.files[0], true); + } + }, + }); + document.body.append(fileInput); + + // Create the button widget for selecting the files + uploadWidget = node.addWidget("button", "choose video file to upload", "Video", () => { + fileInput.click(); + }); + + uploadWidget.serialize = false; + + previewVideo(node, videoWidget.value); + const cb = node.callback; + videoWidget.callback = function () { + previewVideo(node,videoWidget.value); + if (cb) { + return cb.apply(this, arguments); + } + }; + + return { widget: uploadWidget }; +} + +ComfyWidgets.VIDEOPLOAD = videoUpload; + +app.registerExtension({ + name: "IP_LAP.UploadVideo", + async beforeRegisterNodeDef(nodeType, nodeData, app) { + if (nodeData?.name == "LoadVideo") { + nodeData.input.required.upload = ["VIDEOPLOAD"]; + } + }, +}); +