import os
import sys

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

import tqdm
import numpy as np
from torch.utils import data

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 dataloaders
self_sim_dataset = SelfSimDataset(
    metadata,
    task_indices,
    split,
    self_sim_input_length_s,
    sample_rate,
    num_self_sim_samples,
)
self_sim_dataloader = data.DataLoader(self_sim_dataset, batch_size=16, shuffle=False)

# get model
ditto = Ditto(
    latent_dim=128,
    model_path="/home/minz/logs/ditto_v2_local/epoch=40.pt",
    is_flash=True,
)
ditto = ditto.eval()
ditto = ditto.cuda()

# process self_sim
outs1, outs2 = [], []
for inp1, inp2 in tqdm.tqdm(self_sim_dataloader):
    inp1 = inp1.cuda()
    inp2 = inp2.cuda()
    out1, out2, _ = ditto(inp1, inp2, "self_sim")
    outs1.append(out1.detach().cpu().numpy())
    outs2.append(out2.detach().cpu().numpy())
outs1 = np.concatenate(outs1, axis=0)
outs2 = np.concatenate(outs2, axis=0)
