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~!")