"""Auto-select the best restorer + best generator by latent-FAD (REMIX), and write _best.txt (BEST_RESTORER / BEST_GEN) for the experimental GAN to consume. Decode-free. Restorers: ranked here by REMIX latent-FAD (distributional quality) with REMIX-STFT shown too. Generators: read from the phase-2 gen-FAD comparison log (_gen_fad_compare.log). """ from __future__ import annotations import argparse, glob, random, re from collections import defaultdict from pathlib import Path import torch from .config import Cfg, STEMS, STEM_ID from . import data as D, fad from .quality_probe import load_run BASE = Path("/media/maindisk/ksoil/repos/Restoration/Dac_demucs_lyra") def restorer_remix_fad(run, dev, pairs, cfg): m, st = load_run(run, dev) cl_set, rr_set = [], [] with torch.inference_mode(): for (_pid, its) in pairs: clat = rlat = None for it in its: s = it["stem"]; sid = torch.tensor([STEM_ID[s]], device=dev) rc = D._fit_T(torch.load(it["clean"], map_location="cpu").float(), cfg.T) rd = D._fit_T(torch.load(it["deg"], map_location="cpu").float(), cfg.T) rm = D._fit_T(torch.load(it["mix"], map_location="cpu").float(), cfg.T) sd = ((rd - st[s]["mu"][:, None]) / st[s]["sd"][:, None]).to(dev)[None] sm = ((rm - st["__mix__"]["mu"][:, None]) / st["__mix__"]["sd"][:, None]).to(dev)[None] rr = D._fit_T(_destd_cpu(m(sd, sid, sm)[0], st, s), cfg.T) clat = rc if clat is None else clat + rc rlat = rr if rlat is None else rlat + rr if clat is not None: cl_set.append(clat); rr_set.append(rlat) return fad.fad(fad.latent_set(cl_set), fad.latent_set(rr_set)) def _destd_cpu(z, stats, stem): return (z * stats[stem]["sd"][:, None].to(z.device) + stats[stem]["mu"][:, None].to(z.device)).detach().cpu() def main(): p = argparse.ArgumentParser() p.add_argument("--device", default="cuda"); p.add_argument("--n-pairs", type=int, default=150) a = p.parse_args(); dev = a.device; random.seed(0) cfg = Cfg(val_cache_roots=(str(BASE / "demucs_results_val_full"),), use_mix=True, model_kind="attn", T=32, device=dev) 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); pairs = pairs[:a.n_pairs] rest_runs = [Path(d).name for d in sorted(glob.glob(str(BASE / "restoflow_runs/restorer_*"))) if (Path(d) / "ckpt_best.pt").exists()] print(f"=== restorer REMIX latent-FAD (lower=better) over {len(pairs)} pairs ===") scores = {} for r in rest_runs: try: scores[r] = restorer_remix_fad(r, dev, pairs, cfg); print(f" {r:32} {scores[r]:.3f}") except Exception as e: print(f" {r:32} FAILED {e}") best_rest = min(scores, key=scores.get) if scores else "restorer_attn_v1" # generators: parse phase-2 comparison gfad = {}; clog = BASE / "restoflow_runs/_gen_fad_compare.log" if clog.exists(): cur = None for ln in clog.read_text().splitlines(): mrun = re.search(r"#####\s+(\S+)", ln) if mrun: cur = mrun.group(1) mf = re.search(r"FAD-REMIX\s*\|\s*([0-9.]+)", ln) if mf and cur: gfad[cur] = float(mf.group(1)) print(f"\n=== generator REMIX-FAD (from _gen_fad_compare.log) ===") for g, v in gfad.items(): print(f" {g:32} {v:.3f}") best_gen = min(gfad, key=gfad.get) if gfad else "gen_distvar_baseline" out = BASE / "restoflow_runs/_best.txt" out.write_text(f"BEST_RESTORER={best_rest}\nBEST_GEN={best_gen}\n") print(f"\nBEST_RESTORER={best_rest}\nBEST_GEN={best_gen}\n-> {out}") if __name__ == "__main__": main()