"""FINAL external baselines on MUSDB18-HQ TEST (full mixes), robustdeg_v1 degradation. The honest, re-pinned bar (replaces the stale drums/lowpass BASELINES.md). Fixed, reproducible eval set: the official 50 test songs, deterministic high-energy 9s excerpts, degraded with robustdeg_v1 (bandwidth + hiss/hum/clicks/tape). Methods compared vs the CLEAN full mix: degraded · OURS(advramp pipeline) · AudioSR · FlashSR · codec-floor(SAME encode->decode of clean). Metrics: FAD-CLAP (PRIMARY), FAD-VGGish, LSD, LSD-HF. AudioSR/FlashSR run as isolated-venv subprocesses. NOTE: robustdeg adds hiss/hum/clicks that AudioSR/FlashSR (bandwidth-SR only) are NOT built to remove — so they set the "pure-SR" bar; we expect to beat them on the full restoration task. Run on the eval GPU: python -m restoflow.eval_baselines --device cuda:1 --excerpts-per-song 6 [--no-audiosr --no-flashsr] """ from __future__ import annotations import argparse, glob, os, random, subprocess, tempfile from collections import defaultdict from pathlib import Path import numpy as np import soundfile as sf import librosa BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") MUSDB_TEST = BASE / "data" / "musdb18hq" / "test" AUDIOSR_VENV = "/tmp/audiosr_env/bin/python" FLASHSR_DIR = "/tmp/flashsr" STEMS = ("other", "vocals", "drums", "bass") def _read(fp, SR): a, sr = sf.read(str(fp), dtype="float32"); a = a.T if a.ndim == 2 else np.stack([a, a]) if a.ndim == 1: a = np.stack([a, a]) if sr != SR: a = np.stack([librosa.resample(a[c], orig_sr=sr, target_sr=SR) for c in range(2)]) return np.ascontiguousarray(a.astype(np.float32)) def _load_out(d, names, SR): out = [] for fn in names: fp = Path(d) / fn if not fp.exists(): out.append(None); continue out.append(_read(fp, SR)) return out def main(): p = argparse.ArgumentParser() p.add_argument("--device", default="cuda:1") p.add_argument("--excerpts-per-song", type=int, default=12) p.add_argument("--max-songs", type=int, default=0) p.add_argument("--excerpt-s", type=float, default=3.0, help="3s = the native SAME unit (faithful to training)") p.add_argument("--hf-hz", type=float, default=4000.0); p.add_argument("--ddim", type=int, default=50) p.add_argument("--export-n", type=int, default=3); p.add_argument("--out-tag", default="eval_baselines") p.add_argument("--no-audiosr", action="store_true"); p.add_argument("--no-flashsr", action="store_true") p.add_argument("--write-md", action="store_true", help="write the pinned BASELINES.md") a = p.parse_args() os.environ["RESTOFLOW_DEVICE"] = a.device import restoflow.app as APP from restoflow import lit_eval as LE, fad as FAD from restoflow.robustdeg import degrade, RobustDegConfig, ROBUSTDEG_VERSION import torch APP.models(); SR = APP.SR; random.seed(0); dcfg = RobustDegConfig() expdir = BASE / "restoflow_runs" / "_logs" / a.out_tag; expdir.mkdir(parents=True, exist_ok=True) work = Path(tempfile.mkdtemp(prefix="evalbase_")) L = int(a.excerpt_s * SR) # ---------- STAGE 1: fixed MUSDB-test eval set (robustdeg) + OURS + codec ---------- songs = sorted([d for d in MUSDB_TEST.glob("*") if d.is_dir()]) if a.max_songs: songs = songs[:a.max_songs] degdir = work / "deg"; degdir.mkdir(parents=True) clean_l, deg_l, ours_l, cf_l, names = [], [], [], [], [] for si, song in enumerate(songs): stems = {s: _read(song / f"{s}.wav", SR) for s in STEMS if (song / f"{s}.wav").exists()} if len(stems) < len(STEMS): continue n = min(v.shape[1] for v in stems.values()) mixp = song / "mixture.wav" mix = _read(mixp, SR)[:, :n] if mixp.exists() else np.ascontiguousarray(sum(v[:, :n] for v in stems.values())) # deterministic high-energy excerpts nseg = max(1, n // L); rng = random.Random(si) cand = [i * L for i in range(nseg) if float(np.sqrt(np.mean(mix[:, i*L:(i+1)*L]**2))) > 1e-3] rng.shuffle(cand) for s0 in sorted(cand[:a.excerpts_per_song]): clean = np.ascontiguousarray(mix[:, s0:s0 + L]) deg, _ = degrade(clean, SR, random.Random(si * 1000 + s0), dcfg) fn = f"{len(names)}.wav"; sf.write(degdir / fn, deg.T, SR, subtype="FLOAT") out = APP.run_upload(str(degdir / fn)) if not out or out[1] is None: (degdir / fn).unlink(missing_ok=True); continue ours = np.asarray(out[1][1]).T clean_l.append(clean); deg_l.append(deg); ours_l.append(ours) cf_l.append(LE.codec_passthrough(clean, APP)); names.append(fn) print(f"[eval] MUSDB-test: built {len(names)} excerpts from {len(songs)} songs (robustdeg={ROBUSTDEG_VERSION})", flush=True) # ---------- STAGE 2: free our models ---------- APP._M.clear(); import gc; gc.collect(); torch.cuda.empty_cache() # ---------- STAGE 3/4: AudioSR + FlashSR subprocesses ---------- gpu = a.device.split(":")[-1] if ":" in a.device else "1" # inherit external CUDA_VISIBLE_DEVICES isolation if set (so subprocs land on the SAME physical GPU as # this process); else pin to --device's index. base_env = {**os.environ, "CUDA_VISIBLE_DEVICES": os.environ.get("CUDA_VISIBLE_DEVICES", gpu), "PYTORCH_CUDA_ALLOC_CONF": "expandable_segments:True", "HF_HUB_OFFLINE": "1", "TRANSFORMERS_OFFLINE": "1"} audiosr = [None] * len(names); flashsr = [None] * len(names) if not a.no_audiosr: ad = work / "audiosr" r = subprocess.run([AUDIOSR_VENV, str(BASE / "restoflow/run_audiosr.py"), "--in-dir", str(degdir), "--out-dir", str(ad), "--ddim", str(a.ddim), "--chunk-s", "3.0", "--device", "cuda"], env=base_env, capture_output=True, text=True) print(f"[eval] AudioSR: {r.stdout.strip().splitlines()[-1] if r.stdout.strip() else r.stderr[-200:]}") audiosr = _load_out(ad, names, SR) if not a.no_flashsr: fd = work / "flashsr" r = subprocess.run([AUDIOSR_VENV, str(BASE / "restoflow/run_flashsr.py"), "--in-dir", str(degdir), "--out-dir", str(fd), "--device", "cuda"], env={**base_env, "PYTHONPATH": FLASHSR_DIR}, cwd=FLASHSR_DIR, capture_output=True, text=True) print(f"[eval] FlashSR: {r.stdout.strip().splitlines()[-1] if r.stdout.strip() else r.stderr[-200:]}") flashsr = _load_out(fd, names, SR) # ---------- STAGE 5: metrics ---------- data = {"degraded": deg_l, "OURS": ours_l, "AudioSR": audiosr, "FlashSR": flashsr, "codec-floor": cf_l} cols = list(data); N = len(names) M = {c: defaultdict(list) for c in cols} for i in range(N): for c in cols: if data[c][i] is not None: LE.add_metrics(M, c, clean_l[i], data[c][i], SR, a.hf_hz) to_t = lambda Lst: [torch.from_numpy(np.ascontiguousarray(x)).float() for x in Lst if x is not None] rows = {} for label, klass in (("FAD-CLAP", "CLAP"), ("FAD-VGGish", "VGGish")): try: emb = getattr(FAD, klass)(); ec = emb.embed_set(to_t(clean_l), SR); D = ec.shape[1]; valid = len(ec) >= 2 * D rows[label] = {} for c in cols: tl = to_t(data[c]) rows[label][c] = FAD.fad(ec, emb.embed_set(tl, SR), shrink=not valid) if tl else float("nan") rows[label]["_flag"] = "valid" if valid else f"N={len(ec)}<2*{D}" except Exception as e: rows[label] = {"_flag": f"skip:{type(e).__name__}"} for m in ("LSD-HF", "LSD", "SI-SDR"): rows[m] = {c: (float(np.mean(M[c][m])) if M[c][m] else float("nan")) for c in cols} print(f"\n=== MUSDB18-HQ TEST · {N} excerpts · robustdeg_v1 · vs clean full mix ===") print(f"{'metric':12} " + " ".join(f"{c:>12}" for c in cols)) for m in ("FAD-CLAP", "FAD-VGGish", "LSD-HF", "LSD", "SI-SDR"): line = " ".join(f"{rows[m].get(c, float('nan')):>12.3f}" for c in cols) print(f"{m:12} {line} {rows[m].get('_flag','')}") # exports for i in range(min(a.export_n, N)): sf.write(expdir / f"ex{i}_clean.wav", np.asarray(clean_l[i]).T, SR, subtype="FLOAT") for c in cols: if data[c][i] is not None: sf.write(expdir / f"ex{i}_{c}.wav", np.asarray(data[c][i]).T, SR, subtype="FLOAT") print(f"[eval] exports -> {expdir}") if a.write_md: md = [f"# Baselines — PINNED on MUSDB18-HQ test (full mixes), {ROBUSTDEG_VERSION}. FAD-CLAP = primary.\n", f"Measured by `restoflow/eval_baselines.py`. N={N} excerpts ({a.excerpts_per_song}/song × official 50 test).", "Task = real lo-fi restoration (bandwidth + hiss/hum/clicks/tape). AudioSR/FlashSR are bandwidth-SR", "only (don't remove hiss) -> they set the pure-SR bar. codec-floor = SAME encode->decode of clean.\n", "| metric | " + " | ".join(cols) + " |", "|---|" + "|".join(["---:"] * len(cols)) + "|"] for m in ("FAD-CLAP", "FAD-VGGish", "LSD-HF", "LSD", "SI-SDR"): md.append(f"| {m} | " + " | ".join(f"{rows[m].get(c, float('nan')):.3f}" for c in cols) + " |") (BASE / "BASELINES.md").write_text("\n".join(md) + "\n") print(f"[eval] wrote BASELINES.md") print("EVAL_BASELINES_DONE") if __name__ == "__main__": main()