import os
import sys

sys.path.append("/home/minz/neon/ditto-training")

import tqdm
import random
import glob
import torch
import numpy as np
from torch.utils import data
from suno_utils.audio import Audio
from einops import rearrange

from ditto_v2.data_loaders.self_sim import SelfSimDataset
from ditto_v2.data_loaders.self_vox_sim import SelfVoxSimDataset
from ditto_v2.data_loaders.self_lyric_sim import SelfLyricSimDataset
from ditto_v2.data_loaders.artist_sim import ArtistSimDataset
from ditto_v2.data_loaders.artist_vox_sim import ArtistVoxSimDataset
from ditto_v2.data_loaders.album_sim import AlbumSimDataset
from ditto_v2.data_loaders.genre_sim import GenreSimDataset
from ditto_v2.data_loaders.lyric_sim import LyricSimDataset
from ditto_v2.models.ditto import Ditto

from suno_utils.utils.text import read_jsonl, read_json


split = "valid"
self_sim_input_length_s = 15.0
self_vox_sim_input_length_s = 15.0
self_lyric_sim_input_length = 300
artist_sim_input_length_s = 15.0
artist_vox_sim_input_length_s = 15.0
album_sim_input_length_s = 15.0
genre_sim_input_length_s = 15.0
genre_tag_input_length = 300
lyric_sim_input_length_s = 15.0
lyric_sim_input_text_length = 300
sample_rate = 24000
num_samples = 1000
num_self_sim_samples = -1
num_vox_sim_samples = -1
num_self_lyric_sim_samples = -1
num_artist_sim_samples = -1
num_artist_vox_sim_samples = -1
num_album_sim_samples = -1
num_genre_sim_samples = -1
num_lyric_sim_samples = -1
metadata = read_jsonl("/app/suno/data/v2_audio/metadata/metas_%s.jsonl" % split)
task_indices = read_json("/app/suno/data/v2_audio/metadata/task_indices.json")

# get model
ditto = Ditto(
    latent_dim=128,
    model_path="/home/minz/logs/ditto_v2_local_8gpu/epoch=29.ckpt",
    is_flash=False,
)
ditto = ditto.eval()
ditto = ditto.cuda()

# preprocess data
(
    self_sim_emb,
    artist_sim_emb,
    album_sim_emb,
    genre_sim_emb,
    self_vox_sim_emb,
    artist_vox_sim_emb,
) = [], [], [], [], [], []


class GeniusDataset(data.Dataset):
    def __init__(self, filelist, sample_rate=24000):
        self.filelist = filelist
        self.sample_rate = sample_rate

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

    def __getitem__(self, idx):
        filepath = self.filelist[idx]
        mix_tensor = self.get_tensor(filepath)
        vox_filepath = os.path.join(
            "/app/suno/data/v2_audio/genius_vox/", os.path.basename(filepath)
        )
        vox_tensor = self.get_tensor(vox_filepath)
        return mix_tensor, vox_tensor

    def get_tensor(self, filepath):
        inp_wav = Audio.from_file(filepath, sample_rate=self.sample_rate)
        inp_array = inp_wav.array_float
        total_duration = len(inp_array)
        chunk_duration = self.sample_rate * 15

        # Calculate 4 evenly spaced start points
        start_points = [
            int(i * (total_duration - chunk_duration) / 3) for i in range(4)
        ]

        chunks = [inp_array[start : start + chunk_duration] for start in start_points]
        return torch.tensor(np.vstack(chunks))


def get_emb(filepath, task):
    inp = self.get_tensor(filepath)
    out = ditto.music_to_latent(inp.cuda(), task)
    out = out.mean(dim=0).detach().cpu().numpy()
    return out


genius_filelist = glob.glob("/app/suno/data/v2_audio/genius_hq/*.wav")
random.seed(134)
random.shuffle(genius_filelist)
genius_filelist = genius_filelist[:10000]
dataset = GeniusDataset(genius_filelist)
batch_size = 64
dataloader = data.DataLoader(
    dataset, batch_size=batch_size, num_workers=4, pin_memory=True
)

for mix_batch, vox_batch in tqdm.tqdm(dataloader):
    mix_batch = mix_batch.cuda()
    vox_batch = vox_batch.cuda()

    mix_batch = rearrange(mix_batch, "b c t -> (b c) t")
    vox_batch = rearrange(vox_batch, "b c t -> (b c) t")

    # Process mix
    out = ditto.music_to_latent(mix_batch, "self_sim")
    out = rearrange(out, "(b c) e -> b c e", b=batch_size)
    out = out.mean(dim=1).detach().cpu().numpy()
    self_sim_emb.append(out)

    out = ditto.music_to_latent(mix_batch, "artist_sim")
    out = rearrange(out, "(b c) e -> b c e", b=batch_size)
    out = out.mean(dim=1).detach().cpu().numpy()
    artist_sim_emb.append(out)

    out = ditto.music_to_latent(mix_batch, "album_sim")
    out = rearrange(out, "(b c) e -> b c e", b=batch_size)
    out = out.mean(dim=1).detach().cpu().numpy()
    album_sim_emb.append(out)

    out = ditto.music_to_latent(mix_batch, "genre_sim")
    out = rearrange(out, "(b c) e -> b c e", b=batch_size)
    out = out.mean(dim=1).detach().cpu().numpy()
    genre_sim_emb.append(out)

    # Process vox
    out = ditto.music_to_latent(vox_batch, "self_vox_sim")
    out = rearrange(out, "(b c) e -> b c e", b=batch_size)
    out = out.mean(dim=1).detach().cpu().numpy()
    self_vox_sim_emb.append(out)

    out = ditto.music_to_latent(vox_batch, "artist_vox_sim")
    out = rearrange(out, "(b c) e -> b c e", b=batch_size)
    out = out.mean(dim=1).detach().cpu().numpy()
    artist_vox_sim_emb.append(out)

# Concatenate results
self_sim_emb = np.concatenate(self_sim_emb)
artist_sim_emb = np.concatenate(artist_sim_emb)
album_sim_emb = np.concatenate(album_sim_emb)
genre_sim_emb = np.concatenate(genre_sim_emb)
self_vox_sim_emb = np.concatenate(self_vox_sim_emb)
artist_vox_sim_emb = np.concatenate(artist_vox_sim_emb)

np.save(open("/home/minz/temp/genius_self_sim_emb.npy", "wb"), self_sim_emb)
np.save(open("/home/minz/temp/genius_artist_sim_emb.npy", "wb"), artist_sim_emb)
np.save(open("/home/minz/temp/genius_album_sim_emb.npy", "wb"), album_sim_emb)
np.save(open("/home/minz/temp/genius_genre_sim_emb.npy", "wb"), genre_sim_emb)
np.save(open("/home/minz/temp/genius_self_vox_sim_emb.npy", "wb"), self_vox_sim_emb)
np.save(open("/home/minz/temp/genius_artist_vox_sim_emb.npy", "wb"), artist_vox_sim_emb)
