File size: 4,874 Bytes
7153194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
"""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()