# This file takes audio/audios and encode them through the codec to generate codec labels import os import argparse from suno_utils.audio import Audio from suno_utils.tasks import dac 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_speech" OUT_AUDIO_DIR = os.path.join(OUT_DATA_DIR, "audio") OUT_AUDIO_DEMUC_DIR = os.path.join(OUT_DATA_DIR, "audio_demuc") OUT_AUDIO_DEMUC_LABEL_DIR = os.path.join(OUT_DATA_DIR, "audio_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" CODEC_CKPT_PATH = "/home/tony/Data/MERT/codec.pt" MAX_SHAPE = 300 DATASET_TYPE = "valid" def encode_audio_file(input_audio_file): # pre-pend demuc dir input_audio_file_path = os.path.join(OUT_AUDIO_DIR, input_audio_file) input_audio_raw_name = input_audio_file.replace(".wav", "") output_file_name = os.path.join( OUT_AUDIO_DEMUC_LABEL_DIR, f"{DATASET_TYPE}_{input_audio_raw_name}.npy" ) if os.path.exists(output_file_name): return # print(input_audio_files) audio = Audio.from_file(input_audio_file_path) output_arrays = dac.encode(audio) # output_arrays = dac.encode_files(input_audio_files, batch_size=BATCH_SIZE, n_gpus=1) # print(output_arrays) # output_arrays = output_arrays[0].astype(np.int16) # final_arrays = np.zeros((len(input_audio_files), MAX_SHAPE, 8), dtype=np.int16) # for n_row, arr in enumerate(output_arrays): # if arr is None: # print(f"WTF arr is None, {n_row}, {input_audio_files[n_row]}") # raise ValueError # # this job will not finish! # return # # arr = np.ones((MAX_SHAPE, 8), dtype=np.int16) # # arr *= -1 # arr = arr.astype(np.int16) # if arr.shape[0] < MAX_SHAPE: # arr = np.pad( # arr, # ((0, MAX_SHAPE - arr.shape[0]), (0, 0)), # "constant", # constant_values=-1, # ) # final_arrays[n_row, :] = arr.astype(np.int16) if os.path.exists(output_file_name): print(f"{output_file_name} is already written") return np.save( os.path.join( OUT_AUDIO_DEMUC_LABEL_DIR, f"{DATASET_TYPE}_{input_audio_raw_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) dac.load_model(checkpoint_filepath=CODEC_CKPT_PATH, device="cuda") # load current data input_args = parse_args() tsv_info = [] # check GPU with open(os.path.join(OUT_TSV_DIR, f"{DATASET_TYPE}.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] # .replace(".wav", "_vocals.wav") for x in tsv_data[input_args.start_index : input_args.end_index] ] # batches = [ # work_items[x : x + BATCH_SIZE] for x in range(0, len(work_items), BATCH_SIZE) # ] print(len(work_items), "work_items") # , len(batches), "batches") thread_map(encode_audio_file, work_items, max_workers=4, chunksize=1) # thread_map(encode_audio_file, batches, max_workers=8, chunksize=1) # for batch in batches: # encode_audio_file(batch) print(f"Done~!, {time() - start_time:.3f} seconds.")