import os import time import glob import torch import argparse import torchaudio import numpy as np import scipy.signal import pyloudnorm as pyln import multiprocessing as mp from tqdm import tqdm from suno_amp.v0 import apply_postprocessing VERSIONS = [0] def run(filepath: str, output_dir: str, version: int = 0): audio, sample_rate = torchaudio.load(filepath) filename = os.path.basename(filepath).split(".")[0] print(filename) # start = time.perf_counter() if version == 0: output_audio = apply_postprocessing(audio, sample_rate) else: raise ValueError(f"Invalid version: {version}. Must be one of {VERSIONS}.") # end = time.perf_counter() # elapsed = timings.append(end - start) # print(np.mean(elapsed)) # loudnorm the input for comparision # loudness normalize to target LUFS dB ffmpeg_filter = f"loudnorm=I=-16.0" effector = torchaudio.io.AudioEffector(ffmpeg_filter) audio_norm = effector.apply(audio.T, sample_rate).T # save output out_filepath_input = os.path.join(output_dir, filename + ".wav") out_filepath_output = os.path.join(output_dir, filename + "-output.wav") torchaudio.save(out_filepath_input, audio_norm, sample_rate) torchaudio.save(out_filepath_output, output_audio, sample_rate) if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument( "input_dir", help="Path to directory containing audio files to process." ) parser.add_argument("--version", default=0, type=int) args = parser.parse_args() # make directory to save outputs args.input_dir = args.input_dir.rstrip("/") # remove trailing / dirname = os.path.dirname(args.input_dir) basename = os.path.basename(args.input_dir) output_dir = os.path.join(dirname, basename + f"+postprocess-v{args.version}") os.makedirs(output_dir, exist_ok=True) print(f"Saving outputs in {output_dir}...") # find all files # find all audio files filepaths = [] for ext in ["flac", "wav", "mp3", "ogg"]: filepaths += glob.glob(os.path.join(args.input_dir, f"*.{ext}")) timings = [] filepaths = filepaths[:20] # run algo on each audio file args = [(filepath, output_dir, args.version) for filepath in filepaths] with mp.Pool(32) as pool: pool.starmap(run, args)