"""Literature-comparable evaluation: place our restoration on the SAME axes the BWE / music-restoration papers report, so we can see where we stand. Shared baselines & what they report: - AERO (ICASSP'23, spectral audio super-res): LSD, ViSQOL, MUSHRA — VCTK / MUSDB18. - BigWavGAN ('23, wave GAN music super-res): LSD, SI-SDR, ViSQOL — MUSDB18. - BABE-2 / Diffusion Generative Equalizer (DAFx'24, Moliner): FAD, LSD — music restoration. Shared metrics we compute here on OUR val set, for the BWE-comparable unit = the mix of present (restoration-class) stems: degraded (lower bound) and OURS (advramp restoration), each vs clean (oracle): - LSD log-spectral distance (lower=better; the BWE standard) - LSD-HF LSD restricted to >hf_hz (the high band — our brightness/hiss question, where BWE lives) - SI-SDR scale-invariant SDR (higher=better; fidelity) - FAD Frechet Audio Distance (VGGish + CLAP-music; lower=better; the generative-restoration metric) CAVEAT printed in the report: our dataset + degradation differ from each paper's, so absolute numbers are NOT a head-to-head ranking — they place us in the same REGIME/units. For a true head-to-head, re-run on MUSDB18-HQ with a matched low-pass degradation (--musdb path, future). Run: python -m restoflow.lit_eval --device cuda --n-pairs 120 """ from __future__ import annotations import argparse, csv, math, random from collections import defaultdict from pathlib import Path import numpy as np import torch from .config import Cfg, STEMS from . import eval as E BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") EPS = 1e-8 def _stft_mag(x, n_fft=2048, hop=512): w = torch.hann_window(n_fft) X = torch.stft(torch.from_numpy(x).float(), n_fft, hop_length=hop, window=w, return_complex=True) return X.abs().numpy() # [F, T] def lsd(ref, est, sr=44100, n_fft=2048, hop=512, hf_hz=None): """Log-spectral distance (dB-domain power), mono. hf_hz -> restrict to the high band only.""" r = _stft_mag(ref.mean(0), n_fft, hop); e = _stft_mag(est.mean(0), n_fft, hop) T = min(r.shape[1], e.shape[1]); r, e = r[:, :T], e[:, :T] lr = np.log10(r ** 2 + EPS); le = np.log10(e ** 2 + EPS) if hf_hz is not None: f0 = int(round(hf_hz / (sr / 2) * (r.shape[0] - 1))) lr, le = lr[f0:], le[f0:] return float(np.mean(np.sqrt(np.mean((lr - le) ** 2, axis=0)))) def si_sdr(ref, est): """Scale-invariant SDR (dB), mono, length-aligned.""" r = ref.mean(0).astype(np.float64); e = est.mean(0).astype(np.float64) n = min(len(r), len(e)); r, e = r[:n], e[:n] r = r - r.mean(); e = e - e.mean() alpha = np.dot(e, r) / (np.dot(r, r) + EPS) s = alpha * r; noise = e - s return float(10 * np.log10((np.dot(s, s) + EPS) / (np.dot(noise, noise) + EPS))) def _best_lag(r, e, max_lag): """FFT cross-correlation lag of est vs ref within +-max_lag (no scipy dep).""" n = 1 << int(math.ceil(math.log2(2 * len(r) + 1))) cc = np.fft.irfft(np.fft.rfft(e, n) * np.conj(np.fft.rfft(r, n)), n) cc = np.concatenate([cc[-max_lag:], cc[:max_lag + 1]]) return int(np.arange(-max_lag, max_lag + 1)[int(np.argmax(cc))]) def align(ref, est, sr=44100, max_lag_s=0.05): """Paper-faithful: correct ONLY the constant codec/window time-latency (sub-50ms cross-correlation) so LSD/SI-SDR measure spectral/waveform quality, not a systematic offset. NO level-norm — LSD should count level differences, as the SR papers do (their models output level-matched audio).""" n = min(ref.shape[1], est.shape[1]); ref, est = ref[:, :n], est[:, :n] lag = _best_lag(ref.mean(0).astype(np.float64), est.mean(0).astype(np.float64), int(max_lag_s * sr)) if lag > 0: est = est[:, lag:]; ref = ref[:, :est.shape[1]] elif lag < 0: ref = ref[:, -lag:]; est = est[:, :ref.shape[1]] L = min(ref.shape[1], est.shape[1]); return ref[:, :L], est[:, :L] def add_metrics(M, tag, ref, est, sr, hf_hz): ref, est = align(ref, est, sr) M[tag]["LSD"].append(lsd(ref, est, sr)); M[tag]["LSD-HF"].append(lsd(ref, est, sr, hf_hz=hf_hz)) M[tag]["SI-SDR"].append(si_sdr(ref, est)) def codec_passthrough(audio, APP): """SAME encode->decode of a mix (3s windows, no separation/restoration) = the autoencoder codec FLOOR our system cannot beat. Lets MUSDB show how much of the gap is the codec vs the restoration.""" SR = APP.SR; chunk = SR * 3; outs = [] for i in range(0, audio.shape[1], chunk): seg = audio[:, i:i + chunk]; o = seg.shape[1] if o < chunk: seg = np.pad(seg, ((0, 0), (0, chunk - o))) outs.append(np.asarray(APP.decode(APP.encode(seg)))) return np.concatenate(outs, 1) if outs else audio def _fad_rows(aud, sr): """FAD-VGGish + FAD-CLAP-music for {deg,ours} vs clean. Returns list of (name, deg, ours) or notes.""" rows = [] try: from . import fad as FAD except Exception as e: return [("FAD", f"skipped ({e})", "")] to_t = lambda L: [torch.from_numpy(np.ascontiguousarray(x)).float() for x in L] # embed_set wants tensors for name, klass in (("VGGish", "VGGish"), ("CLAP", "CLAP")): try: emb = getattr(FAD, klass)() ec, ed, eo = (emb.embed_set(to_t(aud[k]), sr) for k in ("clean", "deg", "ours")) D = ec.shape[1]; Nmin = min(len(ec), len(ed), len(eo)) valid = Nmin >= 2 * D # need embeddings >> dim for a full-rank covariance sh = not valid # empirical cov when valid; shrink ONLY (flagged) when under-sampled tag = f"FAD-{name} (N={Nmin},D={D}{' ✓' if valid else ' ⚠shrink'})" rows.append((tag, f"{FAD.fad(ec, ed, shrink=sh):.3f}", f"{FAD.fad(ec, eo, shrink=sh):.3f}")) except Exception as e: rows.append((f"FAD-{name}", f"skip:{type(e).__name__}", "")) return rows def _print_table(label, M, aud, sr, n, tags=("deg", "ours")): titles = {"deg": "degraded", "ours": "OURS", "codecfloor": "codec-floor"} print(f"\n=== {label} · {n} items (vs clean oracle) ===") print(f"{'metric':12} " + " ".join(f"{titles.get(t, t):>12}" for t in tags)) for m in ("LSD", "LSD-HF", "SI-SDR"): arrow = "↑" if m == "SI-SDR" else "↓" print(f"{m:12} " + " ".join(f"{np.mean(M[t][m]):>12.3f}" for t in tags) + f" ({arrow} better)") for name, dv, ov in _fad_rows(aud, sr): # FAD only for deg/ours vs clean set print(f" {name:34} deg={dv:>9} ours={ov:>9}") def run_musdb(a): """Full-pipeline eval on MUSDB18-HQ: clean mixture -> deterministic low-pass degrade -> our app pipeline (separate->restore->remix) -> metrics vs clean. Same axes as the BWE papers' MUSDB tables.""" import os, glob, tempfile, soundfile as sf, librosa os.environ.setdefault("RESTOFLOW_DEVICE", a.device) import restoflow.app as APP APP.models() SR = APP.SR; exc = int(a.excerpt_s * SR) M = {t: defaultdict(list) for t in ("deg", "ours", "codecfloor")}; aud = {"clean": [], "deg": [], "ours": []} tracks = sorted([d for d in glob.glob(str(Path(a.musdb_root) / "*")) if Path(d).is_dir()]) random.shuffle(tracks); tracks = tracks[:a.musdb_n]; used = 0 for td in tracks: 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 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)]) # low-pass 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 # skip silent/invalid excerpts tmp = tempfile.mktemp(suffix=".wav"); sf.write(tmp, deg.T, SR, subtype="FLOAT") out = APP.run_upload(tmp) if not out or out[1] is None: # run_upload bailed (sep failed) continue ours = np.asarray(out[1][1]).T codec = codec_passthrough(clean, APP) # SAME codec floor (no sep/restore) add_metrics(M, "deg", clean, deg, SR, a.hf_hz) add_metrics(M, "ours", clean, ours, SR, a.hf_hz) add_metrics(M, "codecfloor", clean, codec, SR, a.hf_hz) aud["clean"].append(clean); aud["deg"].append(deg); aud["ours"].append(ours); used += 1 _print_table(f"MUSDB18-HQ (raw-clean ref; low-pass {a.lp_sr}Hz; full pipeline incl. separation+codec)", M, aud, SR, used, tags=("deg", "ours", "codecfloor")) print(" (codec-floor = SAME encode->decode of clean, NO separation/restoration = the ceiling our latent") print(" pipeline can reach; the ours->codecfloor gap is restoration+separation, codecfloor->0 is the codec tax.)") def main(): p = argparse.ArgumentParser() p.add_argument("--device", default="cuda"); p.add_argument("--n-pairs", type=int, default=120) p.add_argument("--hf-hz", type=float, default=4000.0, help="LSD-HF band cutoff") p.add_argument("--val-root", default=str(BASE / "demucs_results_val_full")) p.add_argument("--musdb-root", default="/media/maindisk/melkor169/drums_dereverb/data/gmd_musdb18hq_stereo") p.add_argument("--musdb-n", type=int, default=25); p.add_argument("--excerpt-s", type=float, default=9.0) p.add_argument("--lp-sr", type=int, default=8000, help="low-pass cutoff = lp_sr/2 (BWE degradation)") p.add_argument("--skip-ourval", action="store_true"); p.add_argument("--skip-musdb", action="store_true") a = p.parse_args(); dev = a.device; random.seed(0) cfg = Cfg(device=dev, T=32) sa, sr = E.load_same(cfg) root = Path(a.val_root) # ---- (1) OUR val set (cached latents; restoration-class stem mix) ---- if not a.skip_ourval: rows = list(csv.DictReader(open(root / "metadata.csv"))) by_pair = defaultdict(list) for r in rows: by_pair[r["pair_id"]].append(r) pairs = list(by_pair.items()); random.shuffle(pairs); pairs = pairs[:a.n_pairs] M = {"deg": defaultdict(list), "ours": defaultdict(list)}; aud = {"clean": [], "deg": [], "ours": []}; used = 0 with torch.inference_mode(): for pid, srows in pairs: cl = dg = rr = None for r in srows: s = r.get("stem") if s not in STEMS or r.get("class") != "restoration": continue cp = r.get("clean_latent_path"); dp = r.get("deg_latent_path") rp = root / pid / "latents" / "restored" / f"{s}.pt" if not (cp and dp and (root / cp).exists() and (root / dp).exists() and rp.exists()): continue c = E._decode(sa, torch.load(root / cp, map_location="cpu").float()) d = E._decode(sa, torch.load(root / dp, map_location="cpu").float()) o = E._decode(sa, torch.load(rp, map_location="cpu").float()) cl = c if cl is None else cl[:, :c.shape[1]] + c[:, :cl.shape[1]] dg = d if dg is None else dg[:, :d.shape[1]] + d[:, :dg.shape[1]] rr = o if rr is None else rr[:, :o.shape[1]] + o[:, :rr.shape[1]] if cl is None: continue cl, dg, rr = (x.numpy() if hasattr(x, "numpy") else x for x in (cl, dg, rr)) n = min(cl.shape[1], dg.shape[1], rr.shape[1]); cl, dg, rr = cl[:, :n], dg[:, :n], rr[:, :n] for tag, est in (("deg", dg), ("ours", rr)): add_metrics(M, tag, cl, est, sr, a.hf_hz) # aligned + level-matched aud["clean"].append(cl); aud["deg"].append(dg); aud["ours"].append(rr); used += 1 # NOTE: here clean/deg/ours are ALL SAME-decoded -> same codec domain -> the delta is fair; the # absolute LSD is vs codec-clean (not raw), so it isn't directly the papers' axis (they use raw clean). _print_table("OUR val (codec-domain ref; restoration delta is fair, absolute != papers)", M, aud, sr, used) # ---- (2) MUSDB18-HQ (shared dataset; low-pass degrade; full pipeline) ---- if not a.skip_musdb and Path(a.musdb_root).exists(): run_musdb(a) print("\nCurrent (2024-2026) baselines on this task family + what they report:") print(" AudioSR (2024) latent-diffusion versatile SR->48k | LSD,LSD-HF,ViSQOL,FAD | THE std baseline; ALSO latent => shares our codec floor") print(" FlashSR (2025) 1-step diffusion-distill, SOTA music SR | LSD,ViSQOL,FAD | per-cutoff music SR") print(" UniverSR (2025) vocoder-free FLOW-MATCHING SR, SOTA | LSD,LSD-HF,2f-model | 8/12/16/24->48k (no codec => no floor)") print(" BABE-2 (2024) BLIND generative music restoration | FAD(CLAP),LSD | closest to OUR blind/unknown-degradation setting") print(" CAVEAT: those are KNOWN-downsample, full-mix, single-stage SR; OURS is blind + per-stem + invents") print(" missing instruments => no exact head-to-head. Shared axes = LSD / LSD-HF / FAD. (ViSQOL/2f-model: TODO.)") if __name__ == "__main__": main()