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()