import os import glob import torch import argparse import torchaudio from suno_boost.utils import ( load_diffusion_upsample_model, apply_normalization, block_based_inference, ) if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("input", help="Path to input audio file to upsample.") parser.add_argument("-o", "--output", help="Filepath to save output") parser.add_argument( "--ckpt_path", help="Path to pretrained model checkpoint.", default="/app/suno/christian/boost-checkpoints/epoch=10-step=68750.ckpt", ) parser.add_argument( "--block_size", help="Overlap-add block size for block-based processing", default=262144, ) parser.add_argument( "--num_steps", help="Number of diffusion sampling steps", default=100, type=int, ) parser.add_argument( "--batch_size", default=1, type=int, ) parser.add_argument("--use_gpu", action="store_true") args = parser.parse_args() # load model model = load_diffusion_upsample_model(args.ckpt_path) # load audio file input_audio, input_sample_rate = torchaudio.load(args.input) # resample if required if input_sample_rate != model.sample_rate: input_audio = torchaudio.functional.resample( input_audio, input_sample_rate, model.sample_rate ) # run sampling process boosted = block_based_inference( input_audio, model, block_size=args.block_size, overlap=args.block_size // 2, num_steps=args.num_steps, use_gpu=args.use_gpu, batch_size=args.batch_size, ) # optional normalization boosted /= boosted.abs().max().clamp(1e-8) # save audio output to disk if args.output is None: output_dir = os.getcwd() basename = os.path.basename(args.input).split(".")[0] output_filepath = os.path.join(output_dir, f"{basename}-boosted.wav") else: output_filepath = args.output torchaudio.save(output_filepath, boosted, model.sample_rate)