import os
import argparse

from suno_utils.audio import Audio
from suno_utils.tasks import mert_25
from multiprocessing import Pool
from tqdm.contrib.concurrent import process_map, thread_map
from time import time
import numpy as np

SAMPLE_RATE = 24_000
EMBEDDING_RATE = 25
N_CODEBOOKS = 8
BATCH_SIZE = 8

OUT_DATA_DIR = "/app/suno/data/mert_25hz_short"
OUT_AUDIO_DIR = os.path.join(OUT_DATA_DIR, "audio")
OUT_AUDIO_DEMUC_DIR = os.path.join(OUT_DATA_DIR, "audio_demuc")
OUT_AUDIO_MERT_LABEL_DIR = os.path.join(OUT_DATA_DIR, "audio_mert_label")
OUT_TSV_DIR = os.path.join(OUT_DATA_DIR, "audio_tsv")
OUT_LABEL_DIR = os.path.join(OUT_DATA_DIR, "label")
OUT_TEMP_DIR = os.path.join(OUT_DATA_DIR, "temp")
# CODEC_CKPT_PATH = "s3://suno-data/georg/trained_models/chirp_v1/codec.pt"
MERT_CKPT_PATH = "/home/tony/Data/MERT/mert_test_8x_400k.pt"
MERT_CENTER_PATH = "/home/tony/Data/MERT/cluster_centers/default.npy"
MAX_SHAPE = 300  # 12 sec * 25 hz
DATASET = "train"


def encode_audio_file(input_audio_file):
    input_audio_path = os.path.join(OUT_AUDIO_DIR, input_audio_file)
    output_file_name = os.path.join(
        OUT_AUDIO_MERT_LABEL_DIR,
        f"{DATASET}_{input_audio_file.replace('.wav', '')}.npy",
    )
    if os.path.exists(output_file_name):
        return
    audio = Audio.from_file(input_audio_path)
    output_arrays = mert_25.encode(
        audio,
        do_clustering=True,
        n_codebooks=1,
    ).astype(np.int16)
    if os.path.exists(output_file_name):
        print(f"{output_file_name} is already written")
        return
    np.save(output_file_name, output_arrays)
    return


def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument("--start_index", type=int, default=0)
    parser.add_argument("--end_index", type=int, default=100)
    args = parser.parse_args()
    return args


if __name__ == "__main__":
    start_time = time()
    avialbe_device = f"cuda:{os.environ['CUDA_VISIBLE_DEVICES']}"
    print(avialbe_device)
    _ = mert_25.preload_models(
        checkpoint_filepath=MERT_CKPT_PATH,
        centroids_filepath=MERT_CENTER_PATH,
    )
    # load current data
    input_args = parse_args()
    tsv_info = []
    # check GPU
    with open(os.path.join(OUT_TSV_DIR, f"{DATASET}.tsv"), "r") as f:
        for line in f.read().strip().split("\n"):
            if len(line.strip()) == 0:
                continue
            tsv_info.append(line.strip().split("\t"))
    tsv_data_path = tsv_info[0]
    tsv_data = tsv_info[1:]
    print(len(tsv_data), input_args.start_index, input_args.end_index)
    work_items = [x[0] for x in tsv_data[input_args.start_index : input_args.end_index]]
    print(len(work_items), "work_items")
    thread_map(encode_audio_file, work_items, max_workers=8, chunksize=1)
    print(f"Done~!, {time() - start_time:.3f} seconds.")
