"""Run AudioSR (2024 baseline) over a folder of degraded wavs -> 48kHz restored wavs. RUNS IN THE ISOLATED VENV (/tmp/audiosr_env) — only audiosr + soundfile + numpy here, NOT our SAME-L env. Invoked by baseline_eval.py via subprocess. Usage: /tmp/audiosr_env/bin/python -m restoflow.run_audiosr --in-dir X --out-dir Y --ddim 50 (or run as a plain script path; no restoflow imports so it works under the isolated venv) """ import argparse, gc, glob, os, tempfile import numpy as np import soundfile as sf import torch import torchaudio def _sf_load(filepath, *args, **kwargs): """soundfile-backed torchaudio.load (this torchaudio build routes load through torchcodec, which isn't installed and we have no network). Returns (wav[ch,frames] float32, sr).""" x, sr = sf.read(filepath, dtype="float32", always_2d=True) # [frames, ch] return torch.from_numpy(x.T.copy()), sr torchaudio.load = _sf_load from audiosr import super_resolution def build_model_lowmem(model_name="basic", device="cuda"): """Like audiosr.build_model but loads the 6.18GB checkpoint on CPU then moves ONLY the model to GPU. Stock build_model does torch.load(map_location=device) -> checkpoint AND model co-resident on GPU (~9.4GB) -> OOMs a 10GB card at LOAD time (before any inference). Model-alone is ~3-4GB fp32, so CPU-load + move fits 10GB.""" import yaml # noqa from audiosr.latent_diffusion.models.ddpm import LatentDiffusion from audiosr.utils import default_audioldm_config, download_checkpoint print(f"[lowmem] Loading AudioSR {model_name} (ckpt on CPU -> model to {device})") ckpt_path = download_checkpoint(model_name) config = default_audioldm_config(model_name) config["model"]["params"]["device"] = device # self.device for runtime tensor creation ld = LatentDiffusion(**config["model"]["params"]) # params init on CPU ckpt = torch.load(ckpt_path, map_location="cpu") # 6.18GB -> CPU RAM, not GPU ld.load_state_dict(ckpt["state_dict"], strict=False) del ckpt; gc.collect() ld.eval() ld = ld.to(device) # move ~3-4GB model only if torch.cuda.is_available(): torch.cuda.empty_cache() return ld def _mono(x): x = np.asarray(x).squeeze() return x.mean(0) if (x.ndim == 2 and x.shape[0] < x.shape[-1]) else (x.mean(-1) if x.ndim == 2 else x) def process_file(model, f, ddim, guidance, chunk_s): """CHUNK to chunk_s (matches our 3s native unit). AudioSR pads each chunk up to its internal ~5.12s segment and emits THAT length at 48k -> crop each output back to the chunk's true duration (scaled to 48k) so the concatenated 48k output aligns 1:1 with the 44.1k input.""" wav, sr = sf.read(f, dtype="float32"); wav = _mono(wav) W = int(chunk_s * sr); outs = [] for st in range(0, max(1, len(wav)), W): seg = wav[st:st + W]; true_len = len(seg) # unpadded length at sr if len(seg) < W: # pad tail (AudioSR needs steady length) seg = np.pad(seg, (0, W - len(seg))) tmp = tempfile.mktemp(suffix=".wav"); sf.write(tmp, seg, sr) o = _mono(super_resolution(model, tmp, seed=42, ddim_steps=ddim, guidance_scale=guidance)) keep = int(round(true_len / sr * 48000)) # crop AudioSR's padded output to true 48k length outs.append(o[:keep]); os.remove(tmp) if torch.cuda.is_available(): torch.cuda.empty_cache() return np.concatenate(outs) # 48 kHz mono, len == input_dur*48k def main(): ap = argparse.ArgumentParser() ap.add_argument("--in-dir", required=True); ap.add_argument("--out-dir", required=True) ap.add_argument("--ddim", type=int, default=50); ap.add_argument("--guidance", type=float, default=3.5) ap.add_argument("--chunk-s", type=float, default=3.0) # match our 3s native processing unit ap.add_argument("--device", default="cuda") a = ap.parse_args() os.makedirs(a.out_dir, exist_ok=True) model = build_model_lowmem(model_name="basic", device=a.device) # CPU-load -> fits 10GB GPU files = sorted(glob.glob(os.path.join(a.in_dir, "*.wav"))) print(f"[audiosr] {len(files)} files, ddim={a.ddim}, chunk={a.chunk_s}s") for i, f in enumerate(files): try: w = process_file(model, f, a.ddim, a.guidance, a.chunk_s) sf.write(os.path.join(a.out_dir, os.path.basename(f)), w.astype(np.float32), 48000) except Exception as e: print(f"[audiosr] FAIL {os.path.basename(f)}: {type(e).__name__}: {e}") if torch.cuda.is_available(): torch.cuda.empty_cache() if (i + 1) % 5 == 0: print(f"[audiosr] {i+1}/{len(files)}") print("AUDIOSR_DONE") if __name__ == "__main__": main()