import os
import sys
import json
import torch
import argparse
import librosa
import pandas as pd
import numpy as np
from tqdm.contrib.concurrent import thread_map
from datetime import datetime
from suno_utils.gpt.chirp_v2_5 import (
    preload_codec_models,
    codec_decode,
)

# sys.path.append("/home/minz/neon/ditto-cover")
# from ditto.models.ditto import Ditto
sys.path.append("/home/minz/neon/ditto-training")
from ditto_v2.models.ditto import Ditto


def get_emb(wav, task="cover"):
    # resample
    wav = librosa.resample(wav[: 48000 * 360], orig_sr=48000, target_sr=24000)
    expected_len = 24000 * 15
    if len(wav) < expected_len:
        wav = np.pad(wav, (0, expected_len - len(wav)), mode="constant")

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

    # get embedding
    inp = torch.tensor(chunks).unsqueeze(1).cuda()
    #
    # self_sim: audio, random aug, two random aug for same song
    # artist_sim: audio, multiple artists, multiple songs, 2 for each artist,
    # alubm_sim: audio, same album
    # artist_vox_sim: two song per artist, source sep vocal, sim
    # genre_sim: genre_text to audio
    if task == "cover":
        emb = ditto.music_to_latent(inp, "self_sim")
    elif task == "artist":
        emb = ditto.music_to_latent(inp, "artist_sim")

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


def get_similarity(input_job):
    s3_id, task = input_job
    assert task in ["cover", "artist"]
    try:
        NPZ_DIR = "/app/suno/data/dpo/30b_npz"
        local_path = (
            f"{NPZ_DIR}/{s3_id + ('_gen_cycle' if 'cycle' in NPZ_DIR else '')}.npz"
        )
        if not os.path.exists(local_path):
            raise ValueError()
        try:
            temp_npz = np.load(local_path)
            if "v4.0_raw" in temp_npz:
                arr = temp_npz["v4.0_raw"]
            elif "v3.5_raw" in temp_npz:
                if "cycle" not in NPZ_DIR:
                    print(f"weird, {local_path}, with only v3.5")
                arr = temp_npz["v3.5_raw"]
            elif "v3.0_raw" in temp_npz:
                if "cycle" not in NPZ_DIR:
                    print(f"weird, {local_path}, with only v3.0")
                arr = temp_npz["v3.0_raw"]
            else:
                raise ValueError()
        except Exception as e:
            print(local_path)
            raise eval
        assert arr.shape[1] == 13
        if task == "cover":
            prompt_arr = temp_npz["cover_arr"]
        elif task == "artist":
            prompt_arr = temp_npz["artist_arr"]
        assert prompt_arr.shape[1] == 13

        try:
            seed_codec = prompt_arr[:, 1:]
            seed_audio = codec_decode(torch.tensor(seed_codec).long())
            seed_emb = get_emb(seed_audio.array_float.mean(axis=0), task=task)
        except Exception as e:
            print("error in seed embedding --> ", e)
            return s3_id, 2

        # seed_codec = prompt_arr[:, 1:]
        # seed_audio = codec_decode(torch.tensor(seed_codec).long())
        # seed_emb = get_emb(seed_audio.array_float.mean(axis=0))

        codec_arr = arr[:, 1:]
        audio = codec_decode(torch.tensor(codec_arr).long())
        audio_emb = get_emb(audio.array_float.mean(axis=0))
        output_npz_path = f"/app/suno/data/dpo/ditto/{s3_id}_ditto_emb.npz"
        np.savez(output_npz_path, audio_emb=audio_emb, seed_emb=seed_emb)
        sim = np.dot(seed_emb, audio_emb)
        return s3_id, float(sim)
    except Exception as e:
        print("error in processing --> ", e)
        return s3_id, 2


if __name__ == "__main__":
    print(f"working with GPU:{os.environ['CUDA_VISIBLE_DEVICES']} ")
    parser = argparse.ArgumentParser()
    parser.add_argument("--task_index", type=int, default=0)
    parser.add_argument("--max_index", type=int, default=4)
    args = parser.parse_args()

    # Ensure task_index is valid
    if args.task_index < 0 or args.task_index >= args.max_index:
        raise ValueError(f"task_index must be between 0 and {args.max_index - 1}")

    # load models
    print("loading models...")
    _ = preload_codec_models("/app/suno/models/chirp_v2/dac_2c_25x12.pt", device="cuda")
    ditto = Ditto(
        latent_dim=128,
        model_path="/home/minz/logs/ditto_v2_local_8gpu_cont/ditto_v2_epoch_57.pt",
        is_flash=False,
    )
    ditto = ditto.eval()
    ditto = ditto.cuda()
    print("Finished loading models!")

    # read data
    print("loading data...")
    input_df = pd.read_pickle(
        "/home/tony/Data/Preference/30b_v6/interesting_clips_v4_h_t_6_20250115_full.pkl"
    )
    cover_ids = sorted(list(input_df[input_df["task"] == "cover"]["s3_id"].unique()))
    artist_ids = sorted(
        list(input_df[input_df["task"] == "artist_consistency"]["s3_id"].unique())
    )
    print(f"Total cover_ids: {len(cover_ids)}")
    print(f"Total artist_ids: {len(artist_ids)}")
    input_tasks = []
    for cover_id in cover_ids:
        input_tasks.append((cover_id, "cover"))
    for artist_id in artist_ids:
        input_tasks.append((artist_id, "artist"))

    # Split the data into chunks
    chunk_size = len(input_tasks) // args.max_index
    start_index = args.task_index * chunk_size
    end_index = (
        start_index + chunk_size
        if args.task_index < args.max_index - 1
        else len(input_tasks)
    )
    chunk_input_ids = input_tasks[start_index:end_index]
    print(
        f"Processing chunk {args.task_index + 1}/{args.max_index} with {len(chunk_input_ids)} ids"
    )
    print("done!")

    print("get similarity...")
    num_threads = 10
    today_str = datetime.now().strftime("%Y%m%d")
    output_json_path = os.path.join(
        f"/home/tony/Data/Preference/30b_v6/similarity_{today_str}_chunk{args.task_index + 1}of{args.max_index}.json"
    )
    analyzed_outputs = thread_map(
        get_similarity,
        chunk_input_ids,
        max_workers=num_threads,
        chunksize=1,
    )
    out = {input_id: sim for input_id, sim in analyzed_outputs if sim is not None}
    with open(output_json_path, "w") as fp:
        json.dump(out, fp)
    print(f"done! Output saved to {output_json_path}")
