"""Cheap-auxiliary-loss validation: does a cheap LATENT-domain loss actually track the REAL decoded metrics we care about (FAD-CLAP primary, FAD-VGGish, LSD, LSD-HF)? Only auxes that prove meaningful get adopted (objective recipe — no instinct-picked auxiliaries). Method (FAD is distributional -> measured over a quality SWEEP; LSD is per-sample -> measured per-sample, both aggregated across sweep conditions): 1. encode N clean 3s clips -> clean latents; ref audio = decode(clean) (codec-clean -> isolates the perturbation effect from the codec). 2. for each perturbation family {hiss, dull, collapse, atten} x strength level: perturb the clean latents, compute CHEAP latent losses (var_match, mmd, latent-FAD, L2) vs clean, AND decode -> REAL metrics (FAD-CLAP, FAD-VGGish set-level; LSD, LSD-HF per-sample mean) vs ref. 3. across all (family,level) rows, Spearman+Pearson corr of each cheap loss vs each real metric. A cheap loss is MEANINGFUL only if |corr| is strong across DIFFERENT families (robust, not 1-family). Startup STRESS PROBE: runs the single most demanding condition first (full N, strongest perturb, full decode + CLAP) and prints peak GPU mem, so an unattended run won't silently OOM. Run (after the baseline pin frees cuda:1): python -m restoflow.aux_corr --device cuda:1 --n 256 """ from __future__ import annotations import argparse, glob, os, random from pathlib import Path import numpy as np import soundfile as sf import torch SYNTH_VAL = Path("/media/maindisk/melkor169/MSRKit/xlance-msr/stereo_a2sb/synthetic_out/val/clean") # READ-ONLY def _spear(x, y): x, y = np.asarray(x, float), np.asarray(y, float) rx = np.argsort(np.argsort(x)); ry = np.argsort(np.argsort(y)) return float(np.corrcoef(rx, ry)[0, 1]) def _pear(x, y): return float(np.corrcoef(np.asarray(x, float), np.asarray(y, float))[0, 1]) def main(): p = argparse.ArgumentParser() p.add_argument("--device", default="cuda:1"); p.add_argument("--n", type=int, default=256) p.add_argument("--levels", type=int, default=6); p.add_argument("--hf-hz", type=float, default=4000.0) p.add_argument("--dec-bs", type=int, default=16) a = p.parse_args() os.environ["RESTOFLOW_DEVICE"] = a.device import restoflow.app as APP from restoflow import distloss as DL, fad as FAD, lit_eval as LE APP.models(); SR = APP.SR; dev = a.device random.seed(0) # ---- clean latents + codec-clean ref audio ---- files = sorted(glob.glob(str(SYNTH_VAL / "*.wav"))); random.shuffle(files); files = files[: a.n] auds = [] for f in files: w, sr = sf.read(f, dtype="float32"); w = w.T if w.ndim == 2 else np.stack([w, w]) if sr != SR: # synthetic is 44.1k already; guard anyway import librosa; w = np.stack([librosa.resample(w[c], orig_sr=sr, target_sr=SR) for c in range(2)]) auds.append(np.ascontiguousarray(w)) with torch.inference_mode(): lat_list = [] # batch the encode (256-at-once SAME fwd OOMs) for i in range(0, len(auds), a.dec_bs): lat_list.extend(APP.encode_batch(auds[i:i + a.dec_bs])) clean_lat = torch.stack(lat_list).to(dev).float() # [N,256,T] ref_audio = [np.asarray(x) for x in _decode_all(APP, clean_lat, a.dec_bs)] print(f"[aux] N={len(files)} clean latents {tuple(clean_lat.shape)}") fam_std = clean_lat.std().item() def perturb(name, lat, s): if name == "hiss": return lat + s * fam_std * torch.randn_like(lat) # additive noise -> hiss if name == "dull": m = lat.mean(dim=2, keepdim=True); return m + (1 - s) * (lat - m) # temporal flatten if name == "collapse": m = lat.mean(dim=(0, 2), keepdim=True); return m + (1 - s) * (lat - m) # mode collapse if name == "atten": return lat * (1 - 0.6 * s) # energy loss return lat fams = ["hiss", "dull", "collapse", "atten"] levels = [i / (a.levels - 1) for i in range(a.levels)] # 0..1 ; 0 = identity (sanity anchor) clap, vgg = FAD.CLAP(), FAD.VGGish() clap_clean = clap.embed_set([torch.from_numpy(x).float() for x in ref_audio], SR) vgg_clean = vgg.embed_set([torch.from_numpy(x).float() for x in ref_audio], SR) lat_clean_emb = FAD.latent_set(clean_lat.cpu()) # ---- STRESS PROBE: most demanding condition first (full N, strongest hiss, full decode+CLAP) ---- if torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats(dev) _ = _row(APP, DL, FAD, LE, perturb, "hiss", 1.0, clean_lat, ref_audio, clap, vgg, clap_clean, vgg_clean, lat_clean_emb, SR, a.hf_hz, a.dec_bs) if torch.cuda.is_available(): print(f"[aux] STRESS PROBE ok — peak GPU {torch.cuda.max_memory_allocated(dev)/1e9:.2f} GB on {dev}") # ---- full sweep ---- rows = [] for fam in fams: for s in levels: r = _row(APP, DL, FAD, LE, perturb, fam, s, clean_lat, ref_audio, clap, vgg, clap_clean, vgg_clean, lat_clean_emb, SR, a.hf_hz, a.dec_bs) r.update({"family": fam, "level": s}); rows.append(r) print(f" {fam:9} s={s:.2f} | var={r['var']:.3f} mmd={r['mmd']:.3f} lfad={r['lfad']:.2f} L2={r['l2']:.3f}" f" || CLAP={r['fad_clap']:.3f} VGG={r['fad_vgg']:.3f} LSD={r['lsd']:.3f} LSD-HF={r['lsd_hf']:.3f}") # ---- correlation matrix ---- cheap = ["var", "mmd", "lfad", "l2"]; real = ["fad_clap", "fad_vgg", "lsd", "lsd_hf"] print("\n=== Spearman corr (cheap aux ↓ vs real metric →) — across all families+levels ===") print(f"{'aux':6} " + " ".join(f"{r:>10}" for r in real)) for c in cheap: print(f"{c:6} " + " ".join(f"{_spear([row[c] for row in rows], [row[r] for row in rows]):>10.3f}" for r in real)) print("\n=== per-FAMILY Spearman vs FAD-CLAP (robustness: meaningful = strong across ALL families) ===") print(f"{'aux':6} " + " ".join(f"{f:>10}" for f in fams)) for c in cheap: line = [] for f in fams: fr = [row for row in rows if row["family"] == f] line.append(f"{_spear([row[c] for row in fr], [row['fad_clap'] for row in fr]):>10.3f}") print(f"{c:6} " + " ".join(line)) print("\nVERDICT RULE: adopt an aux only if |Spearman vs FAD-CLAP| is high AND consistent across families") print("(a 1-family-only correlation = artifact, not a real surrogate). Same check vs LSD/LSD-HF.") print("AUX_CORR_DONE") def _decode_all(APP, lat, bs): outs = [] for i in range(0, lat.shape[0], bs): d = APP.decode_batch([lat[j] for j in range(i, min(i + bs, lat.shape[0]))]) # decode_batch wants a list outs.extend([np.asarray(x) for x in d]) return outs def _row(APP, DL, FAD, LE, perturb, fam, s, clean_lat, ref_audio, clap, vgg, clap_clean, vgg_clean, lat_clean_emb, SR, hf_hz, bs): import torch, numpy as np with torch.inference_mode(): pl = perturb(fam, clean_lat, s) var = float(DL.var_match_loss(pl, clean_lat)); mmd = float(DL.mmd_loss(pl, clean_lat)) l2 = float(((pl - clean_lat) ** 2).mean()) lfad = float(FAD.fad(lat_clean_emb, FAD.latent_set(pl.cpu()))) pert_audio = _decode_all(APP, pl, bs) cp = clap.embed_set([torch.from_numpy(x).float() for x in pert_audio], SR) vg = vgg.embed_set([torch.from_numpy(x).float() for x in pert_audio], SR) fad_clap = float(FAD.fad(clap_clean, cp, shrink=len(clap_clean) < 2 * clap_clean.shape[1])) fad_vgg = float(FAD.fad(vgg_clean, vg, shrink=len(vgg_clean) < 2 * vgg_clean.shape[1])) lsds, lsdhfs = [], [] for r_, e_ in zip(ref_audio, pert_audio): rr, ee = LE.align(np.asarray(r_), np.asarray(e_), SR) lsds.append(LE.lsd(rr, ee, SR)); lsdhfs.append(LE.lsd(rr, ee, SR, hf_hz=hf_hz)) return {"var": var, "mmd": mmd, "lfad": lfad, "l2": l2, "fad_clap": fad_clap, "fad_vgg": fad_vgg, "lsd": float(np.mean(lsds)), "lsd_hf": float(np.mean(lsdhfs))} if __name__ == "__main__": main()