4L>74C2`LUk3H!
z1Pm^~$kkjZSxC*%me!%q!8}io|1sBuEPASt)#`(;CT2lf9DxDw`Vvky9-jsY_kHF1
zK4jb|M!Mzuw?;u+uL#W6Ds2SR1BT@UkR`
zQs3WcFP%B<7mz%0FHeU%HpB*r?vXlZhKv%fw5)(~we^~K6>sd}3HsHQ%XOg5sytM6
z%Uh`G))oLFOM3@7D=_z}o5h=`tgYpqRd9%34hRXh_k|n#lF@uMOF2FB&FlL4zz;4{`p<6(^8$)79%9UAC$@TKh
zEO%RznsJ1hU^d>&EQuf_R%j3ZVSd1@azz9Y5Ly!65KR!l_uZTIk|nnn&Wb#F@5!5a
zZ=HM2cOKcVR(%b}+PSl>-_s%>@+L({dc
z71#yG37p7?%#Pdj0#Dc8(jqIeKh`2UHnvK;X5fS4L@qdP!6}2|MI~@b1*ZazACQnfR4>b?en#T|V;~pNaI_@~odeq`cM-
zAs?--=iYm9(u>o}*P3KcA6C&xE}#)>fgWjr5tV+awe-l0tdC298QGEZu@+d78+o|e
zztk)(aBgdLKlg9+dk?R#zZbvr!J&y42m_?a?HL^&8a(H5Cf@`}wN%fvU1QJSf>U$L
z%Jj_0%*@K{NRNysPNs!+YVJ9iv5%3Vn`wKV{Bs)SG}mt)Sb6z=EP8P#OxoLV?j)&b
zM{({B8d22lZRFOyb|a~mayuCgVj=h!b1w|ry>=RgxzX$84%fXF(Ml%iv4Uyn4UhLwr&V|7vQNheow##VD>5guUf8O6JPDS5d
zn(Wbqkr0~Pr_sT-jq?W1p)+xTY_`u`pw)Dac3wN
z(K|52l4O1)Ud2_B*T`{OeF33Fub`7qIedLVxAg7W)XIt>O)Bi@m|ob^e+QQf+t=RI
zetMQm%k;ga>2_Pk~O1m|_KQofz=Vx*<5B
z+{k)Kl-a&Y!flc=kZM1r5gvB?VF)~s@~u140t8Jl@Vf0tp+pI|K^%w^vC0>nCXo`C
zb|S^j8NPs1ihWLB#t){uG9#Unt(Q9?t(UK!ST9;Il9Ant&)_EhBIPN*PV!GAv$iUg
zeCEl
z13$yq(RT<^3K=flHFnG$3+P}ebV$w=Z3{^8#w?GtyFK?fLuKuKqP`$9Pbp*0qe=niS_jIj_>xqQU?Z@dlP96(o1wY4_rWP|)Et
z27|a4i7)dufkG^RTE3m?G6sJgg(1yR2Z8!C#H>Sgf%a=a6bPUf2k3O2C++y5}G1a~GLnH^<;bSbKtwBE~Zd+Uf
zPeyDbYol2!YvVu79U*faDGR8A=w7}mFIU!7CQz2E`oI?x&bOA+FC!xUEPx4o7ZmFHed5ctfK>j}49g9i9pvZDd&RYRa^vL%$0T
zw`1=(@JzE~+#N5okMK%6e&!tEm3JyB3^#LT@4zUl(tN8h-zxm4omI1PR%xji%iNj6
zT|9%k^Vxhhx2o+d$WaSfDf7pp7BR}r7BPw~tO%`IR?8MKeo2mB%G}Hwk6*^9Qnrjy
z%g6$kv&EJPU_gXUz{y>D(KtyFs#{EZ>b`gneDMp%)~di_Uj}Q_>_`m0P_o+$tXOkJsTb|4I4{y5*&%QXWuu0V`31sY!dgB
zy;a0L5nruNXI5Xep^w>F`^a_K=H^)~L`IfUJptgsUB;N*iP_;!rqi6K5~y^0KBbS5)4?
z;JB46~h!+(ejkIjt5i9jpq$M{}m7bjH+lbQ9u9>s9_7yOVKim)+n$?
zg@f`I_enY=%$P=tCVU9#%a7<5jrCVJr5dQvLS|4)(@TC?*7+N*Lgr7a)31LK>lFXS
zXPFQ_-%ftZ5w}3LSEi?U5sRQU{148FMR=YIiy*&=bLQY?@fd|DOE&?jLlNH`CW!D)
zi0c{oFlJ2t6Xd966vY9S2WE={i0WUVTjd=FDJ(a-s&K$NvZOG5&%?bED;G}~`uK?c
z%QE_}K!2*E=)VqYkw!s#kT!D2VCtrzOMf4p3&q1h2L-(PoXQZo95=BQm$q`RbAM7=
zn5bWFvAil+FB0#9!jjC`kmZ;sDo|mJuOksRkqzLtXCs&~$5{f(H3%da
z7_jAzTPj!}ucUaXw&$}njL-yuM`j^}a5L79CIvh*YIh|#)yv{GmXn*!Mk)qqO@#+C
zWRNB;1u5-XaLD^0bx(p(V78J0NOKTaExtpV_zy8a`Sl`6^XS)Q(dQWJ|38dQmO*Op
z&!hB$0`{0kA(GwQKcX%Kkq9=REm5QA1B4eO70d#k567>9Ep;;itI06JO9<@%hhFLuh>BimrsqT10V()+!|xkV
z>mOqV2r%zxxrWc
zE6H=c@1HS3=E*XhG3IpsP(xTK#5h3*Lf_~h_KeZt*XX=hA_J?YKPD0wNMiWS!2>#!
z(4zJ0NMcw`VCqp$oh=s$2M|W35Rvu7+9^q{e}wBxZutne%ycCy%QV9-(hTojyCVyB
z@1sUa1jmQQ(f;&Pf~Fc?|OW1>
z*H?@ceFfL;`O&t$EWBZss_`@iXH?SyaP|}tV(#kF5dc~U{!9qubs5XV;QfyEwI8)_W@f{3BsLRht
z!MwaoJ+P=4)8mh9<6Cfm9P(E4Ps)1P`0%VJl&IwUL}B01IXP5QD`>BqydP1dl)A
zEWj3Db$o9gU`swfBy}73xEacq9I}8R%PsB4uH+u+yxG8aUo{t(~b|t(~i_dFMTFnR_cpwEp@JXa32I
literal 0
HcmV?d00001
diff --git a/musetalk/utils/face_parsing/__pycache__/resnet.cpython-310.pyc b/musetalk/utils/face_parsing/__pycache__/resnet.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..fcb8507cecb5886749ac4059f904bab36da08bff
GIT binary patch
literal 3644
zcma)9TW=f36`q;BaJeKUO7g93ng~hJgl#NEa_XdMn!0wNBt=RT2-{!8-|pilYhD*_rR0^UYz_YWWP$
z_SX9H6Mts;APty2*EfT*j=a6xPzL}c-=ZQqU*~Q
zX-nq??>4?*(v`IrOx9v?*gUbiEsS}xjxGVbS_D$dim*!uY9ucECFZbvfS7#0U-
z91BDiEmJ+z_E7Gjfn(K<>%?7Xyf<0w{X+z;|(9BVf(R3c+ttEOuEVG`yYp+%Z%_Z}A9mK{si
zlQbCi!!*vdpN)#Dt;n`7nQEc>_n$Wpp4}`4!_9lyp!e+CuQsDmUSxwH#fkH|^5*?f
z9)A<|kAj`OFE3+|8obY;4?*cP;5$oFDGn;iCpjXrW-7rt0yZtOWGB8E+WoR{O
z*HH+>{O=G2D`12ZF%{$hg*AaOCf12HVS!yZQ+Fa}7!PU~5fgDD(Dp9q>xFaZOKj(q
zW;G;(Y|R=z#4q;+Je9n5cG-V*7h1uZsx>daS}(V9)k;Cd(>clVnP*R>>tiK4hP
zbjMN-_G(O|$cH-1+Q~=5SgAI8+6#gtO^P5eIBJJL-coc<-MAYT(ZOz}2HR3w57Vu3
zber1S+I|p!{jIiD+#hM5d^`^a!+xy&z4*~6PKzY$D`JF>t?fMS_h?^1bP%T6BhR5$
z={H>0gCOdMd5%qx3n{X&T7$TXlD`SD#Fu!}{Fb=O+uY_JYWtNZ{Quam#%t$x+nApx
zB0rx$1OXKjK;NM;VN-<9j30}FPl4hTHkz3}+BUV1N$XTCehjM~x{{5!dIx&fybF6w
z#C3LTLn9_O-mxe49(&1mJ2iC$U(tL={TQ{Pm{r$F+#vB2h>oLvic0gR+5-(FauoKp
zmmg%Rh(?9B08)v@0uo*8W$H<&bH_p&;$=qo;
zh=V};K>!jO^{L$qf=Bp#Ibx7k@6p~`m0_Z|-=7=DY-brJ2ZyXe%ZnJs7IJY#i&g
z(4lBu@Yfx53K#$`ZApG=?E^Lh5_s$e825Ab*XFS^v8MuaI2caQoEez-(NbX@0^q-q
zt6RGrUivhd8CV);c2t<|gUrI=FixfMV)X`%yNbL>y#)bgOGg9qu6Fz3vsmStCzqjK
zjZ_>KabP9{4~Z6T;pLo8(FW?*fVU#e|1b~H13meN5Od}+yo0m6@W49=V&Wa?NEfvu
zYtlmvZn9aohO^bR_n^8LUUOnPDSh7%&_B?)OhAV=(Fv`O8mphWpa-Y0rmkF2(V>Si
zo4$EUL{%?*ydjLHehNnp8b;$oW6M(0ga&u%1x*W@hHPHYEFlDzDJbaH#u-|>ine)n
zheVBpxpowhXz7a=d5I=;g!+I6$yd&vlK`o*gyFqTTdmR%AzFGT@snW}6K0Bd!o--a
z-w&VuG0Xa3tlA)aIz-c$sWSCeRc}}It*ZV|Hw+6tiIe?<0$VYm>6=>`S5a}G*69LpYI2gAp
z*jG%g83Q@AzZ1vyv4h0NehiyGYE3L)V^5p}^Fb0VVBio2l_QE>EqW1{&oRlDwkZ|U
zeqK;?%OonalN9kF*Y0R2q0`GWAvhlgcaSrv-;?+fqPtAj*@YRJA8G#hG9Y&_^#MxW
zgaA1ks&2zm##a_#T=2PwlD_#vU@QQ6jLtrOLL8Tu5;H*D1Q3>CPP-6?7K+_NlA+$4
zZ~$N*x<nz^*z1%RZj8IbDWTYAgG{M!>z)WNs@4D9<6I~I@Xlx+0G+;tkdMHhpt
zt_x&-2{JPPlp@D4iGurelE$Gb19B7F)ovL`CgCvA3EY#s&_2>StW|(?=l5UziKJvN
zgb5mHcHXh+quMR+s7v5#e?La@Y2NDMdnO9W=~j^=_$cCe6Fu?Vh+ro2(|(>Tkt
zV-&i+&~!c1#8|hfs@o(OseYN}x%bW0Z*EP)k0rMfDB8k55=}R$Obr33cS-9ERgxA+
zU`=A7umvzmqzZQHk2UH`4sn)g0sa1;;pYjXQv{{@sD
BKdt}(
literal 0
HcmV?d00001
diff --git a/musetalk/utils/face_parsing/model.py b/musetalk/utils/face_parsing/model.py
new file mode 100755
index 0000000..13d13e0
--- /dev/null
+++ b/musetalk/utils/face_parsing/model.py
@@ -0,0 +1,283 @@
+#!/usr/bin/python
+# -*- encoding: utf-8 -*-
+
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torchvision
+
+from .resnet import Resnet18
+# from modules.bn import InPlaceABNSync as BatchNorm2d
+
+
+class ConvBNReLU(nn.Module):
+ def __init__(self, in_chan, out_chan, ks=3, stride=1, padding=1, *args, **kwargs):
+ super(ConvBNReLU, self).__init__()
+ self.conv = nn.Conv2d(in_chan,
+ out_chan,
+ kernel_size = ks,
+ stride = stride,
+ padding = padding,
+ bias = False)
+ self.bn = nn.BatchNorm2d(out_chan)
+ self.init_weight()
+
+ def forward(self, x):
+ x = self.conv(x)
+ x = F.relu(self.bn(x))
+ return x
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+class BiSeNetOutput(nn.Module):
+ def __init__(self, in_chan, mid_chan, n_classes, *args, **kwargs):
+ super(BiSeNetOutput, self).__init__()
+ self.conv = ConvBNReLU(in_chan, mid_chan, ks=3, stride=1, padding=1)
+ self.conv_out = nn.Conv2d(mid_chan, n_classes, kernel_size=1, bias=False)
+ self.init_weight()
+
+ def forward(self, x):
+ x = self.conv(x)
+ x = self.conv_out(x)
+ return x
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+ def get_params(self):
+ wd_params, nowd_params = [], []
+ for name, module in self.named_modules():
+ if isinstance(module, nn.Linear) or isinstance(module, nn.Conv2d):
+ wd_params.append(module.weight)
+ if not module.bias is None:
+ nowd_params.append(module.bias)
+ elif isinstance(module, nn.BatchNorm2d):
+ nowd_params += list(module.parameters())
+ return wd_params, nowd_params
+
+
+class AttentionRefinementModule(nn.Module):
+ def __init__(self, in_chan, out_chan, *args, **kwargs):
+ super(AttentionRefinementModule, self).__init__()
+ self.conv = ConvBNReLU(in_chan, out_chan, ks=3, stride=1, padding=1)
+ self.conv_atten = nn.Conv2d(out_chan, out_chan, kernel_size= 1, bias=False)
+ self.bn_atten = nn.BatchNorm2d(out_chan)
+ self.sigmoid_atten = nn.Sigmoid()
+ self.init_weight()
+
+ def forward(self, x):
+ feat = self.conv(x)
+ atten = F.avg_pool2d(feat, feat.size()[2:])
+ atten = self.conv_atten(atten)
+ atten = self.bn_atten(atten)
+ atten = self.sigmoid_atten(atten)
+ out = torch.mul(feat, atten)
+ return out
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+
+class ContextPath(nn.Module):
+ def __init__(self, resnet_path, *args, **kwargs):
+ super(ContextPath, self).__init__()
+ self.resnet = Resnet18(resnet_path)
+ self.arm16 = AttentionRefinementModule(256, 128)
+ self.arm32 = AttentionRefinementModule(512, 128)
+ self.conv_head32 = ConvBNReLU(128, 128, ks=3, stride=1, padding=1)
+ self.conv_head16 = ConvBNReLU(128, 128, ks=3, stride=1, padding=1)
+ self.conv_avg = ConvBNReLU(512, 128, ks=1, stride=1, padding=0)
+
+ self.init_weight()
+
+ def forward(self, x):
+ H0, W0 = x.size()[2:]
+ feat8, feat16, feat32 = self.resnet(x)
+ H8, W8 = feat8.size()[2:]
+ H16, W16 = feat16.size()[2:]
+ H32, W32 = feat32.size()[2:]
+
+ avg = F.avg_pool2d(feat32, feat32.size()[2:])
+ avg = self.conv_avg(avg)
+ avg_up = F.interpolate(avg, (H32, W32), mode='nearest')
+
+ feat32_arm = self.arm32(feat32)
+ feat32_sum = feat32_arm + avg_up
+ feat32_up = F.interpolate(feat32_sum, (H16, W16), mode='nearest')
+ feat32_up = self.conv_head32(feat32_up)
+
+ feat16_arm = self.arm16(feat16)
+ feat16_sum = feat16_arm + feat32_up
+ feat16_up = F.interpolate(feat16_sum, (H8, W8), mode='nearest')
+ feat16_up = self.conv_head16(feat16_up)
+
+ return feat8, feat16_up, feat32_up # x8, x8, x16
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+ def get_params(self):
+ wd_params, nowd_params = [], []
+ for name, module in self.named_modules():
+ if isinstance(module, (nn.Linear, nn.Conv2d)):
+ wd_params.append(module.weight)
+ if not module.bias is None:
+ nowd_params.append(module.bias)
+ elif isinstance(module, nn.BatchNorm2d):
+ nowd_params += list(module.parameters())
+ return wd_params, nowd_params
+
+
+### This is not used, since I replace this with the resnet feature with the same size
+class SpatialPath(nn.Module):
+ def __init__(self, *args, **kwargs):
+ super(SpatialPath, self).__init__()
+ self.conv1 = ConvBNReLU(3, 64, ks=7, stride=2, padding=3)
+ self.conv2 = ConvBNReLU(64, 64, ks=3, stride=2, padding=1)
+ self.conv3 = ConvBNReLU(64, 64, ks=3, stride=2, padding=1)
+ self.conv_out = ConvBNReLU(64, 128, ks=1, stride=1, padding=0)
+ self.init_weight()
+
+ def forward(self, x):
+ feat = self.conv1(x)
+ feat = self.conv2(feat)
+ feat = self.conv3(feat)
+ feat = self.conv_out(feat)
+ return feat
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+ def get_params(self):
+ wd_params, nowd_params = [], []
+ for name, module in self.named_modules():
+ if isinstance(module, nn.Linear) or isinstance(module, nn.Conv2d):
+ wd_params.append(module.weight)
+ if not module.bias is None:
+ nowd_params.append(module.bias)
+ elif isinstance(module, nn.BatchNorm2d):
+ nowd_params += list(module.parameters())
+ return wd_params, nowd_params
+
+
+class FeatureFusionModule(nn.Module):
+ def __init__(self, in_chan, out_chan, *args, **kwargs):
+ super(FeatureFusionModule, self).__init__()
+ self.convblk = ConvBNReLU(in_chan, out_chan, ks=1, stride=1, padding=0)
+ self.conv1 = nn.Conv2d(out_chan,
+ out_chan//4,
+ kernel_size = 1,
+ stride = 1,
+ padding = 0,
+ bias = False)
+ self.conv2 = nn.Conv2d(out_chan//4,
+ out_chan,
+ kernel_size = 1,
+ stride = 1,
+ padding = 0,
+ bias = False)
+ self.relu = nn.ReLU(inplace=True)
+ self.sigmoid = nn.Sigmoid()
+ self.init_weight()
+
+ def forward(self, fsp, fcp):
+ fcat = torch.cat([fsp, fcp], dim=1)
+ feat = self.convblk(fcat)
+ atten = F.avg_pool2d(feat, feat.size()[2:])
+ atten = self.conv1(atten)
+ atten = self.relu(atten)
+ atten = self.conv2(atten)
+ atten = self.sigmoid(atten)
+ feat_atten = torch.mul(feat, atten)
+ feat_out = feat_atten + feat
+ return feat_out
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+ def get_params(self):
+ wd_params, nowd_params = [], []
+ for name, module in self.named_modules():
+ if isinstance(module, nn.Linear) or isinstance(module, nn.Conv2d):
+ wd_params.append(module.weight)
+ if not module.bias is None:
+ nowd_params.append(module.bias)
+ elif isinstance(module, nn.BatchNorm2d):
+ nowd_params += list(module.parameters())
+ return wd_params, nowd_params
+
+
+class BiSeNet(nn.Module):
+ def __init__(self, resnet_path='models/resnet18-5c106cde.pth', n_classes=19, *args, **kwargs):
+ super(BiSeNet, self).__init__()
+ self.cp = ContextPath(resnet_path)
+ ## here self.sp is deleted
+ self.ffm = FeatureFusionModule(256, 256)
+ self.conv_out = BiSeNetOutput(256, 256, n_classes)
+ self.conv_out16 = BiSeNetOutput(128, 64, n_classes)
+ self.conv_out32 = BiSeNetOutput(128, 64, n_classes)
+ self.init_weight()
+
+ def forward(self, x):
+ H, W = x.size()[2:]
+ feat_res8, feat_cp8, feat_cp16 = self.cp(x) # here return res3b1 feature
+ feat_sp = feat_res8 # use res3b1 feature to replace spatial path feature
+ feat_fuse = self.ffm(feat_sp, feat_cp8)
+
+ feat_out = self.conv_out(feat_fuse)
+ feat_out16 = self.conv_out16(feat_cp8)
+ feat_out32 = self.conv_out32(feat_cp16)
+
+ feat_out = F.interpolate(feat_out, (H, W), mode='bilinear', align_corners=True)
+ feat_out16 = F.interpolate(feat_out16, (H, W), mode='bilinear', align_corners=True)
+ feat_out32 = F.interpolate(feat_out32, (H, W), mode='bilinear', align_corners=True)
+ return feat_out, feat_out16, feat_out32
+
+ def init_weight(self):
+ for ly in self.children():
+ if isinstance(ly, nn.Conv2d):
+ nn.init.kaiming_normal_(ly.weight, a=1)
+ if not ly.bias is None: nn.init.constant_(ly.bias, 0)
+
+ def get_params(self):
+ wd_params, nowd_params, lr_mul_wd_params, lr_mul_nowd_params = [], [], [], []
+ for name, child in self.named_children():
+ child_wd_params, child_nowd_params = child.get_params()
+ if isinstance(child, FeatureFusionModule) or isinstance(child, BiSeNetOutput):
+ lr_mul_wd_params += child_wd_params
+ lr_mul_nowd_params += child_nowd_params
+ else:
+ wd_params += child_wd_params
+ nowd_params += child_nowd_params
+ return wd_params, nowd_params, lr_mul_wd_params, lr_mul_nowd_params
+
+
+if __name__ == "__main__":
+ net = BiSeNet(19)
+ net.cuda()
+ net.eval()
+ in_ten = torch.randn(16, 3, 640, 480).cuda()
+ out, out16, out32 = net(in_ten)
+ print(out.shape)
+
+ net.get_params()
diff --git a/musetalk/utils/face_parsing/resnet.py b/musetalk/utils/face_parsing/resnet.py
new file mode 100755
index 0000000..e2e5d87
--- /dev/null
+++ b/musetalk/utils/face_parsing/resnet.py
@@ -0,0 +1,109 @@
+#!/usr/bin/python
+# -*- encoding: utf-8 -*-
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import torch.utils.model_zoo as modelzoo
+
+# from modules.bn import InPlaceABNSync as BatchNorm2d
+
+resnet18_url = 'https://download.pytorch.org/models/resnet18-5c106cde.pth'
+
+
+def conv3x3(in_planes, out_planes, stride=1):
+ """3x3 convolution with padding"""
+ return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride,
+ padding=1, bias=False)
+
+
+class BasicBlock(nn.Module):
+ def __init__(self, in_chan, out_chan, stride=1):
+ super(BasicBlock, self).__init__()
+ self.conv1 = conv3x3(in_chan, out_chan, stride)
+ self.bn1 = nn.BatchNorm2d(out_chan)
+ self.conv2 = conv3x3(out_chan, out_chan)
+ self.bn2 = nn.BatchNorm2d(out_chan)
+ self.relu = nn.ReLU(inplace=True)
+ self.downsample = None
+ if in_chan != out_chan or stride != 1:
+ self.downsample = nn.Sequential(
+ nn.Conv2d(in_chan, out_chan,
+ kernel_size=1, stride=stride, bias=False),
+ nn.BatchNorm2d(out_chan),
+ )
+
+ def forward(self, x):
+ residual = self.conv1(x)
+ residual = F.relu(self.bn1(residual))
+ residual = self.conv2(residual)
+ residual = self.bn2(residual)
+
+ shortcut = x
+ if self.downsample is not None:
+ shortcut = self.downsample(x)
+
+ out = shortcut + residual
+ out = self.relu(out)
+ return out
+
+
+def create_layer_basic(in_chan, out_chan, bnum, stride=1):
+ layers = [BasicBlock(in_chan, out_chan, stride=stride)]
+ for i in range(bnum-1):
+ layers.append(BasicBlock(out_chan, out_chan, stride=1))
+ return nn.Sequential(*layers)
+
+
+class Resnet18(nn.Module):
+ def __init__(self, model_path):
+ super(Resnet18, self).__init__()
+ self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3,
+ bias=False)
+ self.bn1 = nn.BatchNorm2d(64)
+ self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
+ self.layer1 = create_layer_basic(64, 64, bnum=2, stride=1)
+ self.layer2 = create_layer_basic(64, 128, bnum=2, stride=2)
+ self.layer3 = create_layer_basic(128, 256, bnum=2, stride=2)
+ self.layer4 = create_layer_basic(256, 512, bnum=2, stride=2)
+ self.init_weight(model_path)
+
+ def forward(self, x):
+ x = self.conv1(x)
+ x = F.relu(self.bn1(x))
+ x = self.maxpool(x)
+
+ x = self.layer1(x)
+ feat8 = self.layer2(x) # 1/8
+ feat16 = self.layer3(feat8) # 1/16
+ feat32 = self.layer4(feat16) # 1/32
+ return feat8, feat16, feat32
+
+ def init_weight(self, model_path):
+ state_dict = torch.load(model_path) #modelzoo.load_url(resnet18_url)
+ self_state_dict = self.state_dict()
+ for k, v in state_dict.items():
+ if 'fc' in k: continue
+ self_state_dict.update({k: v})
+ self.load_state_dict(self_state_dict)
+
+ def get_params(self):
+ wd_params, nowd_params = [], []
+ for name, module in self.named_modules():
+ if isinstance(module, (nn.Linear, nn.Conv2d)):
+ wd_params.append(module.weight)
+ if not module.bias is None:
+ nowd_params.append(module.bias)
+ elif isinstance(module, nn.BatchNorm2d):
+ nowd_params += list(module.parameters())
+ return wd_params, nowd_params
+
+
+if __name__ == "__main__":
+ net = Resnet18()
+ x = torch.randn(16, 3, 224, 224)
+ out = net(x)
+ print(out[0].size())
+ print(out[1].size())
+ print(out[2].size())
+ net.get_params()
diff --git a/musetalk/utils/preprocessing.py b/musetalk/utils/preprocessing.py
new file mode 100644
index 0000000..5c92acc
--- /dev/null
+++ b/musetalk/utils/preprocessing.py
@@ -0,0 +1,157 @@
+import sys
+from face_detection import FaceAlignment,LandmarksType
+from os import listdir, path
+import subprocess
+import numpy as np
+import cv2
+import pickle
+import os
+import json
+from mmpose.apis import inference_topdown, init_model
+from mmpose.structures import merge_data_samples
+import torch
+from tqdm import tqdm
+parent_directory = os.path.dirname(os.path.abspath(__file__))
+# initialize the mmpose model
+device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+'''
+config_file = os.path.join(parent_directory,"dwpose/rtmpose-l_8xb32-270e_coco-ubody-wholebody-384x288.py")
+checkpoint_file = './models/dwpose/dw-ll_ucoco_384.pth'
+model = init_model(config_file, checkpoint_file, device=device)
+'''
+
+
+# initialize the face detection model
+device = "cuda" if torch.cuda.is_available() else "cpu"
+fa = FaceAlignment(LandmarksType._2D, flip_input=False,device=device)
+
+# maker if the bbox is not sufficient
+coord_placeholder = (0.0,0.0,0.0,0.0)
+
+def resize_landmark(landmark, w, h, new_w, new_h):
+ w_ratio = new_w / w
+ h_ratio = new_h / h
+ landmark_norm = landmark / [w, h]
+ landmark_resized = landmark_norm * [new_w, new_h]
+ return landmark_resized
+
+def read_imgs(img_list):
+ frames = []
+ print('reading images...')
+ for img_path in tqdm(img_list):
+ frame = cv2.imread(img_path)
+ frames.append(frame)
+ return frames
+
+def get_bbox_range(model,img_list,batch_size_fa,upperbondrange =0):
+ frames = read_imgs(img_list)
+ # batch_size_fa = 1
+ batches = [frames[i:i + batch_size_fa] for i in range(0, len(frames), batch_size_fa)]
+ coords_list = []
+ landmarks = []
+ if upperbondrange != 0:
+ print('get key_landmark and face bounding boxes with the bbox_shift:',upperbondrange)
+ else:
+ print('get key_landmark and face bounding boxes with the default value')
+ average_range_minus = []
+ average_range_plus = []
+ for fb in tqdm(batches):
+ results = inference_topdown(model, np.asarray(fb)[0])
+ results = merge_data_samples(results)
+ keypoints = results.pred_instances.keypoints
+ face_land_mark= keypoints[0][23:91]
+ face_land_mark = face_land_mark.astype(np.int32)
+
+ # get bounding boxes by face detetion
+ bbox = fa.get_detections_for_batch(np.asarray(fb))
+
+ # adjust the bounding box refer to landmark
+ # Add the bounding box to a tuple and append it to the coordinates list
+ for j, f in enumerate(bbox):
+ if f is None: # no face in the image
+ coords_list += [coord_placeholder]
+ continue
+
+ half_face_coord = face_land_mark[29]#np.mean([face_land_mark[28], face_land_mark[29]], axis=0)
+ range_minus = (face_land_mark[30]- face_land_mark[29])[1]
+ range_plus = (face_land_mark[29]- face_land_mark[28])[1]
+ average_range_minus.append(range_minus)
+ average_range_plus.append(range_plus)
+ if upperbondrange != 0:
+ half_face_coord[1] = upperbondrange+half_face_coord[1] #手动调整 + 向下(偏29) - 向上(偏28)
+
+ text_range=f"Total frame:「{len(frames)}」 Manually adjust range : [ -{int(sum(average_range_minus) / len(average_range_minus))}~{int(sum(average_range_plus) / len(average_range_plus))} ] , the current value: {upperbondrange}"
+ return text_range
+
+
+def get_landmark_and_bbox(model,img_list,batch_size_fa,upperbondrange =0):
+ frames = read_imgs(img_list)
+ # batch_size_fa = 1
+ batches = [frames[i:i + batch_size_fa] for i in range(0, len(frames), batch_size_fa)]
+ coords_list = []
+ landmarks = []
+ if upperbondrange != 0:
+ print('get key_landmark and face bounding boxes with the bbox_shift:',upperbondrange)
+ else:
+ print('get key_landmark and face bounding boxes with the default value')
+ average_range_minus = []
+ average_range_plus = []
+ for fb in tqdm(batches):
+ results = inference_topdown(model, np.asarray(fb)[0])
+ results = merge_data_samples(results)
+ keypoints = results.pred_instances.keypoints
+ face_land_mark= keypoints[0][23:91]
+ face_land_mark = face_land_mark.astype(np.int32)
+
+ # get bounding boxes by face detetion
+ bbox = fa.get_detections_for_batch(np.asarray(fb))
+
+ # adjust the bounding box refer to landmark
+ # Add the bounding box to a tuple and append it to the coordinates list
+ for j, f in enumerate(bbox):
+ if f is None: # no face in the image
+ coords_list += [coord_placeholder]
+ continue
+
+ half_face_coord = face_land_mark[29]#np.mean([face_land_mark[28], face_land_mark[29]], axis=0)
+ range_minus = (face_land_mark[30]- face_land_mark[29])[1]
+ range_plus = (face_land_mark[29]- face_land_mark[28])[1]
+ average_range_minus.append(range_minus)
+ average_range_plus.append(range_plus)
+ if upperbondrange != 0:
+ half_face_coord[1] = upperbondrange+half_face_coord[1] #手动调整 + 向下(偏29) - 向上(偏28)
+ half_face_dist = np.max(face_land_mark[:,1]) - half_face_coord[1]
+ upper_bond = half_face_coord[1]-half_face_dist
+
+ f_landmark = (np.min(face_land_mark[:, 0]),int(upper_bond),np.max(face_land_mark[:, 0]),np.max(face_land_mark[:,1]))
+ x1, y1, x2, y2 = f_landmark
+
+ if y2-y1<=0 or x2-x1<=0 or x1<0: # if the landmark bbox is not suitable, reuse the bbox
+ coords_list += [f]
+ w,h = f[2]-f[0], f[3]-f[1]
+ print("error bbox:",f)
+ else:
+ coords_list += [f_landmark]
+
+ print("********************************************bbox_shift parameter adjustment**********************************************************")
+ print(f"Total frame:「{len(frames)}」 Manually adjust range : [ -{int(sum(average_range_minus) / len(average_range_minus))}~{int(sum(average_range_plus) / len(average_range_plus))} ] , the current value: {upperbondrange}")
+ print("*************************************************************************************************************************************")
+ return coords_list,frames
+
+
+if __name__ == "__main__":
+ img_list = ["./results/lyria/00000.png","./results/lyria/00001.png","./results/lyria/00002.png","./results/lyria/00003.png"]
+ crop_coord_path = "./coord_face.pkl"
+ coords_list,full_frames = get_landmark_and_bbox(img_list)
+ with open(crop_coord_path, 'wb') as f:
+ pickle.dump(coords_list, f)
+
+ for bbox, frame in zip(coords_list,full_frames):
+ if bbox == coord_placeholder:
+ continue
+ x1, y1, x2, y2 = bbox
+ crop_frame = frame[y1:y2, x1:x2]
+ print('Cropped shape', crop_frame.shape)
+
+ #cv2.imwrite(path.join(save_dir, '{}.png'.format(i)),full_frames[i][0][y1:y2, x1:x2])
+ print(coords_list)
diff --git a/musetalk/utils/utils.py b/musetalk/utils/utils.py
new file mode 100644
index 0000000..1cc3706
--- /dev/null
+++ b/musetalk/utils/utils.py
@@ -0,0 +1,61 @@
+import os
+import cv2
+import numpy as np
+import torch
+
+ffmpeg_path = os.getenv('FFMPEG_PATH')
+if ffmpeg_path is None:
+ print("please download ffmpeg-static and export to FFMPEG_PATH. \nFor example: export FFMPEG_PATH=/musetalk/ffmpeg-4.4-amd64-static")
+elif ffmpeg_path not in os.getenv('PATH'):
+ print("add ffmpeg to path")
+ os.environ["PATH"] = f"{ffmpeg_path}:{os.environ['PATH']}"
+
+
+from ..whisper.audio2feature import Audio2Feature
+from ..models.vae import VAE
+from ..models.unet import UNet,PositionalEncoding
+
+def load_all_model(base_dir):
+ audio_processor = Audio2Feature(model_path=os.path.join(base_dir,"whisper/tiny.pt"))
+ vae = VAE(model_path =os.path.join(base_dir,"sd-vae-ft-mse/"))
+ unet = UNet(unet_config=os.path.join(base_dir,"musetalk/musetalk.json"),
+ model_path =os.path.join(base_dir,"musetalk/pytorch_model.bin"))
+ pe = PositionalEncoding(d_model=384)
+ return audio_processor,vae,unet,pe
+
+def get_file_type(video_path):
+ _, ext = os.path.splitext(video_path)
+
+ if ext.lower() in ['.jpg', '.jpeg', '.png', '.bmp', '.tif', '.tiff']:
+ return 'image'
+ elif ext.lower() in ['.avi', '.mp4', '.mov', '.flv', '.mkv']:
+ return 'video'
+ else:
+ return 'unsupported'
+
+def get_video_fps(video_path):
+ video = cv2.VideoCapture(video_path)
+ fps = video.get(cv2.CAP_PROP_FPS)
+ video.release()
+ return fps
+
+def datagen(whisper_chunks,vae_encode_latents,batch_size=8,delay_frame = 0):
+ whisper_batch, latent_batch = [], []
+ for i, w in enumerate(whisper_chunks):
+ idx = (i+delay_frame)%len(vae_encode_latents)
+ latent = vae_encode_latents[idx]
+ whisper_batch.append(w)
+ latent_batch.append(latent)
+
+ if len(latent_batch) >= batch_size:
+ whisper_batch = np.asarray(whisper_batch)
+ latent_batch = torch.cat(latent_batch, dim=0)
+ yield whisper_batch, latent_batch
+ whisper_batch, latent_batch = [], []
+
+ # the last batch may smaller than batch size
+ if len(latent_batch) > 0:
+ whisper_batch = np.asarray(whisper_batch)
+ latent_batch = torch.cat(latent_batch, dim=0)
+
+ yield whisper_batch, latent_batch
\ No newline at end of file
diff --git a/musetalk/whisper/__pycache__/audio2feature.cpython-310.pyc b/musetalk/whisper/__pycache__/audio2feature.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..8e3f393121cfc4af810fd9531ad03525abeb9806
GIT binary patch
literal 3543
zcmai0TW=f36`q;BaCb>cluYYlsVVEaO_|1&>n1g76m5zGPV-=)fZIMSnl5P0P@+sO
zso7Oxu}h$UoVU7vAP68oroW`WVP5-C=ntfT9@>6qmXswqKu6lcGiUDSeCIN**CU2@
z|Mpw`zpXO%PwHH56?8tpmz{%1Ci#%XObY1?`AEcqmog`I?U)z)_O}vO>~9bUDmVyF
zBk8{4oPErsC;b;p`X?f;s+z1k6Z+4pDgzn5$7U&JNd9eB=`^u+1Fkd{Nz1P*2YMk(AuAWN%gdIL6+HRiV{cQJ)
zgSNSehU2bGMq{alQ_<=4Xz+4#(D?vg_6HEfVvcKzkcUDxq$AxIJa(W8Z%;_SFJwgq
zkhqs5p^PBG44TvDvMOun`?4;VUN9(k8CO~{^(>u!Y@hC~5y=)=8
z`}inRzv~X4C6Au`W@m&>j?Tpsy7r;{@}Zr{u?do7kPh-BA^2u4i0>@%>Dm>IbQY7&
z=&5I$AUzLc(duAmsa3+?!Q@cG37>g;Z06&ukUAu9vcla0v}QpO6pj>|Ogc@r!DgZD
zyXd>?R<_P&Q7($Gh)%>u?9=U8HOKs-dLqENg(uw&R`#Io>L#1j3jfqAYUCV61zeT5
z!?$z)pi=m7Gw-FC`Ir+(UshgOzNqvCOcR{{BWlgt$0~1S!$D8U)?7td>uEPbcbvAm
zt-Zl#Ds2r?sh(dfc{u6n?r1SR>1y3Qe%QJ`oaJ4eCjH8*be`5-&{FWT8(FxU=_gT0J1
zNOR-qZn~$8GwMDE8KlPVP9`dqMx+x{?Tyo3H&-b>Q_(8h?@m;^u0!l=>hl|yPd1g$
z29h3VR^GU>ar%?Y_~rdg0~lAmTuJ*Xa<|%r1nc9dr_x+$yH_w&eXF)O)Puc!%i`WR
zg+I4GQIZA`x?hNr_8pQ#S}7ugJ1+f~4(0*!MKpJdpa`#WhfUg#
z&Lq0dzJN=Xd&RLrP5Lk4NJ!KxAiE_pks+C;XsK_4w5xiJ
z_zfcKL@2`OTST@&OeIyH6Nrtk)daqvw`k})M992YPmt8#CDJ1DJtDV>e4of2B0m5z
zLgxC1)cX+;GUX)*ZN6#wB!hov>c^Ovy#rDrXu^GTI^K-<K7iR>{+E?2LdMb3(SFO
zkgjI|qN{uA6oECNFTDqd$wdXQLWBm?-U|_xk)BuDmEC`p*ZAZEYqj$bKCQ-xQ?2uj
z)}!A(dADVKy9Ea%;rsU>D}#*tZMtuJ#+ztFnbM|hY+f3UjNOhd+00ssY-W}gIkLBZ
zls?PM+M{U7B1zxnOAtH
zDEi>CDBxbXw(m#>D?C`lm6c7hqi^}Z@?K!7|-g|ILH8+e6$_t06%e*
zBTr;KJ$R~2X!*%FLr%yJkCd9K(!ka!$zQE!&JAoLiy2$Uko#t^VI0|OB*Q_Lep){dSgi_Yq$OeD*UbKTVCIrqt4pJJnof=Nxd))&b`YxuMAABLZ9
zTs5-2qFuV%7Mdy>O-|Sjs60x$Bb6j3N~mr-8WLYilEb6!upH4;{b-BsyTm;pLYTO8
zLmOx)O=)s=3RW2f$AAeTc=VrDe|4*IyVkIybmpi#NRwoGr_ACtN5VJhTfgbl%mW{Rk
z6%CP|jAyGRP1Oq!NJT~&yWgW3?-Q{}f$&l4`3`Xu&~=I*ecF*#y5Rt@TSJXqM4V
j0i&>7{!!`=dJ&*Nq1#3S6bBRyZt-o*sX->r-EaN}LY7dA
literal 0
HcmV?d00001
diff --git a/musetalk/whisper/audio2feature.py b/musetalk/whisper/audio2feature.py
new file mode 100644
index 0000000..908db3d
--- /dev/null
+++ b/musetalk/whisper/audio2feature.py
@@ -0,0 +1,124 @@
+import os
+from .whisper import load_model
+import soundfile as sf
+import numpy as np
+import time
+import sys
+sys.path.append("..")
+
+class Audio2Feature():
+ def __init__(self,
+ whisper_model_type="tiny",
+ model_path="./models/whisper/tiny.pt"):
+ self.whisper_model_type = whisper_model_type
+ self.model = load_model(model_path) #
+
+ def get_sliced_feature(self,feature_array, vid_idx, audio_feat_length= [2,2],fps = 25):
+ """
+ Get sliced features based on a given index
+ :param feature_array:
+ :param start_idx: the start index of the feature
+ :param audio_feat_length:
+ :return:
+ """
+ length = len(feature_array)
+ selected_feature = []
+ selected_idx = []
+
+ center_idx = int(vid_idx*50/fps)
+ left_idx = center_idx-audio_feat_length[0]*2
+ right_idx = center_idx + (audio_feat_length[1]+1)*2
+
+ for idx in range(left_idx,right_idx):
+ idx = max(0, idx)
+ idx = min(length-1, idx)
+ x = feature_array[idx]
+ selected_feature.append(x)
+ selected_idx.append(idx)
+
+ selected_feature = np.concatenate(selected_feature, axis=0)
+ selected_feature = selected_feature.reshape(-1, 384)# 50*384
+ return selected_feature,selected_idx
+
+ def get_sliced_feature_sparse(self,feature_array, vid_idx, audio_feat_length= [2,2],fps = 25):
+ """
+ Get sliced features based on a given index
+ :param feature_array:
+ :param start_idx: the start index of the feature
+ :param audio_feat_length:
+ :return:
+ """
+ length = len(feature_array)
+ selected_feature = []
+ selected_idx = []
+
+ for dt in range(-audio_feat_length[0],audio_feat_length[1]+1):
+ left_idx = int((vid_idx+dt)*50/fps)
+ if left_idx<1 or left_idx>length-1:
+ left_idx = max(0, left_idx)
+ left_idx = min(length-1, left_idx)
+
+ x = feature_array[left_idx]
+ x = x[np.newaxis,:,:]
+ x = np.repeat(x, 2, axis=0)
+ selected_feature.append(x)
+ selected_idx.append(left_idx)
+ selected_idx.append(left_idx)
+ else:
+ x = feature_array[left_idx-1:left_idx+1]
+ selected_feature.append(x)
+ selected_idx.append(left_idx-1)
+ selected_idx.append(left_idx)
+ selected_feature = np.concatenate(selected_feature, axis=0)
+ selected_feature = selected_feature.reshape(-1, 384)# 50*384
+ return selected_feature,selected_idx
+
+
+ def feature2chunks(self,feature_array,fps,audio_feat_length = [2,2]):
+ whisper_chunks = []
+ whisper_idx_multiplier = 50./fps
+ i = 0
+ print(f"video in {fps} FPS, audio idx in 50FPS")
+ while 1:
+ start_idx = int(i * whisper_idx_multiplier)
+ selected_feature,selected_idx = self.get_sliced_feature(feature_array= feature_array,vid_idx = i,audio_feat_length=audio_feat_length,fps=fps)
+ #print(f"i:{i},selected_idx {selected_idx}")
+ whisper_chunks.append(selected_feature)
+ i += 1
+ if start_idx>len(feature_array):
+ break
+
+ return whisper_chunks
+
+ def audio2feat(self,audio_path):
+ # get the sample rate of the audio
+ result = self.model.transcribe(audio_path)
+ embed_list = []
+ for emb in result['segments']:
+ encoder_embeddings = emb['encoder_embeddings']
+ encoder_embeddings = encoder_embeddings.transpose(0,2,1,3)
+ encoder_embeddings = encoder_embeddings.squeeze(0)
+ start_idx = int(emb['start'])
+ end_idx = int(emb['end'])
+ emb_end_idx = int((end_idx - start_idx)/2)
+ embed_list.append(encoder_embeddings[:emb_end_idx])
+ concatenated_array = np.concatenate(embed_list, axis=0)
+ return concatenated_array
+
+if __name__ == "__main__":
+ audio_processor = Audio2Feature(model_path="../../models/whisper/whisper_tiny.pt")
+ audio_path = "./test.mp3"
+ array = audio_processor.audio2feat(audio_path)
+ print(array.shape)
+ fps = 25
+ whisper_idx_multiplier = 50./fps
+
+ i = 0
+ print(f"video in {fps} FPS, audio idx in 50FPS")
+ while 1:
+ start_idx = int(i * whisper_idx_multiplier)
+ selected_feature,selected_idx = audio_processor.get_sliced_feature(feature_array= array,vid_idx = i,audio_feat_length=[2,2],fps=fps)
+ print(f"video idx {i},\t audio idx {selected_idx},\t shape {selected_feature.shape}")
+ i += 1
+ if start_idx>len(array):
+ break
diff --git a/musetalk/whisper/whisper/__init__.py b/musetalk/whisper/whisper/__init__.py
new file mode 100644
index 0000000..b925553
--- /dev/null
+++ b/musetalk/whisper/whisper/__init__.py
@@ -0,0 +1,116 @@
+import hashlib
+import io
+import os
+import urllib
+import warnings
+from typing import List, Optional, Union
+
+import torch
+from tqdm import tqdm
+
+from .audio import load_audio, log_mel_spectrogram, pad_or_trim
+from .decoding import DecodingOptions, DecodingResult, decode, detect_language
+from .model import Whisper, ModelDimensions
+from .transcribe import transcribe
+
+
+_MODELS = {
+ "tiny.en": "https://openaipublic.azureedge.net/main/whisper/models/d3dd57d32accea0b295c96e26691aa14d8822fac7d9d27d5dc00b4ca2826dd03/tiny.en.pt",
+ "tiny": "https://openaipublic.azureedge.net/main/whisper/models/65147644a518d12f04e32d6f3b26facc3f8dd46e5390956a9424a650c0ce22b9/tiny.pt",
+ "base.en": "https://openaipublic.azureedge.net/main/whisper/models/25a8566e1d0c1e2231d1c762132cd20e0f96a85d16145c3a00adf5d1ac670ead/base.en.pt",
+ "base": "https://openaipublic.azureedge.net/main/whisper/models/ed3a0b6b1c0edf879ad9b11b1af5a0e6ab5db9205f891f668f8b0e6c6326e34e/base.pt",
+ "small.en": "https://openaipublic.azureedge.net/main/whisper/models/f953ad0fd29cacd07d5a9eda5624af0f6bcf2258be67c92b79389873d91e0872/small.en.pt",
+ "small": "https://openaipublic.azureedge.net/main/whisper/models/9ecf779972d90ba49c06d968637d720dd632c55bbf19d441fb42bf17a411e794/small.pt",
+ "medium.en": "https://openaipublic.azureedge.net/main/whisper/models/d7440d1dc186f76616474e0ff0b3b6b879abc9d1a4926b7adfa41db2d497ab4f/medium.en.pt",
+ "medium": "https://openaipublic.azureedge.net/main/whisper/models/345ae4da62f9b3d59415adc60127b97c714f32e89e936602e85993674d08dcb1/medium.pt",
+ "large": "https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large.pt",
+ "large-v1": "https://openaipublic.azureedge.net/main/whisper/models/e4b87e7e0bf463eb8e6956e646f1e277e901512310def2c24bf0e11bd3c28e9a/large-v1.pt",
+ "large-v2": "https://openaipublic.azureedge.net/main/whisper/models/81f7c96c852ee8fc832187b0132e569d6c3065a3252ed18e56effd0b6a73e524/large-v2.pt",
+ "large-v3": "https://openaipublic.azureedge.net/main/whisper/models/e5b1a55b89c1367dacf97e3e19bfd829a01529dbfdeefa8caeb59b3f1b81dadb/large-v3.pt",
+}
+
+
+def _download(url: str, root: str, in_memory: bool) -> Union[bytes, str]:
+ os.makedirs(root, exist_ok=True)
+
+ expected_sha256 = url.split("/")[-2]
+ download_target = os.path.join(root, os.path.basename(url))
+
+ if os.path.exists(download_target) and not os.path.isfile(download_target):
+ raise RuntimeError(f"{download_target} exists and is not a regular file")
+
+ if os.path.isfile(download_target):
+ model_bytes = open(download_target, "rb").read()
+ if hashlib.sha256(model_bytes).hexdigest() == expected_sha256:
+ return model_bytes if in_memory else download_target
+ else:
+ warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
+
+ with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
+ with tqdm(total=int(source.info().get("Content-Length")), ncols=80, unit='iB', unit_scale=True, unit_divisor=1024) as loop:
+ while True:
+ buffer = source.read(8192)
+ if not buffer:
+ break
+
+ output.write(buffer)
+ loop.update(len(buffer))
+
+ model_bytes = open(download_target, "rb").read()
+ if hashlib.sha256(model_bytes).hexdigest() != expected_sha256:
+ raise RuntimeError("Model has been downloaded but the SHA256 checksum does not not match. Please retry loading the model.")
+
+ return model_bytes if in_memory else download_target
+
+
+def available_models() -> List[str]:
+ """Returns the names of available models"""
+ return list(_MODELS.keys())
+
+
+def load_model(name: str, device: Optional[Union[str, torch.device]] = None, download_root: str = None, in_memory: bool = False) -> Whisper:
+ """
+ Load a Whisper ASR model
+
+ Parameters
+ ----------
+ name : str
+ one of the official model names listed by `whisper.available_models()`, or
+ path to a model checkpoint containing the model dimensions and the model state_dict.
+ device : Union[str, torch.device]
+ the PyTorch device to put the model into
+ download_root: str
+ path to download the model files; by default, it uses "~/.cache/whisper"
+ in_memory: bool
+ whether to preload the model weights into host memory
+
+ Returns
+ -------
+ model : Whisper
+ The Whisper ASR model instance
+ """
+
+ if device is None:
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+ if download_root is None:
+ download_root = os.getenv(
+ "XDG_CACHE_HOME",
+ os.path.join(os.path.expanduser("~"), ".cache", "whisper")
+ )
+
+ if name in _MODELS:
+ checkpoint_file = _download(_MODELS[name], download_root, in_memory)
+ elif os.path.isfile(name):
+ checkpoint_file = open(name, "rb").read() if in_memory else name
+ else:
+ raise RuntimeError(f"Model {name} not found; available models = {available_models()}")
+
+ with (io.BytesIO(checkpoint_file) if in_memory else open(checkpoint_file, "rb")) as fp:
+ checkpoint = torch.load(fp, map_location=device)
+ del checkpoint_file
+
+ dims = ModelDimensions(**checkpoint["dims"])
+ model = Whisper(dims)
+ model.load_state_dict(checkpoint["model_state_dict"])
+
+ return model.to(device)
diff --git a/musetalk/whisper/whisper/__main__.py b/musetalk/whisper/whisper/__main__.py
new file mode 100644
index 0000000..bc9b04a
--- /dev/null
+++ b/musetalk/whisper/whisper/__main__.py
@@ -0,0 +1,4 @@
+from .transcribe import cli
+
+
+cli()
diff --git a/musetalk/whisper/whisper/__pycache__/__init__.cpython-310.pyc b/musetalk/whisper/whisper/__pycache__/__init__.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..f15e9ea13520eb118b0400000d8cecd20a561674
GIT binary patch
literal 5295
zcmb_gTW{RP6(+gdt5z#nwq?t)6WVQnEMRGQ=hA6{#J4!DZ?K%E*`nACIYTXNUN*y7
zD_X45G|Fr0x4yInB%m(^`ZxN}zrfc%`Mpm;5%&x!DM}2afLtzkIM*|0&N=g)Gn>xO
zw*+|fORoo4PYc4oc;)Cb50&@e&CZL0fCVf@!l3+$1F@(}1F85{29@F~4`ldCQ8lg&
zYH@u~j~j!ASk7q}Xvgz|`M5La6m7L=AzmCT#wP|R;-$e-
zd~$HIXs<`7;?slE@$z7~s5hcB@tMIHC^vBn&*3(n#~r+Y7x4+agiqpA`1D`egR^*<
ztl%^7o+an-3Tcqn$Vwos2tfm%!>@fU56(ktjjTY6gwNv(&~gF)9KQiG6jF!rE2R18
zbv|Q7z!wksoDN#}68_oO)xjHBTo<};KIeM7+pXmC?J#3``cVRfu9(a0
zGaNq$MG7#g6{QFd(FBKSelkiohB1kT*_il@rW+K+`P>-#Q#xccjJtBaaGm%m4wH?N
zZgpFzIcU{CY8Gm-VoLvE^Vx
zN4`&x>glfKyEf5n+tm=#Ox*A5dVqWfyI6Oyg?&}^Odsie-NslodMr$~)=088X4Bsv
zQwH18Ovg4&WNCe@>49nzL&tVtc)ATs`$o{m*tCgdxT)D>>s{{vq$3@unG|kgcU?G**$g{BL>Z%p=T`jQfe$e-z*|!bdCWc8$89)g;
z$FxG=S_Z;ufOXeLK32il$R!wAHfTFg1KaZhUAOuku^r#lJ;ycrZr?Gms}Z&D=)EjP
zQM9jxFOMmNOZ>obT-VXDt9r){W_l)^AAL`;lqJ#ak(Tdt{D2>Z6G>5k_*zN49ep_9H#T*J0iC|NF)
z922X3?0edQ6rhOXnu3_HS>h1Y3ryP}UZ2=-_=s)V0h|%XA+D-f8k`jslR)=%(+gAr
z4vP(6?-LjGB18drUcn>(KN74y(OyM}z7~MH+kW5D3F!xZ-_W(bk2i;sY8!Ea`_OBd<(ol*=?tS>6m2h>~$ixAGb#Y(f*i!h`Pq!!sHFUV}IL
z6NFhNJQ8Oj7P0h=BnmTOSHhJ&5z8yWb_Mz>yYfsdzSWtqC+Q+-`eT-4l1`_k>yHu|#jpBzkjJ
zT^DBZy1+WS3k9G0Gigs01-AI;#O@Mb`|0|2t>Dzyk)FQ0U1ulRsab9J^q#Q0jGKGn
zj#y+kBF}NYsXl#uyU`S&zlA${;&yY^%r=wxS$+4+493Dv8@p%O%B;B~^4e$V$B=!=
z@Y|_~LekKL?2nt!&VQf4`%>?-pX~kn(cYijeTDvmHShx07--?22U_SZ-hBvpM}vG0
zY0@x#{H@q~{@sfgFQyljqA62Qf|W2+l9VY(p=4tM4x$8ML~@CG(_8z)E-T)IDQrZP
z^;_@jmaX_B;y=zNv4T@l&W;iDM{fg*Rh(`nJmG@;shCo7+w`WLYiYtr!d7pSWP^>K
z-v{RzNa58G^d)qwd6lJ%=WvOiMp-UTl91&sPKTL~K%#b0#^IANOKB)5f-ogTlBJAL
zrg(%TAWpe9!T(!gzUM;fHRXOp;2?qUjBY7MWiOJhweGoGO0&EXqsI^{Xa?tO9D#81
z7_m_>--AO2Sx8<7xkpjx<+W^t
zK&^RmL_Wvi2H`pFCZY+b5EQ{nc@54cA6lp6vk4zihX&4zRzZzJP{%~!)l{be6vR#B_vMWG3VnN>b9!?&XldtlnB_JT*hw!