import os import argparse from suno_utils.tasks import demucs from suno_utils.audio import Audio 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 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_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") def demuc_audio_file(input_audio_file): output_vocal_file_name = os.path.join( OUT_AUDIO_DEMUC_DIR, input_audio_file.replace(".wav", "_vocals.wav") ) if os.path.exists(output_vocal_file_name): return audio = Audio.from_file(os.path.join(OUT_AUDIO_DIR, input_audio_file)) try: vocals, non_vocals = demucs.split_vocals(audio) vocals.to_wav(output_vocal_file_name) non_vocals.to_wav( os.path.join( OUT_AUDIO_DEMUC_DIR, input_audio_file.replace(".wav", "_bkg.wav") ) ) except Exception as E: print(E, input_audio_file) audio.to_wav(output_vocal_file_name) audio.to_wav( os.path.join( OUT_AUDIO_DEMUC_DIR, input_audio_file.replace(".wav", "_bkg.wav") ) ) 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) demucs.load_model(device="cuda") # load current data input_args = parse_args() tsv_info = [] # check GPU with open(os.path.join(OUT_TSV_DIR, "train.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(demuc_audio_file, work_items, max_workers=10, chunksize=1) print(f"Done~!, {time() - start_time:.3f} seconds.")