import random
import json
import numpy as np
import tqdm
import torch
import funcy
import time
import gc
from scipy.io import wavfile
import tempfile
import collections
from collections import defaultdict
from joblib import Parallel, delayed
import os
import argparse

from suno_utils.utils.s3 import _apply_mp
from suno_utils.audio import Audio
from suno_utils.tasks.data_loader import load_audio_mp
from suno_utils.utils.text import write_jsonl, read_jsonl, write_json, read_json
from suno_utils.utils.s3 import read_from_s3, check_s3_file_exists, open_from_s3
from suno_utils.audio.conversion import convert_audio_files

SAMPLE_RATE = 24_000
EMBEDDING_RATE = 25
N_CODEBOOKS = 8

IN_DATA_DIR = "/app/suno/data/mert_25hz"
IN_AUDIO_DIR = os.path.join(IN_DATA_DIR, "audio")
IN_TSV_DIR = os.path.join(IN_DATA_DIR, "audio_tsv")
IN_LABEL_DIR = os.path.join(IN_DATA_DIR, "label")

OUT_DATA_DIR = "/app/suno/data/mert_25hz_long"
OUT_AUDIO_DIR = os.path.join(OUT_DATA_DIR, "audio")
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")

from multiprocessing import Pool
from tqdm.contrib.concurrent import process_map, thread_map

import json

# write new audios to disk
def _write_new_wav_files(work_item):
    from_fn, to_work_items = work_item
    # print(from_fn, to_work_items)
    expected_first_output = os.path.join(OUT_AUDIO_DIR, to_work_items[0][0])
    if os.path.exists(expected_first_output):
        # print(from_fn, "is already done!")
        return 
    _, audio_arr = wavfile.read(os.path.join(IN_AUDIO_DIR, from_fn))
    for to_fn, (start_idx, end_idx) in to_work_items:
        wavfile.write(
            os.path.join(OUT_AUDIO_DIR, to_fn),
            SAMPLE_RATE,
            audio_arr[start_idx:end_idx],
        )
    del audio_arr

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__":
    with open("train_origin_speech.json", "r") as fp:
        new_tsv_data = json.load(fp)

    input_args = parse_args()
    work_items = new_tsv_data
    # work_items.sort()
    print("original total", len(work_items))
    work_items = work_items[input_args.start_index:input_args.end_index]
    print(len(work_items), "work_items", input_args.start_index)
    # _write_new_wav_files(work_items[0])
    # print("Finished try one")
    # process_map(_write_new_wav_files, work_items, max_workers=58, chunksize=10)
    # thread is like 5 times faster
    thread_map(_write_new_wav_files, work_items, max_workers=10, chunksize=1)
    # with Pool(64) as p:
    #      tqdm.tqdm(p.imap(_write_new_wav_files, work_items))
    print("Done~!")
