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)
