import math import numpy as np from torch import nn from loss import stft_loss_rand, stft_loss from modules.seanet import SEANetEncoder, SEANetDecoder from quantization.vq import ResidualVectorQuantizer SAMPLE_RATE = 48_000 class CirceNet(nn.Module): def __init__( self, dimension=128, n_filters=64, ratios=(8, 5, 4, 4), causal=False, n_codebooks=4, sample_rate=SAMPLE_RATE, skip_quantization=False, ): super().__init__() self.encoder = SEANetEncoder( channels=1, norm="weight_norm", causal=causal, dimension=dimension, n_filters=n_filters, ratios=ratios, true_skip=True, n_residual_layers=1, lstm=2, ) if skip_quantization: self.quantizer = None else: self.quantizer = ResidualVectorQuantizer( dimension=dimension, n_q=n_codebooks, bins=2048, kmeans_iters=50, ) self.decoder = SEANetDecoder( channels=1, norm="weight_norm", causal=causal, dimension=dimension, n_filters=n_filters, ratios=ratios, true_skip=True, n_residual_layers=1, lstm=2, ) self.sample_rate = sample_rate self.frame_rate = math.ceil(self.sample_rate / np.prod(self.encoder.ratios)) print("number of parameters: %.2fM" % (self.get_num_params()/1e6,)) def forward(self, x, y=None, randomize_stft=False): assert x.dim() == 3 length = x.shape[-1] # (B, n_chan, T) emb = self.encoder(x) if self.quantizer is not None: q_res = self.quantizer(emb, self.frame_rate) quant = q_res.x comm_loss = q_res.penalty else: quant = emb comm_loss = None # (B, emb_dim, T*) y_pred = self.decoder(quant) # remove extra padding added by the encoder and decoder assert(y_pred.shape[-1] >= length) y_pred = y_pred[..., :length] if y is not None: stft_f = stft_loss_rand if randomize_stft else stft_loss sc_loss, mag_loss = stft_f(y_pred[:, 0, :], y[:, 0, :]) loss = { "sc_loss": sc_loss, "mag_loss": mag_loss, "comm_loss": comm_loss, } else: loss = None return y_pred, loss def get_num_params(self): n_params = sum(p.numel() for p in self.parameters()) return n_params def estimate_mfu(self, fwdbwd_per_iter, dt): """ estimate model flops utilization (MFU) in arbitrary units""" return fwdbwd_per_iter / dt / 10.0