Spaces:
Sleeping
Sleeping
File size: 8,121 Bytes
7153194 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 | """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()
|