import json
import tqdm
import random
import torch
import numpy as np
from torch import nn
from einops import rearrange

from musicfm.modules.random_quantizer import RandomProjectionQuantizer
from musicfm.modules.features import MelSTFT, CQT
from musicfm.modules.conv import Conv2dSubsampling
from suno_utils.tasks.mert_25 import _array_chunk


class MusicFM_Dual(nn.Module):
    """
    MusicFM with melspec + CQT targets

    Input: 128-band mel spectrogram
    Frontend: 2-layer Residual convolution
    Backend: 12-layer Conformer
    Quantizer: 1 codebooks for mel spectrogram or spectrogram
    """

    def __init__(
        self,
        num_codebooks=1,
        codebook_dim=16,
        codebook_size=4096,
        features=["melspec_2048", "cqt"],
        hop_length=240,
        n_mels=128,
        conv_dim=512,
        encoder_dim=1024,
        encoder_depth=12,
        mask_hop=0.4,
        mask_prob=0.6,
        is_flash=False,
        is_cls=False,
        stat_path=None,
        model_path=None,
    ):
        super(MusicFM_Dual, self).__init__()

        # global variables
        self.hop_length = hop_length
        self.mask_hop = mask_hop
        self.mask_prob = mask_prob
        self.num_codebooks = num_codebooks
        self.codebook_size = codebook_size
        self.features = features

        # load feature mean / std stats
        with open(stat_path, "r") as f:
            self.stat = json.load(f)

        # feature extractor
        self.preprocessor_melspec_2048 = MelSTFT(
            n_fft=2048, hop_length=hop_length, is_db=True
        )
        self.preprocessor_cqt = CQT()

        # random quantizer
        seed = 142
        for feature in self.features:
            for i in range(num_codebooks):
                if feature != "cqt":
                    setattr(
                        self,
                        "quantizer_%s_%d" % (feature, i),
                        RandomProjectionQuantizer(
                            n_mels * 4, codebook_dim, codebook_size, seed=seed + i
                        ),
                    )
                else:
                    setattr(
                        self,
                        "quantizer_%s_%d" % (feature, i),
                        RandomProjectionQuantizer(
                            84 * 4, codebook_dim, codebook_size, seed=seed + i
                        ),
                    )

        # two residual convolution layers + one projection layer
        self.conv = Conv2dSubsampling(
            1, conv_dim, encoder_dim, strides=[2, 2], n_bands=n_mels
        )

        # Conformer
        if is_flash:
            from musicfm.modules.flash_conformer import (
                Wav2Vec2ConformerEncoder,
                Wav2Vec2ConformerConfig,
            )
        else:
            from transformers.models.wav2vec2_conformer.modeling_wav2vec2_conformer import (
                Wav2Vec2ConformerEncoder,
                Wav2Vec2ConformerConfig,
            )
        config = Wav2Vec2ConformerConfig.from_pretrained(
            "facebook/wav2vec2-conformer-rope-large-960h-ft"
        )
        config.num_hidden_layers = encoder_depth
        config.hidden_size = encoder_dim

        self.conformer = Wav2Vec2ConformerEncoder(config)

        # projection
        self.linear = nn.Linear(
            encoder_dim, num_codebooks * codebook_size * len(features)
        )

        # loss function
        self.loss = nn.CrossEntropyLoss()

        # cls token
        if is_cls:
            random.seed(seed)
            self.cls_token = nn.Parameter(torch.randn(encoder_dim))

        # load model
        if model_path:
            S = torch.load(model_path)["state_dict"]
            SS = {k[6:]: v for k, v in S.items()}
            self.load_state_dict(SS, strict=True)

    def masking(self, x):
        """random masking of 400ms with given probability"""
        mx = x.clone()
        b, t = mx.shape
        len_masking_raw = int(24000 * self.mask_hop)
        len_masking_token = int(24000 / self.hop_length / 2 / 2 * self.mask_hop)

        # get random mask indices
        start_indices = torch.rand(b, t // len_masking_raw) < self.mask_prob
        time_domain_masked_indices = torch.nonzero(
            start_indices.repeat_interleave(len_masking_raw, dim=1)
        )
        token_domain_masked_indices = torch.nonzero(
            start_indices.repeat_interleave(len_masking_token, dim=1)
        )

        # mask with random values
        masking_noise = (
            torch.randn(time_domain_masked_indices.shape[0], dtype=x.dtype) * 0.1
        )  # 0 mean 0.1 std
        mx[tuple(time_domain_masked_indices.t())] = masking_noise.to(x.device)

        return mx, token_domain_masked_indices

    @torch.no_grad()
    def preprocessing(self, x, features):
        """extract classic audio features"""
        # check precision
        if x.dtype == torch.float16:
            precision = 16
        elif x.dtype == torch.bfloat16:
            precision = "bf16"
        else:
            precision = 32

        out = {}
        for key in features:
            layer = getattr(self, "preprocessor_%s" % key)
            out[key] = layer.float()(x.float())[..., :-1]
            if precision == 16:
                out[key] = out[key].half()
            elif precision == "bf16":
                out[key] = out[key].bfloat16()
        return out

    def encoder(self, x, is_cls=False):
        """2-layer conv + w2v-conformer"""
        x = self.conv(x)
        if is_cls:
            cls_token = self.cls_token.repeat(x.shape[0], 1, 1)
            x = torch.cat((cls_token, x), dim=1)
        out = self.conformer(x, output_hidden_states=True)
        hidden_emb = out["hidden_states"]
        last_emb = out["last_hidden_state"]
        logits = self.linear(last_emb)
        logit_dict = {}
        ix = 0
        for key in self.features:
            for i in range(self.num_codebooks):
                logit_dict["%s_%d" % (key, i)] = logits[
                    :, :, ix * self.codebook_size : (ix + 1) * self.codebook_size
                ]
        return logit_dict, hidden_emb

    @torch.no_grad()
    def normalize(self, x):
        """normalize the input audio to have zero mean unit variance"""
        for key in x.keys():
            x[key] = (x[key] - self.stat["%s_mean" % key]) / self.stat["%s_std" % key]
        return x

    @torch.no_grad()
    def rearrange(self, x):
        """rearrange the batch to flatten every 4 steps"""
        for key in x.keys():
            if key == "chromagram":
                x[key] = rearrange(x[key], "b f t -> b t f")
            else:
                x[key] = rearrange(x[key], "b f (t s) -> b t (s f)", s=4)
        return x

    @torch.no_grad()
    def tokenize(self, x):
        out = {}
        for key in x.keys():
            for i in range(self.num_codebooks):
                layer = getattr(self, "quantizer_%s_%d" % (key, i))
                out["%s_%d" % (key, i)] = layer(x[key])
        return out

    def get_targets(self, x):
        x = self.preprocessing(x, features=self.features)
        x = self.normalize(x)
        x = self.rearrange(x)
        target_tokens = self.tokenize(x)
        return target_tokens

    def get_predictions(self, x, is_cls=False):
        # preprocessing
        x = self.preprocessing(x, features=["melspec_2048"])
        x = self.normalize(x)

        # encoding
        logits, hidden_emb = self.encoder(x["melspec_2048"], is_cls)

        return logits, hidden_emb

    def get_latent(self, x, layer_ix=12, is_cls=False):
        _, hidden_states = self.get_predictions(x, is_cls)
        emb = hidden_states[layer_ix]
        return emb

    @torch.no_grad()
    def encode_arrays(self, x, batch_size=16, layer_ix=7):
        # make batch
        subsplit_arrays = []
        for n_array, arr in enumerate(x):
            for sub_arr in _array_chunk(24000 * 30, arr, step_size=24000 * 20, dim=1):
                if sub_arr.size(-1) == 24000 * 30:
                    subsplit_arrays.append(sub_arr)

        # encode them
        encoded_arrays = []
        num_iter = len(subsplit_arrays) // batch_size
        if len(subsplit_arrays) % batch_size > 0:
            num_iter += 1
        for i in tqdm.tqdm(range(num_iter)):
            inp = torch.cat(
                subsplit_arrays[i * batch_size : (i + 1) * batch_size]
            ).cuda()
            emb = self.get_latent(inp, layer_ix)
            emb = rearrange(emb, "b t c -> (b t) c")
            encoded_arrays.append(emb.cpu().detach().numpy())
        return np.concatenate(encoded_arrays)

    def get_loss(self, logits, target_tokens, masked_indices):
        losses = {}
        accuracies = {}
        for key in logits.keys():
            masked_logits = logits[key][tuple(masked_indices.t())]
            masked_tokens = target_tokens[key][tuple(masked_indices.t())]
            losses[key] = self.loss(masked_logits, masked_tokens)
            accuracies[key] = (
                torch.sum(masked_logits.argmax(-1) == masked_tokens)
                / masked_tokens.numel()
            )
        return losses, accuracies

    def forward(self, x):
        # get target feature tokens
        target_tokens = self.get_targets(x)

        # masking
        x, masked_indices = self.masking(x)

        # forward
        logits, hidden_emb = self.get_predictions(x)

        # get loss
        losses, accuracies = self.get_loss(logits, target_tokens, masked_indices)

        return logits, hidden_emb, losses, accuracies
