From a0a7851de36a21aa9cccb024b9683ee2686230c8 Mon Sep 17 00:00:00 2001 From: sedherthe Date: Fri, 20 Feb 2026 11:34:29 +0000 Subject: [PATCH] Add codec modules. --- .../__pycache__/decoder.cpython-311.pyc | Bin 0 -> 2125 bytes codec/codec_decoder/decoder.py | 50 ++++ .../encoder/__pycache__/codec.cpython-311.pyc | Bin 0 -> 11364 bytes .../__pycache__/quantizer.cpython-311.pyc | Bin 0 -> 4259 bytes codec/encoder/codec.py | 213 +++++++++++++++ codec/encoder/quantizer.py | 64 +++++ codec_dataset.py | 53 ++++ codec_model.py | 17 ++ codec_train.py | 247 ++++++++++++++++++ 9 files changed, 644 insertions(+) create mode 100644 codec/codec_decoder/__pycache__/decoder.cpython-311.pyc create mode 100644 codec/codec_decoder/decoder.py create mode 100644 codec/encoder/__pycache__/codec.cpython-311.pyc create mode 100644 codec/encoder/__pycache__/quantizer.cpython-311.pyc create mode 100644 codec/encoder/codec.py create mode 100644 codec/encoder/quantizer.py create mode 100644 codec_dataset.py create mode 100644 codec_model.py create mode 100644 codec_train.py diff --git a/codec/codec_decoder/__pycache__/decoder.cpython-311.pyc b/codec/codec_decoder/__pycache__/decoder.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..6b15613e36637ad47a886825072832cc5f9da276 GIT binary patch literal 2125 zcmaJCO>YxN^o>8)X6?oy#ZHV#pftrIA)=+KR4Q7frBy`;s;GMKWwhCyC{EV9?(8}= zu2FM9szFGgiYhtu;8O|;haNfhC)mgmt34t0Qs0aMQpKrn)?O1U5pTS2=FPm%nfIRm zOeErfaQ4i4@nIO?A3^wlv?@|48snW0QBwy{aYsbh zkq|9*gu|*>`XJWA_+E$r{7$yO#@UA*TTBHUv z;B(tie+BU7?rOnWU5qocEp8rM6~&cn=BgTicQrO#W0v|xk0U1az~1W&3K7vMa5=R z?8J>qiR#J*64WYInVZyZ(0w!}!;$@D+znfmGUKj9Y0)U!l)A}E*;2%N zQcG2I#)VIAWT9vEbEX{n!5ZsbeTSU<11l8G)0hJ4HoSxjGcBL*!A z_m42cOApWKqf4f3>kM;oZ4e+IMU^}yP(~<^4Lo;al&ZR7S(IAh%BK3T`TFy@6|=4tVa_fqs%`QOZ&RgOFG1(}^SSZ{Poui(v7!7u=ih{QRqH;Y z|2E;(7{rpdqj#gvq|8$((~>6I(nM35c%D7m%wBr<^ONr${nC^&&WV|Ns4Zn*MIbZr zV552R!e48x^n5!#Uk}}tad&*$$xb^XStmE^q)$30CY|i*ZX!A!e*uVI3WS^qG&i38 zBi163VCsBPzM3X0cL`)Grfs-4vF>kE0X)!m&S(zy|mIYRv z&XjfLDrh%TG_hQhV#yTb|A-=0PuYZcKZ zbR|fn?k(I|*iPNPcK2GHICA3Nl{;6q$@ck{oNLRuI_X9rk$fgkK9wgMq4v~APf|^J zvL#<^%NOh9xjb_3>Yc0GCtLFIwtO5Fj2_#X{cW~ETBE1hqo?Xo=bdBqB!3^x<)d7r zHW&Sgi+bYXKWdR5@o8*6*bWS(#IPO2DJ_{=1=|62WlDdnD0=V2s~tXrLT;j8pDNRJ zO7TkH!(B_H*%#hTaM3J1BoEY`k{tI8J2QjvZs8WS5~hBZ)R&nHl+^ zRTjxs3rmK!QBnogk8B|bp;E4)eO_ct(q-Q_;j261G(8J_eNXU9X)wdFu4xKmM%gZ;$-jthQh0FY2CHBA$ty zk0x}NB5N~BDrPf!&ZrBN#^T{yAXhn!1eL;BD_)grm7e3~_?!GSq;FVx*HGq$RSvzh zQ48(Kcq(kyomyInP3!hVdP;XttHL&gW>N8is3b!azkPXlI-Qh_g6FjEc;S^H{C61#Xg;_<0Z3sBNAb zcPwxh`Ee(DY1h4{;>k=xJ~nzHW-Tq0cHqZ`_Z}jj^VI`<4S6^4Gs2oRY-SGqSM%_k z<}5!wYt=Nz&)KRcVvdVY-7Hj9itf%_h^15KM|8)L!=rEOcFd-kw=?X}9jcs|)a_C{ zsfXezO;(b!6pw0hgmS!*NGuUm)ks7onW89uq=XPr+CXL-tQi^pw1f>CjtM(EN6`SmNM6Pmk+#3Ej1`K1>F8PoLXU)m4$!V_6D#0~4 z`BR+I0;aHCfgG%Rq$=9#EzI<}=&80VuWaZPk(1M4JRM$hcK)Eq z)gTQHXMK@-8e#2al4B5R1=3O!9=-=ssTL%D-0tNhzhsvkk^?`d409zrCOu4HxaR}o zE|@HLmC33|O*$msykp$U)QJ!3M9>4TVv~%?Kh_Q=(y?>VB)V2r##V&BE4hpqBjS{ykfMrWjWT2$jxup&b~A2ySCK$%kauo$JVUrnuThBNw!@>D{fdV`EE-v2jF$nvB7?n?Li%SQVKOKjRnQ# z?2_}GgIJYqRaVE`JKUUIYc^$ZTp5(DS75VfzvQ~%@KIudl zVSmL`z^X(g3Dd2+rj+zdM*ZWj{_D4E|NT$5Uilh;M9!!R6*>)c5>Vt1X5!cVDX__3V})Gdu>+lPPK|IPEiJN0+N-wfy44`ffi zofXdHgfn^J3?ldHn~QGF8(8u!;A5y%2KLB5s4|;0$ZVeo*e5WX{D-#hVLJKRrjs0t z>b1kn?6VqAZjD;BA*^4eB=am=@5K7CuQD~q(;5cH_h}(8)+{7ZDc~vlm_plF*#Q!E zD!cKgdo{pdHIr6lWiLJ3M`S-S*vn>J_h324C74x}`hhiNn9M#TW}B;*RLR1V9Q9`) zu<(sd3m&rTZ(ezGm0uiN8Y8+dOq znpS+PCqE4?9H=4Vp3t=X!6KC|0|j0gL^HyEr5}IF79v|gbax~orDH%0cGwTy$qa(A zP1#;^ora}JD-1j+15_#uQVAGn(Vm_|DtIW#y*OOkqvwpSti>@Mv$o2l4gDmn4#Mm>`-WYyX zAW%!;;XKGj7~mT>M}j9j z?u9SxD=TLyRG9)?sRvvU7+eWJ0XOP_GvuBgjK-O5jo!GNZzWkEg8O7`ob$mYPXM7=BPqIpxNi17LtUgR9(* z+Nk5%zz{22jd|F{8q=M#vXWNyVD-ROpdDYi#d^KLRbg2TlSRf$x5LK9EZdG83-kOQ zOW@uWaO!VjLI0N zQUdjYg+*IUWMtu;s(LL63$_Ks(^0zPFYh%S^M6*+IRLhhjQB1jB#Z3Qy# zjrcu$4dip@Rc;E1?I+zI@|A~3i?&6_qJ7bsaZB8y3-f~!okdSELk?h9{EYCaZD|wQ zVrOP>%E-I7KkO*4)$(Et^Gvsg-MVdPht5X~9tN*Olr+%RC}JuN9#9qqW){Eo?MuW} z4F_(SS7x8DQe_t48rla11#|ieA|`KwguM*{Gg9O{TVXFO66ER|9tXHUN1>^;&@)ol zyuGmXhuN*C|4{#j@$AXd{}9O?I9>EQ1HlIz$Ri>}yH%Rx7+!N-1sZ)5RT;ga0aAVd zqB~`{W|FdpG6)aozGKX>9z`Q{C!oTVj1xcr*pz%x(?gSSxO;{&iV$|yozYAtaS6^g z^R+I>@u_J|4f_lzQLa!V`80USF(PbKNdqWjM3_CLi5JE9-+v#!Y6x3c&~QX9t>HVrN~@}AB@V+%WvX>MJ4ryz6`gt|goxX{$U_Cc;`pg?xm)#7`=fjlAt zULT+_xN9KXwvI0X6|+^Iw;C}BY3W?ulWX3RZ{AX@!_+(=u17==Yxg78?niZj#|>Q1 zb6*PAPu-HT;@+INH!tp`rgs-S!Y95fzU5u{#{QgVOWw03Ykt`@SK(pOU!>U{(%@9< z%{gAPG$6buD_)grX~oHP2v|8?#b@oZHA`cnLMzPiXD#h0kynY@$_KqXaQazG3w<1Z zf@DimT4E8YEGgO7%O&k&{=$y2DdQ}AWj~!GtD3~zwy;3?k9l{Z8mO>4xG`887~!7sUa?8DXbXbYiA4--OVnIvvg6 zIP2h``0AiI8XgocOv{Qa9u(s^g2LI@L{y7S<6Lf59<*v8nU&#H5g<0L0qmw{G!k!u z-$5~&lEhI{mWqQ7+>^*+LQYL-(szv>G zkSd?qOFlD^Uat4K!luoIoi7y{JBp2>b8C?Ux!2ZNbW*~_1?!7$ih0P5_EOBpHMbT0 z6cf0v&BXx4g5*%wQ7pvugs(S$Sy!y5R0EZ_#4_3I$syk1#F$$|s5c#NL=49pks|k< z#L%zQ=@@=iFl`w_w$+Ze1V2H-xxeI^;w1NPc*#>b`}fX!aarhtOAW^x)JfOT<1Xep zyMb&2y8Ezkk$ImGcASF_E8*Pi{`*v8oYeUkPJhlJ!N`3w*9w8pGO2RH$PuAQf7l0U ze~yD^3D=A4aJj+(gLZy`1BNGd0`t{+7TgxhfKyg+-8kn`Y;&%Z1CE-Fj?D1P0v0&f z<|5|6L&qaUy<(m=3rQRv)OxkQ9ck9FUq`T^#K z22}T2H&m5w8ku_p!~bO_VK#KwP(+r?`&SU9ZUWMQ#(&)O1m8^TyB^GmFBx7=$x*x8 z-nIJE>&J54yL0V(^6h)z8U+iTTZpHvaOl14Ksr0{Bl=Dxv#E4$B27>F3%$GOe{*4g zN$W2>*GFmbKihiqZQJj(?YrH!uh7oeznH?w)-29#(FM>{Mr8A=Fw7xZB{ufuLL)l@lI^xUNlK z>|Q&a?|<=@mVfzpuJc5`^8{=A-R91tt96Khg_Q#k3FGt&{3+*&93=7= zMCOS+2U6)+=`Ldj%m&TnW~Bd~_zZVy-3tZ)41wT6ls?P#^vO1D!*^NAx6pq!xGgUX zRstqCSY3Ici@~-#Lf37fYxQJK*qj$Ov#cs$XF;l}38bN--@|V}=tgP{{~NU*=l&WxcY^jcCtl3to84eTf&TMpIW}MzBDO3w+ z%FB7WP`~iFR=m{Kb^#2!?(v7nK literal 0 HcmV?d00001 diff --git a/codec/encoder/__pycache__/quantizer.cpython-311.pyc b/codec/encoder/__pycache__/quantizer.cpython-311.pyc new file mode 100644 index 0000000000000000000000000000000000000000..84e52c33fecafdc161984e3805ed2754d96b00fb GIT binary patch literal 4259 zcmbVPTWk~Q75-=Ju_sP~9mfF^XUUKyXq}r0x3-WC1iH}OQd9~O#Z@~@{EuUUJ(Hao zb1|k$E2{NMskBt7$V#l_N-G)!4|&*1dFW#+@laJB%TnBtkfgnDH5N^?JX;3DJ^{HwRPhqC8M)SHj{rl>S)$W%yb zYF0=Y#;krM7Mn>MsoeBfLe0i#E|HRzd@3g?GqF?W&yCLBG*m5-icM$K=~z}uE3pq@ z%}D2EO^+q@565P2zJ2E4xbFpCbvDXatfByI(_~50;1D^{c!d378-tf8pz{%kC=(Ek z#MvTAf{1tcxQnuPNZd_bKso9L>Ot}1Ug`nLQ!miABDulDeIOEVdEk-fFT9s%o{SfO zS_umO@}CSLzh&+b1HR4n06LGF15(&e3`3wd#s`-=GP;_4eDY zH@>#5;CkI_DFKi0p$~G@xH&7oK>JvpjC@Vb=%--XHw#Xi)(u%PUi=4|hH8(pmW!sd z8mh_ih|>~qMB6P+&jGi^i(*PSpb(hnv@don42+K!D($Bn*d8K0LSG0b$pC9r^@`)^RDd+Z+(`#UsyV~baHv`O83gZO7|DT=E%Ed z|75j)vi45&=}_$P@EXSkzx5RX@->p8XU#*}I-OImtqj_y9fHWm?qeR*O%(T8gy!rir zhid>MIp6;VIhqXpc-HIW34M|`1G`Ic(;qAIv90pBVWVxP51`W;ydl7B3YfPc=zt-e zP>g1TG1|;-c?h*&R?L5FJ|&z~@%4r-q!pS@$hu&tPFFGnSxTe?8^I$2z&|C;>VhK6 z6f>r(36fv{>~!d6WlgvaX;Y{p5Pd+9v(qxAX=O$i@*QJiW5Vd9fS+;DG0J4)yv2>Y zjaOJfgs%<`rHBcLPO>J`O=f6tiiTQ`FrZ~y0LKh(Vt5;F)`x-2k>|n4;_1@Q%;3Rl z@F3uF*Wgm*=!f z^o9?)^840(ND?li-7cdONecRrd;_Rpq+|gXrd9Ron4opQ(NV_E9nm&OEZ|kLJ<`xN z_2}=?b6Jbi)SN;s7D9RuS3dII$W$RtX;aGxw)iF z+A*NXX>m3@b#z1s?2X9YjWzUU<@Yv}|C+UZuv;Gm0_qRudtyw2HI%NV!!zHrkn;LpW6A1x-jtHhANM` zTXw(=0*aCz6 z{|mJ7#zy#nUcG_T5O!?I;Xd#Q?yQCd@XYW>+yD00e|>pO!z49RXMT1p>ZuoH2tLbY zNQ&hKlG4U-?4hy&*&)i+D*)z()o$nfM!n}zlSy4R5Z;jFof59$D_%$$*CmbSw>ABv zK0XX5>z@IEoZs%BKU!`dcoylue{`wq(_>3F%-(OWjGMg^#qZa8!ixt=dp;jtW_}kl zBm2#s{l!yHLpvAVUrd&6nY;Fzp##;>f#Q4513e4A#qq`5r{6IHgVn&`Qo9*Au$%;+ zDo)h`U7wu0ckwqD?_OHCRO-6>)78*WB{XD)MyjC^C^70B7aFC98Q2RG{y>qpeGqD5 zD25d&D?>5t6UD4bb5PfaR);9Ue+e`76}#xH$2A019bV}9+Ce1m0my_Bl7@{6@1IU_-VVE_RW!xYD5S}oda~=3sBcXEZuSWXHt-m$T XVFq5ml@xz8f9Wfn|N1qH_N)F2XbZtL literal 0 HcmV?d00001 diff --git a/codec/encoder/codec.py b/codec/encoder/codec.py new file mode 100644 index 0000000..d92989c --- /dev/null +++ b/codec/encoder/codec.py @@ -0,0 +1,213 @@ +""" +Adapted from https://github.com/gemelo-ai/vocos +""" +from typing import Optional + +import torchaudio +import torch +from torch import nn + +from .quantizer import FSQSTE + +def safe_log(x: torch.Tensor, clip_val: float = 5e-3) -> torch.Tensor: + return torch.log(torch.clip(x, min=clip_val)) + + +class SimpleMLP(nn.Module): + def __init__(self, + dim, + intermediate_dim, + ): + super().__init__() + self.pwconv1 = nn.Linear(dim, intermediate_dim) + self.act = nn.GELU() + self.pwconv2 = nn.Linear(intermediate_dim, dim) + + def forward(self, x): + x = self.pwconv1(x) + x = self.act(x) + x = self.pwconv2(x) + return x + + +class ConvNeXtBlock(nn.Module): + """ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal. + + Args: + dim (int): Number of input channels. + intermediate_dim (int): Dimensionality of the intermediate layer. + layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling. + Defaults to None. + """ + + def __init__( + self, + dim: int, + intermediate_dim: int, + layer_scale_init_value: float, + dw_kernel_size: int = 7, + ): + super().__init__() + self.dwconv = nn.Conv1d(dim, dim, kernel_size=dw_kernel_size, padding=dw_kernel_size//2, groups=dim) # depthwise conv + self.norm = nn.LayerNorm(dim, eps=1e-6) + self.mlp = SimpleMLP(dim, intermediate_dim) + self.gamma = ( + nn.Parameter(layer_scale_init_value * torch.ones(dim), requires_grad=True) + if layer_scale_init_value > 0 + else None + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + residual = x + x = self.dwconv(x) + x = x.transpose(1, 2) # (B, C, T) -> (B, T, C) + x = self.norm(x) + x = self.mlp(x) + if self.gamma is not None: + x = self.gamma * x + x = x.transpose(1, 2) # (B, T, C) -> (B, C, T) + + x = residual + x + return x + + +class VocosBackbone(nn.Module): + """ + Vocos backbone module built with ConvNeXt blocks. + + Args: + input_channels (int): Number of input features channels. + dim (int): Hidden dimension of the model. + intermediate_dim (int): Intermediate dimension used in ConvNeXtBlock. + num_layers (int): Number of ConvNeXtBlock layers. + layer_scale_init_value (float, optional): Initial value for layer scaling. + """ + + def __init__( + self, + input_channels: int, + dim: int, + intermediate_dim: int, + num_layers: int, + input_kernel_size: int = 7, + dw_kernel_size: int = 7, + layer_scale_init_value: Optional[float] = None, + pad: str = 'zeros', + ): + super().__init__() + self.input_channels = input_channels + self.dim = dim + self.embed = nn.Conv1d( + input_channels, + dim, + kernel_size=input_kernel_size, + padding=input_kernel_size//2, + padding_mode=pad + ) + self.norm = nn.LayerNorm(dim, eps=1e-6) + self.convnext = nn.ModuleList([ + ConvNeXtBlock( + dim=dim, + intermediate_dim=intermediate_dim, + dw_kernel_size=dw_kernel_size, + layer_scale_init_value=layer_scale_init_value or 1 / num_layers**0.5, + ) + for _ in range(num_layers) + ]) + self.final_layer_norm = nn.LayerNorm(dim, eps=1e-6) + self.apply(self._init_weights) + + def _init_weights(self, m): + if isinstance(m, (nn.Conv1d, nn.Linear)): + nn.init.trunc_normal_(m.weight, std=0.02) + if m.bias is not None: + nn.init.constant_(m.bias, 0) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """ + Args: + x (Tensor): Input tensor of shape (B, C, L), where B is the batch size, + C denotes output features, and L is the sequence length. + + Returns: + Tensor: Output of shape (B, L, H), where B is the batch size, L is the sequence length, + and H denotes the model dimension. + """ + x = self.embed(x) # (B, C, L) + x = self.norm(x.transpose(1, 2)) + x = x.transpose(1, 2) + for conv_block in self.convnext: + x = conv_block(x) + x = self.final_layer_norm(x.transpose(1, 2)) + x = x.transpose(1, 2) + return x + + +class Encoder(nn.Module): + def __init__(self, + num_input_mels=50, + mel_hop_length=512, + mel_hop_scale=0.25, + encoder_num_layers=8, + encoder_dim=768, + encoder_intermediate_dim=None, + fsq_levels=[8, 8, 5, 5, 5], + dw_kernel=5, + ): + super().__init__() + self.downsample_scale = 2048 // mel_hop_length + self.mel_hop_length = mel_hop_length + self.mel_n_fft = int(mel_hop_length/mel_hop_scale) + self.encoder_dim = encoder_dim + self.encoder_intermediate_dim = encoder_intermediate_dim if encoder_intermediate_dim else encoder_dim*3 + self.encoder_num_layers = encoder_num_layers + self.encoder_initial_channels = num_input_mels + self.bottleneck_channels = 5 + + self.mel_spec = torchaudio.transforms.MelSpectrogram( + sample_rate=32000, + n_fft=self.mel_n_fft, + hop_length=self.mel_hop_length, + n_mels=num_input_mels, + center=True, + power=1, + ) + self.encoder = VocosBackbone(input_channels=self.encoder_initial_channels, + dim=self.encoder_dim, + intermediate_dim=self.encoder_intermediate_dim, + num_layers=self.encoder_num_layers, + input_kernel_size=1, + dw_kernel_size=dw_kernel, + pad='zeros' + ) + self.downsampler = nn.Linear(self.encoder_dim, self.bottleneck_channels) + self.quant = FSQSTE(levels=fsq_levels) + + def encode(self, x): + x = self.encoder(x) + + # import pdb;pdb.set_trace() + + x = x[:, :, ::self.downsample_scale] # What the heck is this? Brute force downsampling from mel -> tokens. + x = x.transpose(1,2) + x = self.downsampler(x) + x = self.quant(x) + return x + + def preprocess(self, audio): + if audio.dim() == 2: # raw audio + x = self.mel_spec(audio) + x = safe_log(x) + elif audio.dim() == 3: # mel spectrogram + x = audio + return x + + def forward(self, audio): + x = self.preprocess(audio) + + # print("done preprocessing: ",x.shape) + # import pdb;pdb.set_trace() + + x = self.encode(x) + codes = self.quant.to_codebook_index(x) + return codes diff --git a/codec/encoder/quantizer.py b/codec/encoder/quantizer.py new file mode 100644 index 0000000..5d36f61 --- /dev/null +++ b/codec/encoder/quantizer.py @@ -0,0 +1,64 @@ +""" +Adapted from https://github.com/duchenzhuang/FSQ-pytorch/blob/main/quantizers/fsq.py#L41 +""" + +import torch +from torch import nn +from einops import rearrange + + +class FSQSTE(nn.Module): + def __init__(self, levels): + super().__init__() + if levels: + self.dim = len(levels) + self._levels = torch.tensor(levels, dtype=torch.int32).view(1, 1, self.dim) + else: + self._levels = levels + + _levels = self._levels + self.register_buffer("levels", _levels, persistent=False) + + _basis = torch.cumprod(torch.tensor([1] + levels[:-1]), + dim=0, + dtype=torch.int32) + self.register_buffer("_basis", _basis, persistent=False) + + def _scale_and_shift(self, zhat_normalized): + half_width = self.levels // 2 + return (zhat_normalized * half_width) + half_width + + def _scale_and_shift_inverse(self, zhat): + half_width = self.levels // 2 + return (zhat - half_width) / half_width + + def indices_to_level_indices(self, indices): + """ Converts indices to indices at each level, perhaps needed for a transformer with factorized embeddings """ + indices = rearrange(indices, '... -> ... 1') + codes_non_centered = (indices // self._basis) % self.levels + return codes_non_centered + + def to_codebook_index(self, zhat): + """ Converts a `code` to an index in the codebook. """ + assert zhat.shape[-1] == self.dim + zhat = self._scale_and_shift(zhat) + indices = (zhat * self._basis).sum(dim = -1).round().to(torch.int32) + return indices + + def from_codebook_index(self, indices): + """ Inverse of `codes_to_indices`. """ + level_indices = self.indices_to_level_indices(indices) + codes = self._scale_and_shift_inverse(level_indices) + return codes + + def forward(self, x): + if self.levels is not None: + + half_levels = (self.levels - 1) * (1 - 1e-3) / 2 + offset = 0.5 - 0.5 * (self.levels % 2) + shift = torch.tan(offset / half_levels) + + x = torch.tanh(x + shift) * half_levels - offset + x = x + (x.round() - x).detach() + x = x / (self.levels // 2) + return x diff --git a/codec_dataset.py b/codec_dataset.py new file mode 100644 index 0000000..e511d90 --- /dev/null +++ b/codec_dataset.py @@ -0,0 +1,53 @@ +import torch +import torchaudio +from torch.utils.data import Dataset +import os +import json + + +class LJSpeechDataset(Dataset): + def __init__(self, root, sample_rate=32000, mode='train'): + """ + root: path to LJSpeech-1.1 directory + """ + self.root = root + self.sample_rate = sample_rate + + # Write code here to handle modes. Take wave files only from mode json. + self.mode = mode + mode_json = os.path.join(root, f"{mode}.json") + with open(mode_json, 'r') as f: + self.dataset = json.load(f) + + # self.wav_dir = os.path.join(root, "wavs") + # self.wav_files = sorted( + # [f for f in os.listdir(self.wav_dir) if f.endswith(".wav")] + # ) + + # assert len(self.wav_files) > 0, "No wav files found!" + + def __len__(self): + return len(self.dataset) + + def __getitem__(self, idx): + # wav_path = os.path.join(self.wav_dir, self.wav_files[idx]) + + # Using train and val json files + item = self.dataset[idx] + text, audio_tokens, wav_path = item + + wav, sr = torchaudio.load(wav_path) + + # mono + if wav.shape[0] > 1: + wav = wav.mean(dim=0, keepdim=True) + + # resample if needed + if sr != self.sample_rate: + wav = torchaudio.functional.resample( + wav, orig_freq=sr, new_freq=self.sample_rate + ) + + # print("wav shape is: ", wav.shape, idx) + + return wav diff --git a/codec_model.py b/codec_model.py new file mode 100644 index 0000000..39949e3 --- /dev/null +++ b/codec_model.py @@ -0,0 +1,17 @@ +import torch +from torch import nn +from codec.encoder.codec import Encoder +from codec.codec_decoder.decoder import SimpleDecoder + + +class FSQAutoEncoder(nn.Module): + def __init__(self, encoder_cfg, decoder_cfg): + super().__init__() + self.encoder = Encoder(**encoder_cfg) + self.decoder = SimpleDecoder(**decoder_cfg) + + def forward(self, audio): + mel = self.encoder.preprocess(audio) + z = self.encoder.encode(mel) # FSQ STE output + mel_hat = self.decoder(z) + return mel_hat, mel diff --git a/codec_train.py b/codec_train.py new file mode 100644 index 0000000..e2ffd19 --- /dev/null +++ b/codec_train.py @@ -0,0 +1,247 @@ +import torch +import torchaudio +from torch.utils.data import DataLoader +from torch.optim import AdamW, Adam +import torch.nn.functional as F +import matplotlib.pyplot as plt +import os +import wandb +from tqdm import tqdm + + +from codec_model import FSQAutoEncoder +from codec_dataset import LJSpeechDataset +from codec.codec_decoder.decoder import SimpleDecoder + + +wandb.init(project="soprano-codec") + + +def pad_collate(batch): + """ + batch: list of tensors [(1, T1), (1, T2), ...] + """ + lengths = torch.tensor([x.shape[-1] for x in batch]) + + max_len = lengths.max().item() + + padded = [ + F.pad(x, (0, max_len - x.shape[-1])) + for x in batch + ] + + audio = torch.stack(padded) # (B, 1, T_max) + return audio, lengths + +def plot_mels(): + pass + + + +dataset = LJSpeechDataset( + root="/home/ubuntu/soma/data/lj_speech/LJSpeech-1.1", + sample_rate=32000, +) + +loader = DataLoader( + dataset, + batch_size=16, + shuffle=True, + drop_last=True, + num_workers=8, + pin_memory=True, + collate_fn=pad_collate +) + +val_dataset = LJSpeechDataset( + root="/home/ubuntu/soma/data/lj_speech/LJSpeech-1.1", + sample_rate=32000, + mode='val' +) + +val_loader = DataLoader( + val_dataset, + batch_size=16, + shuffle=False, + drop_last=True, + num_workers=8, + pin_memory=True, + collate_fn=pad_collate +) +val_loader_it = iter(val_loader) + +# ------------------ +# Config +# ------------------ + +encoder_cfg = dict( + num_input_mels=50, + mel_hop_length=512, + encoder_dim=768, + encoder_num_layers=8, + fsq_levels=[8, 8, 5, 5, 5], +) + +decoder_cfg = dict( + n_mels=50, + encoder_dim=768, + bottleneck_channels=5, + num_layers=8, + upsample_scale=2048 // 512, +) + +device = "cuda" if torch.cuda.is_available() else "cpu" + +# 3. Token rate sanity check +# downsample_scale = 2048 // mel_hop_length +# total_hop = mel_hop_length * downsample_scale = 2048 +# rate = sr / 2048 +sr = 32000 +total_hop = 2048 # This is the token hop rate. The mel hop rate is 512 +print(f"Token rate: {sr / total_hop:.2f} Hz") + + +# ------------------ +# Setup +# ------------------ + +model = FSQAutoEncoder(encoder_cfg, decoder_cfg).to(device) + +freeze_encoder = False +if freeze_encoder: + model_ckpt_path = "/home/ubuntu/soma/ckpt/suprano/suprano_codec/codec_1/step_42000.pt" + + if os.path.exists(model_ckpt_path): + print(f"Loading model from {model_ckpt_path}") + model.load_state_dict(torch.load(model_ckpt_path)) + + # fix encoder. train only the decoder. reset the decoder weights. + for param in model.encoder.parameters(): + param.requires_grad = False + + for name, p in model.named_parameters(): + if "quant" in name: + print(name, p.requires_grad) + + model.decoder = SimpleDecoder(**decoder_cfg).to(device) + +optimizer = Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=1e-4) + +import pdb;pdb.set_trace() + + +root_dir = "/home/ubuntu/soma/ckpt/suprano/suprano_codec" +ckpt_dir = os.path.join(root_dir, "codec_v2") +plot_dir = os.path.join(ckpt_dir, "plots") +os.makedirs(ckpt_dir, exist_ok=True) +os.makedirs(plot_dir, exist_ok=True) + +num_epochs = 100 +step = 0 + +for epoch in range(num_epochs): + + for epoch_step, data in tqdm(enumerate(loader), total=len(loader)): + + step += 1 + + audio, lengths = data + audio = audio.squeeze().to(device) + + # import pdb;pdb.set_trace() + + mel_hat, mel = model(audio) + + # crop to match length (upsampling can overshoot) + T = min(mel_hat.shape[-1], mel.shape[-1]) + mel_hat = mel_hat[..., :T] + mel = mel[..., :T] + + loss = torch.mean(torch.abs(mel_hat - mel)) + + optimizer.zero_grad() + loss.backward() + optimizer.step() + + if step % 100 == 0: + print(f"step {step} | loss {loss.item():.4f}") + wandb.log({"train/loss": loss.item()}, step=step) + + if step % 400 == 0: + val_loss = 0.0 + val_steps = 10 + model.eval() + with torch.no_grad(): + for _ in range(val_steps): + try: + vdata = next(val_loader_it) + except StopIteration: + val_loader_it = iter(val_loader) + vdata = next(val_loader_it) + + vaudio, vlengths = vdata + vaudio = vaudio.squeeze().to(device) + + vmel_hat, vmel = model(vaudio) + + T_v = min(vmel_hat.shape[-1], vmel.shape[-1]) + vmel_hat = vmel_hat[..., :T_v] + vmel = vmel[..., :T_v] + + val_loss += torch.mean(torch.abs(vmel_hat - vmel)).item() + + val_loss /= val_steps + print(f"step {step} | val_loss {val_loss:.4f}") + wandb.log({"val/loss": val_loss}, step=step) + + # Val Plotting + vmel_np = vmel[0].detach().cpu().numpy() + vmel_hat_np = vmel_hat[0].detach().cpu().numpy() + + fig, axs = plt.subplots(2, 1, figsize=(10, 6)) + + axs[0].imshow(vmel_np, aspect="auto", origin="lower") + axs[0].set_title("Val Original Mel") + + axs[1].imshow(vmel_hat_np, aspect="auto", origin="lower") + axs[1].set_title("Val Reconstructed Mel") + + plt.tight_layout() + plt.savefig(f"{plot_dir}/val_step_{step:05d}.png") + wandb.log({"val/reconstruction": wandb.Image(f"{plot_dir}/val_step_{step:05d}.png")}, step=step) + plt.close() + + model.train() + + if step % 400 == 0: + with torch.no_grad(): + # 1. FSQ bin usage + z = model.encoder.encode(mel) + indices = model.encoder.quant.to_codebook_index(z) + total_bins = int(torch.prod(model.encoder.quant.levels)) + unique_bins = len(torch.unique(indices)) + print(f"[FSQ] unique bins: {unique_bins} / {total_bins}. Lens: {indices.shape} {z.shape}") + wandb.log({"train/unique_bins": unique_bins}, step=step) + + if step % 200 == 0: + mel_np = mel[0].detach().cpu().numpy() + mel_hat_np = mel_hat[0].detach().cpu().numpy() + + fig, axs = plt.subplots(2, 1, figsize=(10, 6)) + + axs[0].imshow(mel_np, aspect="auto", origin="lower") + axs[0].set_title("Original Mel") + + axs[1].imshow(mel_hat_np, aspect="auto", origin="lower") + axs[1].set_title("Reconstructed Mel") + + plt.tight_layout() + plt.savefig(f"{plot_dir}/step_{step:05d}.png") + wandb.log({"train/reconstruction": wandb.Image(f"{plot_dir}/step_{step:05d}.png")}, step=step) + plt.close() + + + if step % 1000 == 0: + ckpt_path = os.path.join(ckpt_dir, f"step_{step:05d}.pt") + torch.save(model.state_dict(), ckpt_path) + print(f"Saved checkpoint to {ckpt_path}")