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
