import os import argparse from suno_utils.audio import Audio from suno_utils.tasks import mert_25 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_short" OUT_AUDIO_DIR = os.path.join(OUT_DATA_DIR, "audio") OUT_AUDIO_DEMUC_DIR = os.path.join(OUT_DATA_DIR, "audio_demuc") OUT_AUDIO_MERT_LABEL_DIR = os.path.join(OUT_DATA_DIR, "audio_mert_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" MERT_CKPT_PATH = "/home/tony/Data/MERT/mert_test_8x_400k.pt" MERT_CENTER_PATH = "/home/tony/Data/MERT/cluster_centers/default.npy" MAX_SHAPE = 300 # 12 sec * 25 hz DATASET = "train" def encode_audio_file(input_audio_file): input_audio_path = os.path.join(OUT_AUDIO_DIR, input_audio_file) output_file_name = os.path.join( OUT_AUDIO_MERT_LABEL_DIR, f"{DATASET}_{input_audio_file.replace('.wav', '')}.npy", ) if os.path.exists(output_file_name): return audio = Audio.from_file(input_audio_path) output_arrays = mert_25.encode( audio, do_clustering=True, n_codebooks=1, ).astype(np.int16) if os.path.exists(output_file_name): print(f"{output_file_name} is already written") return np.save(output_file_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) _ = mert_25.preload_models( checkpoint_filepath=MERT_CKPT_PATH, centroids_filepath=MERT_CENTER_PATH, ) # load current data input_args = parse_args() tsv_info = [] # check GPU with open(os.path.join(OUT_TSV_DIR, f"{DATASET}.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(encode_audio_file, work_items, max_workers=8, chunksize=1) print(f"Done~!, {time() - start_time:.3f} seconds.")