from copy import deepcopy
import json
import math
import os
import soxr
import random
import traceback
import contextlib
from tqdm import tqdm

import numpy as np
import torch
import torch.nn.functional as F
from torch.utils.data import IterableDataset

from pretrained_models import load_vae_model, preload_mert_models, preload_musicfm_models
from audioloader import AudioConfig

from suno_utils.tasks.dac_vae_fixed_25hz import encode as encode_vae
from suno_utils.tasks.mert_25 import encode as encode_mert
from suno_utils.tasks.musicfm_v3 import encode as encode_musicfm


def get_batch(
    batch_size: int,
    audio_cfg: AudioConfig,
    sample_generator,
):
    cur_len = 0
    stacked_batch = []
    for sample in tqdm(sample_generator, desc="stacking batch", disable=True):
        if cur_len >= batch_size:
            break
        stacked_batch.append(sample)
        cur_len += 1

    return stacked_batch


def _resample_for_mert(arr):
    assert arr.ndim == 2
    assert arr.shape[0] == 2
    out_arr = soxr.resample(arr.T, 48_000, 24_000).mean(axis=1).astype(np.float32)
    return out_arr[np.newaxis, :]


def _resample_for_musicfm(arr):
    assert arr.ndim == 2
    assert arr.shape[0] == 2
    out_arr = soxr.resample(arr.T, 48_000, 16_000).astype(np.float32)
    return out_arr.T


class PreprocessDataset(IterableDataset):
    def __init__(
        self,
        sample_data_dl,
        batch_size: int,
        audio_cfg: AudioConfig,
    ):
        self.sample_data_dl = sample_data_dl
        self.batch_size = batch_size
        self.audio_cfg = audio_cfg

    def sample_generator_fn(self, sample_data_dl):
        while True:
            try:
                sample_data = next(sample_data_dl)
                yield from self.make_sample(sample_data)
            except Exception as e:
                print(f"Error in sample_generator_fn: {e}")
                print(traceback.format_exc())

    def make_sample(self, sample_data):
        def get_vae(track, vae_scale_factor=0.4):
            load_vae_model()
            wav = torch.from_numpy(track).cuda()
            vae = encode_vae(wav)
            vae = torch.from_numpy(vae) * vae_scale_factor
            return vae.T  # (C, T)

        def get_mert(track):
            preload_mert_models()
            track = _resample_for_mert(track)
            track = torch.from_numpy(track)
            embs = encode_mert(
                track,
                pad_to_chunksize=False,
                batch_size=48,
                do_clustering=False,
            ).T  # (C, T)
            return embs

        def get_musicfm(track):
            preload_musicfm_models()
            track = _resample_for_musicfm(track)
            track = torch.from_numpy(track)
            with open(os.devnull, "w") as devnull:
                with contextlib.redirect_stdout(devnull):
                    embs = encode_musicfm(
                        track, pad_to_chunksize=False, batch_size=48, token_type="emb_pre"
                    ).T  # (C, T)
            return embs

        # process data
        if self.audio_cfg.is_vae:
            sample_data.data_vae = get_vae(sample_data.data_wav)
        if self.audio_cfg.is_mert:
            sample_data.data_mert = get_mert(sample_data.data_wav)
        if self.audio_cfg.is_musicfm:
            sample_data.data_musicfm = get_musicfm(sample_data.data_wav)

        yield sample_data

    def __iter__(self):
        self.sample_generator = self.sample_generator_fn(self.sample_data_dl)
        return self

    def __next__(self):
        batch = get_batch(
            self.batch_size,
            self.audio_cfg,
            self.sample_generator,
        )
        return batch
