import os
import argparse

from suno_utils.tasks import demucs
from suno_utils.audio import Audio
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

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_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")


def demuc_audio_file(input_audio_file):
    output_vocal_file_name = os.path.join(
        OUT_AUDIO_DEMUC_DIR, input_audio_file.replace(".wav", "_vocals.wav")
    )
    if os.path.exists(output_vocal_file_name):
        return
    audio = Audio.from_file(os.path.join(OUT_AUDIO_DIR, input_audio_file))
    try:
        vocals, non_vocals = demucs.split_vocals(audio)
        vocals.to_wav(output_vocal_file_name)
        non_vocals.to_wav(
            os.path.join(
                OUT_AUDIO_DEMUC_DIR, input_audio_file.replace(".wav", "_bkg.wav")
            )
        )
    except Exception as E:
        print(E, input_audio_file)
        audio.to_wav(output_vocal_file_name)
        audio.to_wav(
            os.path.join(
                OUT_AUDIO_DEMUC_DIR, input_audio_file.replace(".wav", "_bkg.wav")
            )
        )
    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)
    demucs.load_model(device="cuda")
    # load current data
    input_args = parse_args()
    tsv_info = []
    # check GPU
    with open(os.path.join(OUT_TSV_DIR, "train.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(demuc_audio_file, work_items, max_workers=10, chunksize=1)
    print(f"Done~!, {time() - start_time:.3f} seconds.")
