"""Validate CHEAP latent-FAD vs REAL CLAP-music-FAD, per stem + REMIX. For a sample of val pairs, builds 4 condition sets vs clean: degraded, det-restored (restorer_attn_v1), flowattn-restored (scout_rest_flowattn), plus a clean-split FLOOR. Reports FAD(clean, X) for both backends. If the two backends RANK the conditions the same, the cheap latent-FAD is a valid stand-in (no decode, no embedder). Run: python -m restoflow.fad_eval --n-pairs 150 --device cuda """ from __future__ import annotations import argparse, random from collections import defaultdict import numpy as np, torch from .config import Cfg, STEMS, STEM_ID from . import data as D, flow as Fl, eval as E, fad from .model import AttnRestorer, MixAttnCondFlow BASE = "/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra" DET = f"{BASE}/restoflow_runs/restorer_attn_v1" FLOW = f"{BASE}/restoflow_runs/scout_rest_flowattn" CONDS = ["degraded", "det", "flowattn"] def _load_det(dev): ck = torch.load(f"{DET}/ckpt_best.pt", map_location="cpu"); c = ck["cfg"] m = AttnRestorer(256, c["hidden"], c["depth"], n_stems=len(STEMS), stem_emb=c["stem_emb_dim"], use_mix=c["use_mix"], heads=c.get("heads", 4)).to(dev).eval() m.load_state_dict(ck["model"]); return m, torch.load(f"{DET}/norm_stats.pt", map_location="cpu") def _load_flow(dev): ck = torch.load(f"{FLOW}/ckpt_best.pt", map_location="cpu"); c = ck["cfg"] m = MixAttnCondFlow(256, c["hidden"], c["depth"], n_stems=len(STEMS), stem_emb=c["stem_emb_dim"], heads=c.get("heads", 8), use_mix=c["use_mix"]).to(dev).eval() m.load_state_dict(ck["model"]); return m, torch.load(f"{FLOW}/norm_stats.pt", map_location="cpu"), c.get("sigma", 0.3) def main(): p = argparse.ArgumentParser() p.add_argument("--n-pairs", type=int, default=150) p.add_argument("--device", default="cuda") p.add_argument("--steps", type=int, default=40) p.add_argument("--no-clap", action="store_true") a = p.parse_args() dev = a.device random.seed(0); torch.manual_seed(0) cfg = Cfg(val_cache_roots=(f"{BASE}/demucs_results_val_full",), use_mix=True, model_kind="attn", T=32, device=dev, sample_steps=a.steps) sa, sr = E.load_same(cfg) det, dstats = _load_det(dev) fl, fstats, sigma = _load_flow(dev) print(f"[fad] models loaded; flow sigma={sigma}; sr={sr}") items = D._index(cfg, cfg.val_cache_roots, "val") by_pair = defaultdict(list) for it in items: by_pair[it["pair_id"]].append(it) pairs = list(by_pair.items()); random.shuffle(pairs) freq = defaultdict(int) # front-load rare-stem (bass/drums) pairs for _, its in pairs: for it in its: freq[it["stem"]] += 1 pairs.sort(key=lambda kv: min(freq[it["stem"]] for it in kv[1])) pairs = pairs[:a.n_pairs] print(f"[fad] {len(pairs)} pairs, {sum(len(v) for _,v in pairs)} stem-instances") # collectors: lat[cond][stem] = list[256,T]; aud[cond][stem] = list[2,S]; remix per pair lat = {c: defaultdict(list) for c in ["clean"] + CONDS} aud = {c: defaultdict(list) for c in ["clean"] + CONDS} rlat = {c: [] for c in ["clean"] + CONDS} raud = {c: [] for c in ["clean"] + CONDS} def destd(std, stats, stem): return E._destd(std[0], stats, stem) with torch.inference_mode(): for pi, (pid, its) in enumerate(pairs): psum = {c: None for c in ["clean"] + CONDS}; pasum = {c: None for c in ["clean"] + CONDS} for it in its: stem = it["stem"]; sid = torch.tensor([STEM_ID[stem]], device=dev) raw_c = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) raw_d = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) raw_m = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) # det normalization dmu, dsd = dstats[stem]["mu"], dstats[stem]["sd"]; mm = dstats["__mix__"] std_d = ((raw_d - dmu[:, None]) / dsd[:, None]).to(dev)[None] std_mix = ((raw_m - mm["mu"][:, None]) / mm["sd"][:, None]).to(dev)[None] raw_det = destd(det(std_d, sid, std_mix), dstats, stem) # flow normalization (own stats) fmu, fsd = fstats[stem]["mu"], fstats[stem]["sd"]; fm = fstats["__mix__"] fstd_d = ((raw_d - fmu[:, None]) / fsd[:, None]).to(dev)[None] fstd_mix = ((raw_m - fm["mu"][:, None]) / fm["sd"][:, None]).to(dev)[None] fstd_r = Fl.sample_mix(fl, fstd_d, fstd_mix, sid, cfg.sample_steps, sigma) raw_flow = destd(fstd_r, fstats, stem) raws = {"clean": raw_c.to(dev), "degraded": raw_d.to(dev), "det": raw_det, "flowattn": raw_flow} for c in ["clean"] + CONDS: lat[c][stem].append(raws[c].cpu()) au = E._decode(sa, raws[c]) aud[c][stem].append(au) psum[c] = raws[c] if psum[c] is None else psum[c] + raws[c] pasum[c] = au if pasum[c] is None else pasum[c] + au for c in ["clean"] + CONDS: if psum[c] is not None: rlat[c].append(psum[c].cpu()); raud[c].append(pasum[c]) if (pi + 1) % 25 == 0: print(f" ...{pi+1}/{len(pairs)} pairs") clap = None if a.no_clap else fad.CLAP(device=dev) def report(title, lat_sets, aud_sets): # lat_sets/aud_sets: dict cond -> list (latents [256,T] / audios [2,S]) n = len(lat_sets["clean"]); half = n // 2 clean_lat = fad.latent_set(lat_sets["clean"]) floor_lat = (fad.fad(fad.latent_set(lat_sets["clean"][:half]), fad.latent_set(lat_sets["clean"][half:])) if half >= 1 else float("nan")) row = {c: fad.fad(clean_lat, fad.latent_set(lat_sets[c])) for c in CONDS} print(f"\n[{title}] LATENT-FAD (cheap) floor(clean-split)={floor_lat:.3f}") print(" " + " ".join(f"{c}={row[c]:.3f}" for c in CONDS)) if clap is not None: clean_au = clap.embed_set(aud_sets["clean"], sr) floor_au = (fad.fad(clap.embed_set(aud_sets["clean"][:half], sr), clap.embed_set(aud_sets["clean"][half:], sr)) if half >= 1 else float("nan")) rowc = {c: fad.fad(clean_au, clap.embed_set(aud_sets[c], sr)) for c in CONDS} print(f"[{title}] CLAP-FAD (real) floor(clean-split)={floor_au:.3f}") print(" " + " ".join(f"{c}={rowc[c]:.3f}" for c in CONDS)) for stem in STEMS: if lat["clean"][stem]: report(f"stem={stem} n={len(lat['clean'][stem])}", {c: lat[c][stem] for c in ["clean"] + CONDS}, {c: aud[c][stem] for c in ["clean"] + CONDS}) report(f"REMIX n={len(rlat['clean'])}", rlat, raud) if __name__ == "__main__": main()