import os

os.environ["CUDA_VISIBLE_DEVICES"] = ""
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
from tqdm.contrib.concurrent import process_map, thread_map
import re

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

with open("hoot/metas_map.json", "r") as fp:
    metas_map = json.load(fp)
with open("hoot/sample_data.json", "r") as fp:
    sample_data = json.load(fp)

def generate_lyrics_data(i):
    if i % 10 != 0:
        data_split = "train"
    else:
        if i % 20 == 0:
            data_split = "dev"
        else:
            data_split = "test"
    k, v = sample_data[i]
    meta = metas_map[k]
    # audio = Audio.from_s3(meta["audio_filepath"])
    data_path = os.path.join(f"/home/tony/Data/Lyrics/{data_split}", f"{i}", "0")
    file_name = f"{i}-0"
    os.makedirs(data_path, exist_ok=True)
    text_lines = []
    for j, s in enumerate(v):
        # print(s["text"], s["start_s"], s["end_s"], len(tokenize(s["text"], max_tokens=512*8)))
        # segment = audio.get_segment(s["start_s"], s["end_s"])
        # segment.to_wav(os.path.join(data_path, f"{file_name}-{audio_name}.wav"))
        audio_name = f"00{str(j) if j > 10 else '0' + str(j)}"
        original_text = s['text']
        # remove sq brackets
        original_text = re.sub("([\[]).*?([\]])", "", original_text)
        # remove digits
        original_text = re.sub(r"\d+", "", original_text)
        cleaned_text = " ".join(re.findall(r"([\w\']+)", original_text)).upper()
        text_lines.append(f"{file_name}-{audio_name} {cleaned_text} \n")
    with open(os.path.join(data_path, f"{file_name}.trans.txt"), "w") as fp:
        fp.writelines(text_lines)

print("START")
thread_map(generate_lyrics_data, list(range(len(sample_data))), max_workers=40, chunksize=1)
print("DONE")