"""Experiment-eval (separate, evolving table) — score OUR restorer variants vs the PINNED baselines. Does NOT recompute AudioSR/FlashSR (those are pinned static in BASELINES.md); only runs the OUR variants through the same pipeline + eval set as the pin, so numbers are directly comparable. Residual-aware: if a variant's ckpt has `rflow` (the generative-residual restorer, exp_genrestore), we monkeypatch APP.restore to apply det-backbone + sampled flow residual on the restore-id stems. PRIMARY metric = FAD-CLAP (guards: FAD-VGGish, LSD-HF). Same fixed set + seed as the pin. Run: python -m restoflow.eval_exp --device cuda:0 --runs exp_genrestore,exp_couple --ours-n 1024 --musdb-n 62 """ from __future__ import annotations import argparse, glob, os, random, 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") SYNTH_VAL = Path("/media/maindisk/melkor169/MSRKit/xlance-msr/stereo_a2sb/synthetic_out/val") # READ-ONLY # PINNED baselines (BASELINES.md, 2026-06-24) shown as static reference rows. [CLAP, VGGish, LSD-HF, LSD] PINNED = { "OURS-val": {"degraded": [0.231, 0.547, 3.108, 2.822], "advramp(pinned)": [0.100, 0.288, 1.430, 1.349], "AudioSR": [0.190, 0.848, 2.017, 1.875], "FlashSR": [0.087, 0.278, 1.828, 1.669], "codec-floor": [0.080, 0.205, 1.342, 1.259]}, "MUSDB": {"degraded": [0.403, 1.560, 4.970, 4.513], "advramp(pinned)": [0.342, 1.585, 2.231, 2.111], "AudioSR": [0.188, 0.209, 1.665, 1.533], "FlashSR": [0.162, 0.265, 1.491, 1.366], "codec-floor": [0.080, 0.670, 1.120, 1.116]}, } def wire_rflow(APP, run, dev): """If `run` has a residual flow, monkeypatch APP.restore = det backbone + sampled residual (no grad).""" import torch from restoflow.model import CondFlow ck = torch.load(BASE / "restoflow_runs" / run / "ckpt_best.pt", map_location="cpu") if "rflow" not in ck: return None rc = ck["rflow_cfg"]; rf = CondFlow(256, rc["hidden"], rc["depth"], n_stems=len(APP.STEMS), stem_emb=rc["stem_emb"]).to(dev).eval() rf.load_state_dict(ck["rflow"]) rids = set(rc["restore_ids"]); steps = rc["gen_steps"] print(f"[exp] {run}: rflow wired ({rc['hidden']}x{rc['depth']}) on stem-ids {sorted(rids)}, {steps} steps") @torch.inference_mode() def restore_resid(latent, stem, mix): m = APP.models(); st = m["rstats"]; mu, sd = st[stem]["mu"], st[stem]["sd"]; mm = st["__mix__"] sd_d = ((latent - mu[:, None]) / sd[:, None]).to(dev)[None] sd_m = ((mix - mm["mu"][:, None]) / mm["sd"][:, None]).to(dev)[None] sid = torch.tensor([APP.STEM_ID[stem]], device=dev) out = m["rest"](sd_d, sid, sd_m) # [1,256,T] std backbone if APP.STEM_ID[stem] in rids: # + sampled residual (Euler, no grad) z = torch.randn_like(out); dt = 1.0 / steps for i in range(steps): t = torch.full((out.shape[0],), i * dt, device=dev) z = z + dt * rf(z, t, out, sid) out = out + z return (out[0].cpu()) * sd[:, None] + mu[:, None] return restore_resid def build_set(APP, name, items): """items: list of (clean[2,T], deg[2,T]) -> run OURS pipeline (current APP restorer). Returns clean, ours.""" work = Path(tempfile.mkdtemp(prefix="evalexp_")); clean_l, ours_l = [], [] for clean, deg in items: fp = work / "d.wav"; sf.write(fp, deg.T, APP.SR, subtype="FLOAT") out = APP.run_upload(str(fp)) if not out or out[1] is None: continue clean_l.append(clean); ours_l.append(np.asarray(out[1][1]).T) return clean_l, ours_l def main(): p = argparse.ArgumentParser() p.add_argument("--device", default="cuda:0"); p.add_argument("--runs", default="exp_genrestore,exp_couple") p.add_argument("--ours-n", type=int, default=1024); p.add_argument("--musdb-n", type=int, default=62) p.add_argument("--excerpt-s", type=float, default=9.0); p.add_argument("--lp-sr", type=int, default=8000) p.add_argument("--hf-hz", type=float, default=4000.0) p.add_argument("--musdb-root", default="/media/maindisk/melkor169/drums_dereverb/data/gmd_musdb18hq_stereo") a = p.parse_args() os.environ["RESTOFLOW_DEVICE"] = a.device # Reuse the shared separation cache (separation is identical across variants -> separate each clip once). os.environ.setdefault("RESTOFLOW_SEP_CACHE", str(BASE / "sep_cache")) import restoflow.app as APP from restoflow import lit_eval as LE, fad as FAD import torch SR = 44100; random.seed(0) # ---- build the fixed eval set ONCE (same seed/selection as the pin) ---- datasets = {} cleans = sorted(glob.glob(str(SYNTH_VAL / "clean" / "*.wav"))); random.shuffle(cleans); items = [] for cp in cleans[: a.ours_n * 2]: dp = cp.replace("/clean/", "/degraded/") if not os.path.exists(dp): continue c, _ = sf.read(cp, dtype="float32"); d, _ = sf.read(dp, dtype="float32") c = c.T if c.ndim == 2 else np.stack([c, c]); d = d.T if d.ndim == 2 else np.stack([d, d]) n = min(c.shape[1], d.shape[1]); items.append((np.ascontiguousarray(c[:, :n]), np.ascontiguousarray(d[:, :n]))) if len(items) >= a.ours_n: break datasets["OURS-val"] = items if Path(a.musdb_root).exists(): exc = int(a.excerpt_s * SR); tracks = sorted([d for d in glob.glob(str(Path(a.musdb_root) / "*")) if Path(d).is_dir()]) random.shuffle(tracks); mitems = [] for td in tracks[: a.musdb_n * 2]: mixp = Path(td) / "mixture.wav" if mixp.exists(): clean, _ = sf.read(mixp, dtype="float32"); clean = clean.T else: sts = [sf.read(Path(td) / f"{s}.wav", dtype="float32")[0].T for s in APP.STEMS if (Path(td) / f"{s}.wav").exists()] if not sts: continue clean = sum(sts) if clean.ndim == 1: clean = np.stack([clean, clean]) if clean.shape[1] < exc: continue st = (clean.shape[1] - exc) // 2; clean = np.ascontiguousarray(clean[:, st:st + exc]) deg = np.stack([librosa.resample(clean[c], orig_sr=SR, target_sr=a.lp_sr) for c in range(2)]) deg = np.ascontiguousarray(np.stack([librosa.resample(deg[c], orig_sr=a.lp_sr, target_sr=SR) for c in range(2)])[:, :clean.shape[1]].astype(np.float32)) if not np.isfinite(deg).all() or float(np.sqrt(np.mean(deg ** 2))) < 1e-5: continue mitems.append((clean, deg)) if len(mitems) >= a.musdb_n: break datasets["MUSDB"] = mitems print(f"[exp] eval set: " + ", ".join(f"{k}={len(v)}" for k, v in datasets.items())) to_t = lambda L: [torch.from_numpy(np.ascontiguousarray(x)).float() for x in L] results = {ds: {} for ds in datasets} # ds -> run -> [CLAP,VGGish,LSD-HF,LSD] for run in [r for r in a.runs.split(",") if r]: os.environ["RESTOFLOW_DEFAULT_REST"] = run; APP._M.clear() import gc; gc.collect(); torch.cuda.empty_cache() APP.models() orig_restore = APP.restore patched = wire_rflow(APP, run, a.device) if patched is not None: APP.restore = patched try: for ds, items in datasets.items(): clean_l, ours_l = build_set(APP, ds, items) M = defaultdict(list) for cl, ou in zip(clean_l, ours_l): r_, e_ = LE.align(cl, ou, SR) M["LSD-HF"].append(LE.lsd(r_, e_, SR, hf_hz=a.hf_hz)); M["LSD"].append(LE.lsd(r_, e_, SR)) clap = FAD.CLAP(); vgg = FAD.VGGish() ec_c = clap.embed_set(to_t(clean_l), SR); eo_c = clap.embed_set(to_t(ours_l), SR) ec_v = vgg.embed_set(to_t(clean_l), SR); eo_v = vgg.embed_set(to_t(ours_l), SR) fad_clap = FAD.fad(ec_c, eo_c, shrink=len(ec_c) < 2 * ec_c.shape[1]) fad_vgg = FAD.fad(ec_v, eo_v, shrink=len(ec_v) < 2 * ec_v.shape[1]) results[ds][run] = [fad_clap, fad_vgg, float(np.mean(M["LSD-HF"])), float(np.mean(M["LSD"]))] print(f"[exp] {run} · {ds} (N={len(clean_l)}): CLAP={fad_clap:.3f} VGG={fad_vgg:.3f} " f"LSD-HF={np.mean(M['LSD-HF']):.3f} LSD={np.mean(M['LSD']):.3f}") finally: APP.restore = orig_restore # ---- table: variants (measured) + pinned baselines (static reference) ---- cols = ["FAD-CLAP*", "FAD-VGGish", "LSD-HF", "LSD"] for ds in datasets: print(f"\n=== {ds} · experiment-eval (variants measured; baselines = pinned static) ===") print(f"{'model':22} " + " ".join(f"{c:>11}" for c in cols)) for run in [r for r in a.runs.split(",") if r]: v = results[ds].get(run) if v: print(f"{run:22} " + " ".join(f"{x:>11.3f}" for x in v) + (" <- vs advramp" if run else "")) for ref, v in PINNED[ds].items(): print(f"{ref+' [pin]':22} " + " ".join(f"{x:>11.3f}" for x in v)) print("\n* FAD-CLAP = PRIMARY (promote rule: beat advramp on CLAP both datasets + no VGGish/LSD-HF regression)") print("EVAL_EXP_DONE") if __name__ == "__main__": main()