"""Correct final eval: load a trained checkpoint and score it over the FULL val cache with stratified (all-stem) coverage. Uses the run's TRAIN norm stats (not rebuilt). python -m restoflow.final_eval --ckpt restoflow_runs/v3_big/ckpt_best.pt \ --val-cache-roots demucs_results_val_full --eval-max-pairs 0 """ from __future__ import annotations import argparse from dataclasses import replace from pathlib import Path import torch from .config import Cfg from . import eval as E from .model import CondFlow, DetRestorer def main(): p = argparse.ArgumentParser() p.add_argument("--ckpt", required=True) p.add_argument("--val-cache-roots", default="demucs_results_val_full") p.add_argument("--eval-max-pairs", type=int, default=0, help="0 = all val pairs") p.add_argument("--device", default="cuda") p.add_argument("--demo-pairs", type=int, default=8) a = p.parse_args() ck = torch.load(a.ckpt, map_location="cpu") base = Cfg() saved = {k: v for k, v in ck["cfg"].items() if hasattr(base, k)} cfg = replace(base, **saved) run_dir = Path(a.ckpt).parent cfg = replace(cfg, val_cache_roots=tuple(x for x in a.val_cache_roots.split(",") if x), eval_max_pairs=a.eval_max_pairs, device=a.device, demo_pairs=a.demo_pairs, out_dir=str(run_dir), stats_path=str(run_dir / "norm_stats.pt")) stats = torch.load(cfg.stats_file(), map_location="cpu") if cfg.model_kind == "det": model = DetRestorer(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), stem_emb=cfg.stem_emb_dim, use_mix=cfg.use_mix) else: model = CondFlow(cfg.latent_dim, cfg.hidden, cfg.depth, n_stems=len(cfg.stems), stem_emb=cfg.stem_emb_dim) model.load_state_dict(ck["model"]); model = model.to(cfg.device) print(f"[final_eval] ckpt={a.ckpt} epoch={ck.get('epoch')} {cfg.model_kind} " f"params={model.num_params()/1e6:.2f}M val={cfg.val_cache_roots} budget={a.eval_max_pairs or 'ALL'}") sa, sr = E.load_same(cfg) E.evaluate(cfg, model, stats, sa, sr, epoch=9999) if __name__ == "__main__": main()