import os
import glob
import glob
import torch
import uuid
import IPython
import numpy as np
import torchaudio
import itertools

from tqdm import tqdm
from dac.model.dac2 import DAC
from dac.model.discriminator2 import Discriminator as Discriminator_import
from dac.nn import loss as loss_import
from dac.utils.accelerator import Accelerator
from dac.utils import load_model

from suno_utils.models.musicfm.modeling_MusicFM import MusicFM_MERTLong

import matplotlib.pyplot as plt
from sklearn.preprocessing import StandardScaler
from sklearn.cluster import KMeans, MiniBatchKMeans

from suno_utils.utils.s3 import read_from_s3

USE_VAL = False
NUM_FRAMES = int(10 * 24_000)
SAMPLE_RATE = 24_000
BATCH_SIZE = 128


class AudioDataset(torch.utils.data.Dataset):
    def __init__(self):
        # get audio files
        if USE_VAL:
            audio_subsets = glob.glob(
                os.path.join(f"/app/suno/data/audio_2ch_24khz_lg/val/**")
            )
        else:
            audio_subsets = glob.glob(
                os.path.join(f"/app/suno/data/audio_2ch_24khz_lg/train/**")
            )

        audio_files = []
        for audio_subset in audio_subsets:
            # find the first MAX_FILES_PER_SUBSET files
            with os.scandir(audio_subset) as filepaths:
                # Use itertools.islice to limit the iterator to the first N entries
                first_n_files = list(itertools.islice(filepaths, 40000))
            first_n_files = [entry.path for entry in first_n_files if entry.is_file()]
            audio_files += first_n_files
            print(len(first_n_files), audio_subset)

        print("Total", len(audio_files))
        print(np.random.choice(audio_files, 5))

        self.audio_files = audio_files

    def __len__(self):
        return len(self.audio_files)

    def __getitem__(self, idx):
        audio_file = self.audio_files[idx]
        num_frames = torchaudio.info(audio_file).num_frames
        if num_frames > NUM_FRAMES:
            frame_offset = np.random.randint(
                0, torchaudio.info(audio_file).num_frames - NUM_FRAMES - 1
            )
        else:
            frame_offset = 0

        audio, sr = torchaudio.load(
            audio_file, frame_offset=frame_offset, num_frames=NUM_FRAMES
        )
        audio = audio.mean(dim=0)

        if audio.shape[-1] < NUM_FRAMES:
            audio = torch.cat(
                [audio, torch.zeros(NUM_FRAMES - audio.shape[-1])], dim=-1
            )

        assert sr == SAMPLE_RATE
        audio = audio / audio.abs().max().clamp(1e-8)
        return idx, audio


if __name__ == "__main__":

    # load MERT
    model_filepath = "s3://suno-data/minz/models/musicfm_concat_epoch=51.pt"
    centroids_filepath = "s3://suno-data/minz/models/musicfm_concat_centroids.npy"
    mert_model = MusicFM_MERTLong(
        is_flash=False,
        stat_path="s3://suno-data/minz/models/mertlong_stats.json",
        model_path=model_filepath,
    )
    mert_model.cuda()
    mert_model.eval()

    # setup data
    if USE_VAL:
        subset_name = "val"
    else:
        subset_name = "train"

    root_dir = "/app/suno/christian/data/mert/"
    out_dir = os.path.join(root_dir, subset_name)
    os.makedirs(out_dir, exist_ok=True)

    dataset = AudioDataset()
    dataloader = torch.utils.data.DataLoader(
        dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=32
    )

    # embed
    for batch in tqdm(dataloader):
        idx, audios = batch
        with torch.no_grad():
            embeddings = mert_model.get_latent(audios.cuda())

        # save embeddings as npz
        for i, (index, emb) in enumerate(zip(idx, embeddings)):

            np.savez(
                f"/app/suno/christian/data/mert/{subset_name}/{index:10d}.npz",
                emb.cpu().numpy(),
            )
