import os
import sys
import json
import torch
import tqdm
import argparse
import librosa
import numpy as np
from tqdm.contrib.concurrent import thread_map
from suno_utils.utils.text import read_json, read_jsonl
from suno_utils.audio import Audio

from suno_utils.gpt.chirp_v2_5 import (
    preload_semantic_models,
    preload_codec_models,
    _get_model_if_needed,
    GenerationConfig,
    generate,
    semantic_encode,
    codec_encode,
    codec_decode_stream_to_full_audio,
    codec_decode
)
sys.path.append("/home/minz/neon/ditto-cover")
from ditto.models.ditto import Ditto


def get_emb(wav):
    # resample
    wav = librosa.resample(wav[:48000*360], orig_sr=48000, target_sr=24000)

    # make chunks
    chunks = librosa.util.frame(wav, frame_length=24000 * 15, hop_length=24000 * 15)
    chunks = chunks.transpose(1, 0)

    # get embedding
    inp = torch.tensor(chunks).unsqueeze(1).cuda()
    emb = ditto.music_to_latent(inp)

    return emb.mean(dim=0).detach().cpu().numpy()

def get_cover_similarity(ix):
    try:
        seed = keys[ix]
        print('processing ', seed)
        seed_codec = data[int(seed)][:, 1:]
        seed_codec = seed_codec[:np.where(seed_codec == 2048)[0][0]]
        seed_audio = codec_decode(torch.tensor(seed_codec).long())
        seed_emb = get_emb(seed_audio.array_float.mean(axis=0))
    except Exception as e:
        print("error in seed", e)
        return None, None
    
    sims = []
    covers = info[seed]
    for cover in covers:
        try:
            cover_codec = data[int(cover)][:, 1:]
            cover_codec = cover_codec[:np.where(cover_codec == 2048)[0][0]]
            cover_audio = codec_decode(torch.tensor(cover_codec).long())
            cover_emb = get_emb(cover_audio.array_float.mean(axis=0))
            sim = np.dot(seed_emb, cover_emb) / (np.linalg.norm(seed_emb) * np.linalg.norm(cover_emb))
            sims.append({"similarity": sim, "seed": seed, "cover": cover})
        except Exception as e:
            print("error in cover", e)
            continue
    print('completed ', seed)
    return seed, sims


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--trial_index", type=int, default=0)
    args = parser.parse_args()
    process_start_index = args.trial_index * 12000

    # load models
    print("loading models...")
    _ = preload_codec_models("/app/suno/models/chirp_v2/dac_2c_25x12.pt")
    ditto = Ditto(model_path="/home/minz/logs/ditto_cover/epoch40.ckpt", is_flash=False)
    ditto = ditto.eval()
    ditto = ditto.cuda()
    print("done!")

    # read data
    print("loading data...")
    SPLIT = "tr" # tr or val
    OUTPUT_PATH = "/app/suno/minz/cover_sim/" + SPLIT
    info = read_json("/app/suno/data/chirp_v4/multi/info_%s.json" % SPLIT)["discogs_covers"]["idx_map"]
    keys = list(info.keys())
    data = np.memmap("/app/suno/data/chirp_v4/multi/data_%s.bin" % SPLIT, dtype=np.uint16, mode="r")
    data = data.reshape(-1, 6016, 13)
    # metadata = read_jsonl("/app/suno/data/chirp_v4/multi/metas_%s.jsonl" % SPLIT)
    print("done!")

    print("get similarity...")
    num_threads = 8
    batch_size = 100
    

    for i in tqdm.tqdm(range(process_start_index, process_start_index + 12000, batch_size)):
        start_index = i
        end_index = min(start_index + batch_size, len(info))
        if start_index > len(info):
            continue
        output_json_path = os.path.join(OUTPUT_PATH, f"batch_{start_index}.json")
        
        with open(output_json_path, "w") as fp:
            fp.write("")

        analyzed_outputs = thread_map(
            get_cover_similarity,
            range(start_index, end_index),
            max_workers=num_threads,
            chunksize=1,
            disable=True
        )
        out = {seed: sims for seed, sims in analyzed_outputs if sims is not None}
        with open(output_json_path, "w") as fp:
            json.dump(out, fp, default=float)
    print("done!")