stem-restoration / restoflow /run_audiosr.py
soilkon's picture
sync app + active models
7153194 verified
Raw History Blame Contribute Delete
4.87 kB
"""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()