stem-restoration / restoflow /aux_corr.py
soilkon's picture
sync app + active models
7153194 verified
Raw History Blame Contribute Delete
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()