# This file takes audio/audios and encode them through the codec to generate codec labels
import os
import argparse

from suno_utils.audio import Audio
from suno_utils.tasks import dac
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_speech"
OUT_AUDIO_DIR = os.path.join(OUT_DATA_DIR, "audio")
OUT_AUDIO_DEMUC_DIR = os.path.join(OUT_DATA_DIR, "audio_demuc")
OUT_AUDIO_DEMUC_LABEL_DIR = os.path.join(OUT_DATA_DIR, "audio_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"
CODEC_CKPT_PATH = "/home/tony/Data/MERT/codec.pt"
MAX_SHAPE = 300

DATASET_TYPE = "valid"


def encode_audio_file(input_audio_file):
    # pre-pend demuc dir
    input_audio_file_path = os.path.join(OUT_AUDIO_DIR, input_audio_file)
    input_audio_raw_name = input_audio_file.replace(".wav", "")
    output_file_name = os.path.join(
        OUT_AUDIO_DEMUC_LABEL_DIR, f"{DATASET_TYPE}_{input_audio_raw_name}.npy"
    )
    if os.path.exists(output_file_name):
        return
    # print(input_audio_files)
    audio = Audio.from_file(input_audio_file_path)
    output_arrays = dac.encode(audio)
    #    output_arrays = dac.encode_files(input_audio_files, batch_size=BATCH_SIZE, n_gpus=1)
    # print(output_arrays)
    #    output_arrays = output_arrays[0].astype(np.int16)
    #     final_arrays = np.zeros((len(input_audio_files), MAX_SHAPE, 8), dtype=np.int16)
    #     for n_row, arr in enumerate(output_arrays):
    #         if arr is None:
    #             print(f"WTF arr is None, {n_row}, {input_audio_files[n_row]}")
    #             raise ValueError
    #             # this job will not finish!
    #             return
    #             # arr = np.ones((MAX_SHAPE, 8), dtype=np.int16)
    #             # arr *= -1
    #         arr = arr.astype(np.int16)
    #         if arr.shape[0] < MAX_SHAPE:
    #             arr = np.pad(
    #                 arr,
    #                 ((0, MAX_SHAPE - arr.shape[0]), (0, 0)),
    #                 "constant",
    #                 constant_values=-1,
    #             )
    #         final_arrays[n_row, :] = arr.astype(np.int16)
    if os.path.exists(output_file_name):
        print(f"{output_file_name} is already written")
        return
    np.save(
        os.path.join(
            OUT_AUDIO_DEMUC_LABEL_DIR, f"{DATASET_TYPE}_{input_audio_raw_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)
    dac.load_model(checkpoint_filepath=CODEC_CKPT_PATH, device="cuda")
    # load current data
    input_args = parse_args()
    tsv_info = []
    # check GPU
    with open(os.path.join(OUT_TSV_DIR, f"{DATASET_TYPE}.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]  # .replace(".wav", "_vocals.wav")
        for x in tsv_data[input_args.start_index : input_args.end_index]
    ]
    #     batches = [
    #         work_items[x : x + BATCH_SIZE] for x in range(0, len(work_items), BATCH_SIZE)
    #     ]
    print(len(work_items), "work_items")  # , len(batches), "batches")
    thread_map(encode_audio_file, work_items, max_workers=4, chunksize=1)
    # thread_map(encode_audio_file, batches, max_workers=8, chunksize=1)
    # for batch in batches:
    #     encode_audio_file(batch)
    print(f"Done~!, {time() - start_time:.3f} seconds.")
