Spaces:
Sleeping
Sleeping
Download restoflow/aux_corr.py from soilkon/stem-restoration: direct link, hf CLI and curl.
- Browser
- Download file 8.12 kB
-
https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/aux_corr.py
- Command line
-
hf download hf://spaces/soilkon/stem-restoration/restoflow/aux_corr.py
-
curl -L -o aux_corr.py https://huggingface.co/spaces/soilkon/stem-restoration/resolve/main/restoflow/aux_corr.py
8.12 kB
| """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() | |