From 263103e9b1e0c67cef3b85c064c594e45a2c7d20 Mon Sep 17 00:00:00 2001 From: joerenner Date: Wed, 5 Jun 2024 23:37:37 +0000 Subject: [PATCH 1/7] var training --- dac/model/dac.py | 87 +++++++++++++++++++++++++++--------------------- 1 file changed, 49 insertions(+), 38 deletions(-) diff --git a/dac/model/dac.py b/dac/model/dac.py index eb754b2..de9cdde 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -4,6 +4,7 @@ import numpy as np import torch +import torch.nn.functional as F from audiotools import AudioSignal from audiotools.ml import BaseModel from torch import nn @@ -12,7 +13,7 @@ from dac.nn.layers import Snake1d from dac.nn.layers import WNConv1d from dac.nn.layers import WNConvTranspose1d -from dac.nn.quantize import ResidualVectorQuantize +from dac.nn.quantize import VectorQuantize def init_weights(m): @@ -144,27 +145,27 @@ def forward(self, x): return self.model(x) -class DAC(BaseModel, CodecMixin): +class DAC(BaseModel): def __init__( self, encoder_dim: int = 64, encoder_rates: List[int] = [2, 4, 8, 8], - latent_dim: int = None, + latent_dim: int = 128, decoder_dim: int = 1536, decoder_rates: List[int] = [8, 8, 4, 2], - n_codebooks: int = 9, codebook_size: int = 1024, codebook_dim: Union[int, list] = 8, - quantizer_dropout: bool = False, sample_rate: int = 44100, + levels: int = 4 ): super().__init__() - + self.levels = levels self.encoder_dim = encoder_dim self.encoder_rates = encoder_rates self.decoder_dim = decoder_dim self.decoder_rates = decoder_rates self.sample_rate = sample_rate + self.sample_rates = [math.ceil(sample_rate / (2**d)) for d in reversed(range(self.levels))] if latent_dim is None: latent_dim = encoder_dim * (2 ** len(encoder_rates)) @@ -172,17 +173,15 @@ def __init__( self.latent_dim = latent_dim self.hop_length = np.prod(encoder_rates) - self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim) + self.encoders = nn.ModuleList([Encoder(encoder_dim, encoder_rates, latent_dim) for _ in range(self.levels)]) + self.phis = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(self.levels)]) - self.n_codebooks = n_codebooks self.codebook_size = codebook_size self.codebook_dim = codebook_dim - self.quantizer = ResidualVectorQuantize( + self.quantizer = VectorQuantize( input_dim=latent_dim, - n_codebooks=n_codebooks, codebook_size=codebook_size, codebook_dim=codebook_dim, - quantizer_dropout=quantizer_dropout, ) self.decoder = Decoder( @@ -193,8 +192,6 @@ def __init__( self.sample_rate = sample_rate self.apply(init_weights) - self.delay = self.get_delay() - def preprocess(self, audio_data, sample_rate): if sample_rate is None: sample_rate = self.sample_rate @@ -208,9 +205,8 @@ def preprocess(self, audio_data, sample_rate): def encode( self, - audio_data: torch.Tensor, - n_quantizers: int = None, - ): + audio_data: torch.Tensor + ): """Encode given audio data and return quantized latent codes Parameters @@ -225,13 +221,10 @@ def encode( ------- dict A dictionary with the following keys: - "z" : Tensor[B x D x T] - Quantized continuous representation of input + "codes" : Tensor[B x N x T] Codebook indices for each codebook (quantized discrete representation of input) - "latents" : Tensor[B x N*D x T] - Projected latents (continuous representation of input before quantization) "vq/commitment_loss" : Tensor[1] Commitment loss to train encoder to predict vectors closer to codebook entries @@ -240,18 +233,32 @@ def encode( "length" : int Number of samples in input audio """ - z = self.encoder(audio_data) - z, codes, latents, commitment_loss, codebook_loss = self.quantizer( - z, n_quantizers - ) - return z, codes, latents, commitment_loss, codebook_loss - - def decode(self, z: torch.Tensor): + f = None + commitment_loss = 0.0 + codebook_loss = 0.0 + indices = [] + for i in range(self.levels): + x = AudioSignal(audio_data, self.sample_rate).resample(self.sample_rates[i]) + x = self.encoders[i](x.audio_data) + if 0 < i < len(self.sample_rates) - 1: + resized_z_q_i = F.interpolate( + z_q_i, + size=x.shape[2], + mode="nearest" + ) + x -= self.phis[i](resized_z_q_i) + z_q_i, commitment_loss_i, codebook_loss_i, indices_i, _ = self.quantizer(x) + indices.append(indices_i) + commitment_loss += commitment_loss_i.mean() + codebook_loss += codebook_loss_i.mean() + return indices, commitment_loss, codebook_loss + + def decode(self, codes): """Decode given latent codes and return audio data Parameters ---------- - z : Tensor[B x D x T] + codes : list of tensors Tensor[B x T] Quantized continuous representation of input length : int, optional Number of samples in output audio, by default None @@ -263,13 +270,21 @@ def decode(self, z: torch.Tensor): "audio" : Tensor[B x 1 x length] Decoded audio data. """ + z = self.phis[-1](self.quantizer.out_proj(self.quantizer.decode_code(codes[-1]))) + for i in range(self.levels-1): + z_q = self.quantizer.out_proj(self.quantizer.decode_code(codes[i])) + z_q = F.interpolate( + z_q, + size=z.shape[2], + mode="nearest" + ) + z += self.phis[i](z_q) return self.decoder(z) def forward( self, audio_data: torch.Tensor, - sample_rate: int = None, - n_quantizers: int = None, + sample_rate: int = None ): """Model forward pass @@ -280,9 +295,6 @@ def forward( sample_rate : int, optional Sample rate of audio data in Hz, by default None If None, defaults to `self.sample_rate` - n_quantizers : int, optional - Number of quantizers to use, by default None. - If None, all quantizers are used. Returns ------- @@ -307,16 +319,14 @@ def forward( """ length = audio_data.shape[-1] audio_data = self.preprocess(audio_data, sample_rate) - z, codes, latents, commitment_loss, codebook_loss = self.encode( - audio_data, n_quantizers + codes, commitment_loss, codebook_loss = self.encode( + audio_data ) - x = self.decode(z) + x = self.decode(codes) return { "audio": x[..., :length], - "z": z, "codes": codes, - "latents": latents, "vq/commitment_loss": commitment_loss, "vq/codebook_loss": codebook_loss, } @@ -352,6 +362,7 @@ def forward( # Make a backward pass out.backward(grad) + exit(0) # Check non-zero values gradmap = x.grad.squeeze(0) From b113ad17a00adc92339381ad46c3973ca87f995c Mon Sep 17 00:00:00 2001 From: joerenner Date: Sun, 9 Jun 2024 19:20:58 +0000 Subject: [PATCH 2/7] adding back RVQ --- dac/model/dac.py | 35 ++++++++++++++++++----------------- scripts/train.py | 3 ++- 2 files changed, 20 insertions(+), 18 deletions(-) diff --git a/dac/model/dac.py b/dac/model/dac.py index de9cdde..a3dfeb9 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -156,16 +156,15 @@ def __init__( codebook_size: int = 1024, codebook_dim: Union[int, list] = 8, sample_rate: int = 44100, - levels: int = 4 + sample_rates: list = [44100] ): super().__init__() - self.levels = levels self.encoder_dim = encoder_dim self.encoder_rates = encoder_rates self.decoder_dim = decoder_dim self.decoder_rates = decoder_rates self.sample_rate = sample_rate - self.sample_rates = [math.ceil(sample_rate / (2**d)) for d in reversed(range(self.levels))] + self.sample_rates = sample_rates if latent_dim is None: latent_dim = encoder_dim * (2 ** len(encoder_rates)) @@ -173,16 +172,17 @@ def __init__( self.latent_dim = latent_dim self.hop_length = np.prod(encoder_rates) - self.encoders = nn.ModuleList([Encoder(encoder_dim, encoder_rates, latent_dim) for _ in range(self.levels)]) - self.phis = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(self.levels)]) + self.encoders = nn.ModuleList([Encoder(encoder_dim, encoder_rates, latent_dim) for _ in range(len(self.sample_rates))]) + self.phis = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(len(self.sample_rates))]) + self.quantizers = nn.ModuleList( + [ + VectorQuantize(latent_dim, codebook_size, codebook_dim) + for i in range(len(self.sample_rates)) + ] + ) self.codebook_size = codebook_size self.codebook_dim = codebook_dim - self.quantizer = VectorQuantize( - input_dim=latent_dim, - codebook_size=codebook_size, - codebook_dim=codebook_dim, - ) self.decoder = Decoder( latent_dim, @@ -233,21 +233,22 @@ def encode( "length" : int Number of samples in input audio """ - f = None commitment_loss = 0.0 codebook_loss = 0.0 indices = [] - for i in range(self.levels): + for i in range(len(self.sample_rates)): x = AudioSignal(audio_data, self.sample_rate).resample(self.sample_rates[i]) x = self.encoders[i](x.audio_data) - if 0 < i < len(self.sample_rates) - 1: + + if i > 0: resized_z_q_i = F.interpolate( z_q_i, size=x.shape[2], mode="nearest" ) x -= self.phis[i](resized_z_q_i) - z_q_i, commitment_loss_i, codebook_loss_i, indices_i, _ = self.quantizer(x) + + z_q_i, commitment_loss_i, codebook_loss_i, indices_i, _ = self.quantizers[i](x) indices.append(indices_i) commitment_loss += commitment_loss_i.mean() codebook_loss += codebook_loss_i.mean() @@ -270,9 +271,9 @@ def decode(self, codes): "audio" : Tensor[B x 1 x length] Decoded audio data. """ - z = self.phis[-1](self.quantizer.out_proj(self.quantizer.decode_code(codes[-1]))) - for i in range(self.levels-1): - z_q = self.quantizer.out_proj(self.quantizer.decode_code(codes[i])) + z = self.phis[-1](self.quantizers[-1].out_proj(self.quantizers[-1].decode_code(codes[-1]))) + for i in range(len(self.sample_rates)-1): + z_q = self.quantizers[i].out_proj(self.quantizers[i].decode_code(codes[i])) z_q = F.interpolate( z_q, size=z.shape[2], diff --git a/scripts/train.py b/scripts/train.py index 646ed57..eca0e6a 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -17,7 +17,7 @@ from audiotools.ml.decorators import Tracker from audiotools.ml.decorators import when from torch.utils.tensorboard import SummaryWriter - +from torch.distributed.elastic.multiprocessing.errors import record import dac warnings.filterwarnings("ignore", category=UserWarning) @@ -347,6 +347,7 @@ def validate(state, val_dataloader, accel): return output + @argbind.bind(without_prefix=True) def train( args, From 52c948e7e7dc2b4aaea2797c54c6326e28f3f537 Mon Sep 17 00:00:00 2001 From: joerenner Date: Mon, 10 Jun 2024 15:47:10 +0000 Subject: [PATCH 3/7] adding classic var with no downsampling --- dac/model/dac.py | 60 ++++++++++++++++++++++++++++-------------------- 1 file changed, 35 insertions(+), 25 deletions(-) diff --git a/dac/model/dac.py b/dac/model/dac.py index a3dfeb9..48ed931 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -156,7 +156,8 @@ def __init__( codebook_size: int = 1024, codebook_dim: Union[int, list] = 8, sample_rate: int = 44100, - sample_rates: list = [44100] + levels: int = 5 + #sample_rates: list = [44100] ): super().__init__() self.encoder_dim = encoder_dim @@ -164,7 +165,8 @@ def __init__( self.decoder_dim = decoder_dim self.decoder_rates = decoder_rates self.sample_rate = sample_rate - self.sample_rates = sample_rates + # self.sample_rates = sample_rates + self.levels = 5 if latent_dim is None: latent_dim = encoder_dim * (2 ** len(encoder_rates)) @@ -172,12 +174,14 @@ def __init__( self.latent_dim = latent_dim self.hop_length = np.prod(encoder_rates) - self.encoders = nn.ModuleList([Encoder(encoder_dim, encoder_rates, latent_dim) for _ in range(len(self.sample_rates))]) - self.phis = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(len(self.sample_rates))]) + self.resolutions = [(1, 2, 2), (2, 4, 4), (4, 8, 8), (4, 12, 12), (4, 16, 16), (4, 32, 32)] + self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim) + # self.encoders = nn.ModuleList([Encoder(encoder_dim, encoder_rates, latent_dim) for _ in range(len(self.sample_rates))]) + self.phis = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(self.levels)]) self.quantizers = nn.ModuleList( [ VectorQuantize(latent_dim, codebook_size, codebook_dim) - for i in range(len(self.sample_rates)) + for i in range(self.levels) ] ) @@ -236,25 +240,31 @@ def encode( commitment_loss = 0.0 codebook_loss = 0.0 indices = [] - for i in range(len(self.sample_rates)): - x = AudioSignal(audio_data, self.sample_rate).resample(self.sample_rates[i]) - x = self.encoders[i](x.audio_data) - - if i > 0: - resized_z_q_i = F.interpolate( - z_q_i, - size=x.shape[2], - mode="nearest" - ) - x -= self.phis[i](resized_z_q_i) - - z_q_i, commitment_loss_i, codebook_loss_i, indices_i, _ = self.quantizers[i](x) + z = self.encoder(audio_data) + residual = z + quantized = 0 + for i in range(self.levels): + resized_z = F.interpolate( + residual, + size=int(residual.shape[2] / (2**(self.levels - i - 1))), + mode="nearest" + ) + z_q_i, commitment_loss_i, codebook_loss_i, indices_i, _ = self.quantizers[i](resized_z) + resized_z_q_i = F.interpolate( + z_q_i, + size=residual.shape[2], + mode="nearest" + ) + resized_z_q_i = self.phis[i](resized_z_q_i) + residual -= resized_z_q_i + quantized = quantized + resized_z_q_i indices.append(indices_i) commitment_loss += commitment_loss_i.mean() codebook_loss += codebook_loss_i.mean() - return indices, commitment_loss, codebook_loss - def decode(self, codes): + return indices, commitment_loss, codebook_loss, quantized + + def decode(self, quantized): """Decode given latent codes and return audio data Parameters @@ -271,7 +281,7 @@ def decode(self, codes): "audio" : Tensor[B x 1 x length] Decoded audio data. """ - z = self.phis[-1](self.quantizers[-1].out_proj(self.quantizers[-1].decode_code(codes[-1]))) + """z = self.phis[-1](self.quantizers[-1].out_proj(self.quantizers[-1].decode_code(codes[-1]))) for i in range(len(self.sample_rates)-1): z_q = self.quantizers[i].out_proj(self.quantizers[i].decode_code(codes[i])) z_q = F.interpolate( @@ -279,8 +289,8 @@ def decode(self, codes): size=z.shape[2], mode="nearest" ) - z += self.phis[i](z_q) - return self.decoder(z) + z += self.phis[i](z_q)""" + return self.decoder(quantized) def forward( self, @@ -320,11 +330,11 @@ def forward( """ length = audio_data.shape[-1] audio_data = self.preprocess(audio_data, sample_rate) - codes, commitment_loss, codebook_loss = self.encode( + codes, commitment_loss, codebook_loss, quantized = self.encode( audio_data ) - x = self.decode(codes) + x = self.decode(quantized) return { "audio": x[..., :length], "codes": codes, From d0fcb9d2a7718bcfd2e57e38dec484df83ffaceb Mon Sep 17 00:00:00 2001 From: joerenner Date: Wed, 12 Jun 2024 16:46:32 +0000 Subject: [PATCH 4/7] var with fsq quantization --- dac/model/dac.py | 65 +++++++++++++++++++++------------------------- dac/nn/quantize.py | 23 ++++++++++++++++ scripts/train.py | 10 +++---- 3 files changed, 57 insertions(+), 41 deletions(-) diff --git a/dac/model/dac.py b/dac/model/dac.py index 48ed931..1a142ef 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -13,7 +13,7 @@ from dac.nn.layers import Snake1d from dac.nn.layers import WNConv1d from dac.nn.layers import WNConvTranspose1d -from dac.nn.quantize import VectorQuantize +from dac.nn.quantize import VectorQuantize, ResidualVectorQuantize, FSQ def init_weights(m): @@ -153,11 +153,11 @@ def __init__( latent_dim: int = 128, decoder_dim: int = 1536, decoder_rates: List[int] = [8, 8, 4, 2], - codebook_size: int = 1024, + codebook_sizes: Union[int, list] = 1024, codebook_dim: Union[int, list] = 8, sample_rate: int = 44100, - levels: int = 5 - #sample_rates: list = [44100] + channel_floats: list = [8, 8, 8, 5, 5, 5], + downsample_rates: list = [2] ): super().__init__() self.encoder_dim = encoder_dim @@ -165,8 +165,9 @@ def __init__( self.decoder_dim = decoder_dim self.decoder_rates = decoder_rates self.sample_rate = sample_rate - # self.sample_rates = sample_rates - self.levels = 5 + self.downsample_rates = downsample_rates + self.codebook_dim = codebook_dim + self.channel_floats = channel_floats if latent_dim is None: latent_dim = encoder_dim * (2 ** len(encoder_rates)) @@ -174,19 +175,17 @@ def __init__( self.latent_dim = latent_dim self.hop_length = np.prod(encoder_rates) - self.resolutions = [(1, 2, 2), (2, 4, 4), (4, 8, 8), (4, 12, 12), (4, 16, 16), (4, 32, 32)] self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim) - # self.encoders = nn.ModuleList([Encoder(encoder_dim, encoder_rates, latent_dim) for _ in range(len(self.sample_rates))]) - self.phis = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(self.levels)]) + self.phis_downsample = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(len(self.downsample_rates))]) + self.phis_upsample = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(len(self.downsample_rates))]) self.quantizers = nn.ModuleList( [ - VectorQuantize(latent_dim, codebook_size, codebook_dim) - for i in range(self.levels) + FSQ(channel_floats, latent_dim) + for i in range(len(self.downsample_rates)) ] ) - self.codebook_size = codebook_size - self.codebook_dim = codebook_dim + self.codebook_sizes = codebook_sizes self.decoder = Decoder( latent_dim, @@ -243,26 +242,32 @@ def encode( z = self.encoder(audio_data) residual = z quantized = 0 - for i in range(self.levels): + quantized_list = [] + for i in range(len(self.downsample_rates)): resized_z = F.interpolate( residual, - size=int(residual.shape[2] / (2**(self.levels - i - 1))), + size=int(residual.shape[2] / self.downsample_rates[i]), mode="nearest" ) - z_q_i, commitment_loss_i, codebook_loss_i, indices_i, _ = self.quantizers[i](resized_z) + resized_z = self.phis_downsample[i](resized_z) + z_q_i = self.quantizers[i].quantize(resized_z) resized_z_q_i = F.interpolate( z_q_i, size=residual.shape[2], mode="nearest" ) - resized_z_q_i = self.phis[i](resized_z_q_i) + resized_z_q_i = self.phis_upsample[i](resized_z_q_i) residual -= resized_z_q_i - quantized = quantized + resized_z_q_i - indices.append(indices_i) - commitment_loss += commitment_loss_i.mean() - codebook_loss += codebook_loss_i.mean() + quantized += resized_z_q_i + quantized_list.append(quantized) + # indices.append(indices_i) + + identity_loss = 0.0 + z_detach = z.detach() + for quantized in quantized_list: + identity_loss = identity_loss + F.mse_loss(quantized, z_detach) - return indices, commitment_loss, codebook_loss, quantized + return quantized, identity_loss def decode(self, quantized): """Decode given latent codes and return audio data @@ -281,15 +286,6 @@ def decode(self, quantized): "audio" : Tensor[B x 1 x length] Decoded audio data. """ - """z = self.phis[-1](self.quantizers[-1].out_proj(self.quantizers[-1].decode_code(codes[-1]))) - for i in range(len(self.sample_rates)-1): - z_q = self.quantizers[i].out_proj(self.quantizers[i].decode_code(codes[i])) - z_q = F.interpolate( - z_q, - size=z.shape[2], - mode="nearest" - ) - z += self.phis[i](z_q)""" return self.decoder(quantized) def forward( @@ -330,16 +326,13 @@ def forward( """ length = audio_data.shape[-1] audio_data = self.preprocess(audio_data, sample_rate) - codes, commitment_loss, codebook_loss, quantized = self.encode( + quantized, identity_loss = self.encode( audio_data ) - x = self.decode(quantized) return { "audio": x[..., :length], - "codes": codes, - "vq/commitment_loss": commitment_loss, - "vq/codebook_loss": codebook_loss, + "fsq/identity_loss": identity_loss, } diff --git a/dac/nn/quantize.py b/dac/nn/quantize.py index b17ff4a..6b1114d 100644 --- a/dac/nn/quantize.py +++ b/dac/nn/quantize.py @@ -94,6 +94,29 @@ def decode_latents(self, latents): return z_q, indices +class FSQ(nn.Module): + def __init__( + self, + channel_floats: int, + input_dim: int, + ): + super().__init__() + _levels = torch.tensor(channel_floats).int() + codebook_dim = len(channel_floats) + self.register_buffer("_levels", _levels, persistent = False) + self.in_proj = WNConv1d(input_dim, codebook_dim, kernel_size=1) + self.out_proj = WNConv1d(codebook_dim, input_dim, kernel_size=1) + + def quantize(self, z, eps=1e-1): + z = self.in_proj(z) + half_l = self._levels.reshape(1, -1, 1) / 2 - eps + z = z.tanh() * half_l + z_q = z.round() + z_q = z + (z_q - z).detach() + z_q = z_q / half_l + return self.out_proj(z_q) + + class ResidualVectorQuantize(nn.Module): """ Introduced in SoundStream: An end2end neural audio codec diff --git a/scripts/train.py b/scripts/train.py index eca0e6a..5bfec1e 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -237,8 +237,7 @@ def train_loop(state, batch, accel, lambdas): with accel.autocast(): out = state.generator(signal.audio_data, signal.sample_rate) recons = AudioSignal(out["audio"], signal.sample_rate) - commitment_loss = out["vq/commitment_loss"] - codebook_loss = out["vq/codebook_loss"] + identity_loss = out["fsq/identity_loss"] with accel.autocast(): output["adv/disc_loss"] = state.gan_loss.discriminator_loss(recons, signal) @@ -260,8 +259,7 @@ def train_loop(state, batch, accel, lambdas): output["adv/gen_loss"], output["adv/feat_loss"], ) = state.gan_loss.generator_loss(recons, signal) - output["vq/commitment_loss"] = commitment_loss - output["vq/codebook_loss"] = codebook_loss + output["fsq/identity_loss"] = identity_loss output["loss"] = sum([v * output[k] for k, v in lambdas.items() if k in output]) state.optimizer_g.zero_grad() @@ -417,9 +415,10 @@ def train( last_iter = ( tracker.step == num_iters - 1 if num_iters is not None else False ) + """ if tracker.step % sample_freq == 0 or last_iter: save_samples(state, val_idx, writer) - + """ if tracker.step % valid_freq == 0 or last_iter: validate(state, val_dataloader, accel) checkpoint(state, save_iters, save_path) @@ -430,6 +429,7 @@ def train( break + if __name__ == "__main__": args = argbind.parse_args() args["args.debug"] = int(os.getenv("LOCAL_RANK", 0)) == 0 From fdb42776023b0f04938b0edab9b25313b5d3ce7d Mon Sep 17 00:00:00 2001 From: joerenner Date: Mon, 17 Jun 2024 18:05:38 +0000 Subject: [PATCH 5/7] adding var config --- conf/var.yml | 106 +++++++++++++++++++++++++++++++++++++++++++++++ dac/model/dac.py | 7 +--- 2 files changed, 107 insertions(+), 6 deletions(-) create mode 100644 conf/var.yml diff --git a/conf/var.yml b/conf/var.yml new file mode 100644 index 0000000..3d0784c --- /dev/null +++ b/conf/var.yml @@ -0,0 +1,106 @@ +# Model setup +DAC.sample_rate: 44100 +DAC.encoder_dim: 256 +DAC.encoder_rates: [2, 4, 8, 8] +DAC.decoder_dim: 1536 +DAC.decoder_rates: [8, 8, 4, 2] +DAC.latent_dim: 768 + +# Quantization +DAC.downsample_rates: [8, 6, 5, 4, 3, 2, 1, 1] +DAC.channel_floats: [8, 8, 8, 5, 5, 5] + +# Discriminator +Discriminator.sample_rate: 44100 +Discriminator.rates: [] +Discriminator.periods: [2, 3, 5, 7, 11] +Discriminator.fft_sizes: [2048, 1024, 512] +Discriminator.bands: + - [0.0, 0.1] + - [0.1, 0.25] + - [0.25, 0.5] + - [0.5, 0.75] + - [0.75, 1.0] + +# Optimization +AdamW.betas: [0.8, 0.99] +AdamW.lr: 0.00004 +ExponentialLR.gamma: 0.999996 + +amp: false +val_batch_size: 50 +device: cuda +num_iters: 400000 +save_iters: [10000, 50000, 100000, 200000] +valid_freq: 1000 +sample_freq: 10000 +num_workers: 32 +val_idx: [0, 1, 2, 3, 4, 5, 6, 7] +seed: 0 +lambdas: + mel/loss: 15.0 + adv/feat_loss: 2.0 + adv/gen_loss: 1.0 + fsq/identity_loss: 1.0 + +VolumeNorm.db: [const, -16] + +# Transforms +build_transform.preprocess: + - Identity +build_transform.augment_prob: 0.0 +build_transform.augment: + - Identity +build_transform.postprocess: + - VolumeNorm + - RescaleAudio + - ShiftPhase + +# Loss setup +MultiScaleSTFTLoss.window_lengths: [2048, 512] +MelSpectrogramLoss.n_mels: [5, 10, 20, 40, 80, 160, 320] +MelSpectrogramLoss.window_lengths: [32, 64, 128, 256, 512, 1024, 2048] +MelSpectrogramLoss.mel_fmin: [0, 0, 0, 0, 0, 0, 0] +MelSpectrogramLoss.mel_fmax: [null, null, null, null, null, null, null] +MelSpectrogramLoss.pow: 1.0 +MelSpectrogramLoss.clamp_eps: 1.0e-5 +MelSpectrogramLoss.mag_weight: 0.0 + +# Data +batch_size: 16 +train/AudioDataset.duration: 0.4 +train/AudioDataset.n_examples: 10000000 + +val/AudioDataset.duration: 2.5 +val/build_transform.augment_prob: 1.0 +val/AudioDataset.n_examples: 250 + +test/AudioDataset.duration: 5.0 +test/build_transform.augment_prob: 1.0 +test/AudioDataset.n_examples: 1000 + +AudioLoader.shuffle: true +AudioDataset.without_replacement: true + +train/build_dataset.folders: + speech_fb: + - /data/shared_data/raw/daps/train + speech_hq: + - /data/shared_data/raw/vctk + - /data/shared_data/raw/vocalset + - /data/shared_data/raw/read_speech + - /data/shared_data/raw/french_speech + speech_uq: + - /data/shared_data/raw/emotional_speech/ + # - /data/shared_data/raw/common_voice/ + - /data/shared_data/raw/german_speech/ + - /data/shared_data/raw/russian_speech/ + - /data/shared_data/raw/spanish_speech/ + +val/build_dataset.folders: + speech_hq: + - /data/shared_data/raw/daps/val + +test/build_dataset.folders: + speech_hq: + - /data/shared_data/raw/daps/test diff --git a/dac/model/dac.py b/dac/model/dac.py index 1a142ef..4816035 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -153,8 +153,6 @@ def __init__( latent_dim: int = 128, decoder_dim: int = 1536, decoder_rates: List[int] = [8, 8, 4, 2], - codebook_sizes: Union[int, list] = 1024, - codebook_dim: Union[int, list] = 8, sample_rate: int = 44100, channel_floats: list = [8, 8, 8, 5, 5, 5], downsample_rates: list = [2] @@ -166,7 +164,6 @@ def __init__( self.decoder_rates = decoder_rates self.sample_rate = sample_rate self.downsample_rates = downsample_rates - self.codebook_dim = codebook_dim self.channel_floats = channel_floats if latent_dim is None: @@ -185,8 +182,6 @@ def __init__( ] ) - self.codebook_sizes = codebook_sizes - self.decoder = Decoder( latent_dim, decoder_dim, @@ -250,7 +245,7 @@ def encode( mode="nearest" ) resized_z = self.phis_downsample[i](resized_z) - z_q_i = self.quantizers[i].quantize(resized_z) + z_q_i = self.quantizers[i].quantize(resized_z) resized_z_q_i = F.interpolate( z_q_i, size=residual.shape[2], From f33a46223cdf4a50ac116b318600fbd4fe5e7ab1 Mon Sep 17 00:00:00 2001 From: joerenner Date: Tue, 18 Jun 2024 23:07:44 +0000 Subject: [PATCH 6/7] adding n_codes while encoding --- dac/model/dac.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/dac/model/dac.py b/dac/model/dac.py index 4816035..fbe7614 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -203,7 +203,8 @@ def preprocess(self, audio_data, sample_rate): def encode( self, - audio_data: torch.Tensor + audio_data: torch.Tensor, + n_codes: int = None ): """Encode given audio data and return quantized latent codes @@ -238,7 +239,9 @@ def encode( residual = z quantized = 0 quantized_list = [] - for i in range(len(self.downsample_rates)): + if n_codes is None: + n_codes = len(self.downsample_rates) + for i in range(n_codes): resized_z = F.interpolate( residual, size=int(residual.shape[2] / self.downsample_rates[i]), @@ -286,7 +289,8 @@ def decode(self, quantized): def forward( self, audio_data: torch.Tensor, - sample_rate: int = None + sample_rate: int = None, + n_codes: int = None ): """Model forward pass @@ -322,7 +326,7 @@ def forward( length = audio_data.shape[-1] audio_data = self.preprocess(audio_data, sample_rate) quantized, identity_loss = self.encode( - audio_data + audio_data, n_codes ) x = self.decode(quantized) return { From 29b00e760c0841655059f0217be191e49f431616 Mon Sep 17 00:00:00 2001 From: joerenner Date: Thu, 27 Jun 2024 18:37:33 +0000 Subject: [PATCH 7/7] adding n_codes in forward, shared quantizer --- conf/var.yml | 13 +++++++------ dac/model/dac.py | 30 +++++++++++++++++++----------- 2 files changed, 26 insertions(+), 17 deletions(-) diff --git a/conf/var.yml b/conf/var.yml index 3d0784c..1c5ccba 100644 --- a/conf/var.yml +++ b/conf/var.yml @@ -4,11 +4,12 @@ DAC.encoder_dim: 256 DAC.encoder_rates: [2, 4, 8, 8] DAC.decoder_dim: 1536 DAC.decoder_rates: [8, 8, 4, 2] -DAC.latent_dim: 768 +DAC.latent_dim: 1024 # Quantization -DAC.downsample_rates: [8, 6, 5, 4, 3, 2, 1, 1] +DAC.downsample_rates: [8, 8, 4, 4, 2, 2, 1, 1] DAC.channel_floats: [8, 8, 8, 5, 5, 5] +DAC.quantizer_dropout: 0.5 # Discriminator Discriminator.sample_rate: 44100 @@ -30,8 +31,8 @@ ExponentialLR.gamma: 0.999996 amp: false val_batch_size: 50 device: cuda -num_iters: 400000 -save_iters: [10000, 50000, 100000, 200000] +num_iters: 500000 +save_iters: [250000] valid_freq: 1000 sample_freq: 10000 num_workers: 32 @@ -67,11 +68,11 @@ MelSpectrogramLoss.clamp_eps: 1.0e-5 MelSpectrogramLoss.mag_weight: 0.0 # Data -batch_size: 16 +batch_size: 8 train/AudioDataset.duration: 0.4 train/AudioDataset.n_examples: 10000000 -val/AudioDataset.duration: 2.5 +val/AudioDataset.duration: 1.5 val/build_transform.augment_prob: 1.0 val/AudioDataset.n_examples: 250 diff --git a/dac/model/dac.py b/dac/model/dac.py index fbe7614..9fde5d5 100644 --- a/dac/model/dac.py +++ b/dac/model/dac.py @@ -155,7 +155,8 @@ def __init__( decoder_rates: List[int] = [8, 8, 4, 2], sample_rate: int = 44100, channel_floats: list = [8, 8, 8, 5, 5, 5], - downsample_rates: list = [2] + downsample_rates: list = [2], + quantizer_dropout: float = 0.0 ): super().__init__() self.encoder_dim = encoder_dim @@ -175,13 +176,8 @@ def __init__( self.encoder = Encoder(encoder_dim, encoder_rates, latent_dim) self.phis_downsample = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(len(self.downsample_rates))]) self.phis_upsample = nn.ModuleList([WNConv1d(latent_dim, latent_dim, kernel_size=3, padding="same") for _ in range(len(self.downsample_rates))]) - self.quantizers = nn.ModuleList( - [ - FSQ(channel_floats, latent_dim) - for i in range(len(self.downsample_rates)) - ] - ) - + self.quantizer = FSQ(channel_floats, latent_dim) + self.quantizer_dropout = quantizer_dropout self.decoder = Decoder( latent_dim, decoder_dim, @@ -241,6 +237,15 @@ def encode( quantized_list = [] if n_codes is None: n_codes = len(self.downsample_rates) + if self.training: + n_quantizers = torch.ones((z.shape[0],)) * len(self.downsample_rates) + 1 + dropout = torch.randint(1, len(self.downsample_rates) + 1, (z.shape[0],)) + n_dropout = int(z.shape[0] * self.quantizer_dropout) + n_quantizers[:n_dropout] = dropout[:n_dropout] + n_quantizers = n_quantizers.to(z.device) + else: + n_quantizers = n_codes + for i in range(n_codes): resized_z = F.interpolate( residual, @@ -248,17 +253,20 @@ def encode( mode="nearest" ) resized_z = self.phis_downsample[i](resized_z) - z_q_i = self.quantizers[i].quantize(resized_z) + z_q_i = self.quantizer.quantize(resized_z) resized_z_q_i = F.interpolate( z_q_i, size=residual.shape[2], mode="nearest" ) resized_z_q_i = self.phis_upsample[i](resized_z_q_i) + mask = ( + torch.full((residual.shape[0],), fill_value=i, device=residual.device) < n_quantizers + ) + residual -= resized_z_q_i - quantized += resized_z_q_i + quantized += resized_z_q_i * mask[:, None, None] quantized_list.append(quantized) - # indices.append(indices_i) identity_loss = 0.0 z_detach = z.detach()