"""Apples-to-apples generator FAD: load saved gen runs and re-run the (now FAD-enabled) eval once each, on identical val pairs, so per-stem + REMIX latent-FAD is comparable across the capacity ladder (conv-10M vs xattn-40M vs xattn-63M). Reuses evaluate_gen / evaluate_gen_x. Run: python -m restoflow.gen_fad_probe --runs gen_distvar_baseline,gen_xattn_40M,gen_xattn_63M --device cuda """ from __future__ import annotations import argparse from pathlib import Path import torch from .config import Cfg, STEMS from .model import CondFlow, AttnCondFlow, XAttnCondFlow from . import eval as E from .gen import index_clips, evaluate_gen, evaluate_gen_x BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") def main(): p = argparse.ArgumentParser() p.add_argument("--runs", default="gen_distvar_baseline,gen_xattn_40M,gen_xattn_63M") p.add_argument("--val-cache-roots", default=str(BASE / "demucs_results_val_full")) p.add_argument("--device", default="cuda") p.add_argument("--eval-pairs", type=int, default=120) a = p.parse_args() dev = a.device cfg = Cfg(device=dev, T=32) sa, sr = E.load_same(cfg) va_roots = [x for x in a.val_cache_roots.split(",") if x] for run in [r for r in a.runs.split(",") if r]: ck_path = BASE / "restoflow_runs" / run / "ckpt_best.pt" if not ck_path.exists(): print(f"\n##### {run}: no ckpt_best.pt — skip"); continue ck = torch.load(ck_path, map_location="cpu"); ga = ck["args"] arch = ga.get("arch", "conv") if arch == "xattn": m = XAttnCondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8)) elif arch == "attn": m = AttnCondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim, heads=ga.get("heads", 8)) else: m = CondFlow(256, ga["hidden"], ga["depth"], n_stems=len(STEMS), stem_emb=Cfg().stem_emb_dim) m.load_state_dict(ck["model"]); m.eval().to(dev) stats = ck["stats"] tgt_stems = [s for s in ga.get("target_stems", "bass").split(",") if s] sources = tuple(x for x in ga.get("context_sources", "restored,degraded").split(",") if x) c = type(cfg)(**{**cfg.__dict__, "sample_steps": ga.get("sample_steps", 40)}) va = index_clips(va_roots, tgt_stems, sources) print(f"\n##### {run} arch={arch} {ga['hidden']}x{ga['depth']} " f"params={sum(p_.numel() for p_ in m.parameters())/1e6:.1f}M val={len(va)} #####") demo = str(BASE / "restoflow_runs" / run / "_fadprobe_demo") if arch == "xattn": evaluate_gen_x(m, va, stats, sa, sr, c, ga.get("cfg_w", 2.0), [ga.get("cfg_rescale", 0.0)], 0, demo, a.eval_pairs, 0) else: evaluate_gen(m, va, stats, sa, sr, c, ga.get("cfg_w", 2.0), 0, demo, a.eval_pairs, 0) if __name__ == "__main__": main()