"""Evaluation script for few-shot RS-IMLE models. Computes FID (5000 samples) and Precision/Recall (1000 samples) at one or more latent noise scales, and saves a sample grid. """ import argparse import json import os import sys import shutil import time sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import imageio import numpy as np import torch import torch.nn.functional as F import torchvision from distutils.util import strtobool from data import set_up_data from hps import Hyperparams, add_imle_arguments, parse_args_and_update_hparams from models import IMLE from sampler import Sampler from cleanfid import fid import cleanfid.features as _cleanfid_feat import cleanfid.inception_torchscript as _cleanfid_incept _inception_cache = os.path.join(os.path.expanduser("~"), ".cache", "cleanfid") os.makedirs(_inception_cache, exist_ok=True) _orig_feature_extractor = _cleanfid_feat.feature_extractor def _patched_feature_extractor(name="torchscript_inception", device=torch.device("cuda"), resize_inside=False, use_dataparallel=True): if name == "torchscript_inception": model = _cleanfid_incept.InceptionV3W( _inception_cache, download=True, resize_inside=resize_inside ).to(device) model.eval() if use_dataparallel: model = torch.nn.DataParallel(model) return lambda x: model(x) return _orig_feature_extractor(name, device, resize_inside, use_dataparallel) _cleanfid_feat.feature_extractor = _patched_feature_extractor # ────────────────────────────────────────────────────────────────────────── # Generation helpers # ────────────────────────────────────────────────────────────────────────── def generate_images_to_dir(model, latent_dim, num_images, out_dir, batch_size=16, noise_scale=1.0): """Generate images with latents sampled as z ~ N(0, noise_scale^2).""" os.makedirs(out_dir, exist_ok=True) device = next(model.parameters()).device idx = 0 with torch.no_grad(): while idx < num_images: bs = min(batch_size, num_images - idx) z = torch.randn(bs, latent_dim, device=device) * noise_scale imgs = model(z, None) imgs = (imgs + 1.0) * 127.5 imgs = imgs.clamp(0, 255).permute(0, 2, 3, 1) imgs = imgs.cpu().numpy().astype(np.uint8) for j in range(bs): imageio.imwrite(os.path.join(out_dir, f"{idx}.png"), imgs[j]) idx += 1 def generate_grid_image(model, latent_dim, num_samples, nrow, out_path, noise_scale=1.0): device = next(model.parameters()).device with torch.no_grad(): z = torch.randn(num_samples, latent_dim, device=device) * noise_scale imgs = model(z, None) imgs = (imgs + 1.0) / 2.0 imgs = imgs.clamp(0.0, 1.0) grid = torchvision.utils.make_grid(imgs, nrow=nrow, padding=2) grid_pil = torchvision.transforms.functional.to_pil_image(grid.cpu()) grid_pil.save(out_path) print(f" Saved grid ({num_samples} images) to {out_path}") def _slerp_train_style(a: torch.Tensor, b: torch.Tensor, t: torch.Tensor) -> torch.Tensor: """Same spherical interpolation as train.py (FID-epoch interp strips).""" a = F.normalize(a, dim=-1) b = F.normalize(b, dim=-1) dot = torch.sum(a * b, dim=-1, keepdim=True).clamp(-1.0, 1.0) omega = torch.acos(dot) sin_omega = torch.sin(omega) t = t.view(-1, 1) factor1 = torch.sin((1.0 - t) * omega) / sin_omega factor2 = torch.sin(t * omega) / sin_omega return factor1 * a + factor2 * b def generate_slerp_grid(model, H, sampler, interp_pairs, interp_steps, nrow, out_path, noise_scale=1.0): """SLERP strips like train.py / supplementary few-shot figure (latent geodesics).""" device = next(model.parameters()).device with torch.no_grad(): z1 = torch.randn(interp_pairs, H.latent_dim, device=device, dtype=torch.float32) * noise_scale z2 = torch.randn(interp_pairs, H.latent_dim, device=device, dtype=torch.float32) * noise_scale t_vals = torch.linspace(0.0, 1.0, interp_steps, device=device, dtype=torch.float32) all_rows = [] for pi in range(interp_pairs): z_interp = _slerp_train_style( z1[pi:pi + 1].repeat(interp_steps, 1), z2[pi:pi + 1].repeat(interp_steps, 1), t_vals, ) snoise_tmp = [s[:interp_steps].normal_() for s in sampler.snoise_tmp] preds = sampler.sample(z_interp, model, snoise_tmp) preds_t = torch.from_numpy(preds).float() / 255.0 preds_t = preds_t.permute(0, 3, 1, 2) all_rows.append(preds_t) tensor = torch.cat(all_rows, dim=0) grid = torchvision.utils.make_grid(tensor, nrow=nrow, padding=2) grid_pil = torchvision.transforms.functional.to_pil_image( grid.cpu().clamp(0.0, 1.0)) grid_pil.save(out_path) n = interp_pairs * interp_steps print(f" Saved SLERP grid ({n} images, {interp_pairs}x{interp_steps}) to {out_path}") # ────────────────────────────────────────────────────────────────────────── # Model loading # ────────────────────────────────────────────────────────────────────────── def override_model_cycles(model, test_H_cycles=None, test_L_cycles=None, test_refinement_steps=None): raw = model.module if hasattr(model, 'module') else model mapper = raw.decoder.mapping_network inner = getattr(mapper, 'trm', None) if inner is None: print(" No TRM inner module found; cycle override skipped.") return old_H = inner.H_cycles old_L = inner.L_cycles old_halt = mapper.refinement_steps if test_H_cycles is not None: inner.H_cycles = test_H_cycles if test_L_cycles is not None: inner.L_cycles = test_L_cycles if test_refinement_steps is not None: mapper.refinement_steps = test_refinement_steps print(f" Cycle override: H {old_H}->{inner.H_cycles}, " f"L {old_L}->{inner.L_cycles}, refinement {old_halt}->{mapper.refinement_steps}") def build_and_load_model(H): """Match helpers.train_helpers.load_imle: strip ``module.``, then DP main only.""" from helpers.train_helpers import restore_params local_rank = getattr(H, 'local_rank', 0) mpi_size = getattr(H, 'mpi_size', 1) strict = bool(H.load_strict) model = IMLE(H) if H.restore_path and os.path.isfile(H.restore_path): print(f" Loading model: {H.restore_path}") restore_params( model, H.restore_path, local_rank, mpi_size, map_cpu=True, strict=strict, ) else: print(" WARNING: restore_path not found or not set, using random weights!") model = torch.nn.DataParallel(model.cuda()) ema_model = None if getattr(H, 'restore_ema_path', None) and os.path.isfile(H.restore_ema_path): print(f" Loading EMA: {H.restore_ema_path}") ema_model = IMLE(H) restore_params( ema_model, H.restore_ema_path, local_rank, mpi_size, map_cpu=True, strict=strict, ) ema_model = ema_model.cuda() ema_model.requires_grad_(False) return model, ema_model # ────────────────────────────────────────────────────────────────────────── # Evaluation at a single noise scale # ────────────────────────────────────────────────────────────────────────── def evaluate_at_scale(model, H, noise_scale, output_dir, batch_size, num_fid, num_pr, num_grid, grid_nrow, sampler=None, grid_mode='iid'): results = {'noise_scale': noise_scale} if num_grid > 0: grid_path = os.path.join(output_dir, f"grid_scale_{noise_scale:.2f}.png") mode = (grid_mode or 'iid').lower() if mode == 'slerp': if sampler is None: raise ValueError( "grid_mode=slerp requires a Sampler (set_up_data + Sampler(...))") if grid_nrow <= 0 or num_grid % grid_nrow != 0: raise ValueError( f"For slerp, num_grid_samples ({num_grid}) must be divisible by " f"grid_nrow ({grid_nrow}); each row is one SLERP chain.") interp_steps = grid_nrow interp_pairs = num_grid // interp_steps generate_slerp_grid( model, H, sampler, interp_pairs, interp_steps, interp_steps, grid_path, noise_scale=noise_scale) else: generate_grid_image(model, H.latent_dim, num_grid, grid_nrow, grid_path, noise_scale=noise_scale) results['grid_path'] = grid_path ref_dir = f'{H.data_root}/img' # FID if num_fid > 0: fid_dir = os.path.join(output_dir, f"_tmp_fid_{noise_scale:.2f}") os.makedirs(fid_dir, exist_ok=True) t0 = time.time() generate_images_to_dir(model, H.latent_dim, num_fid, fid_dir, batch_size=batch_size, noise_scale=noise_scale) cur_fid = fid.compute_fid(ref_dir, fid_dir, verbose=False, num_workers=0) results['fid'] = cur_fid print(f" FID = {cur_fid:.4f} ({time.time() - t0:.1f}s)") shutil.rmtree(fid_dir, ignore_errors=True) # Precision / Recall if num_pr > 0: pr_dir = os.path.join(output_dir, f"_tmp_pr_{noise_scale:.2f}") os.makedirs(pr_dir, exist_ok=True) t0 = time.time() generate_images_to_dir(model, H.latent_dim, num_pr, pr_dir, batch_size=batch_size, noise_scale=noise_scale) try: from helpers.improved_precision_recall import compute_prec_recall precision, recall = compute_prec_recall(ref_dir, pr_dir) results['precision'] = precision results['recall'] = recall print(f" Precision = {precision:.4f}, Recall = {recall:.4f} " f"({time.time() - t0:.1f}s)") except ImportError: print(" WARNING: improved_precision_recall not found, skipping P/R") shutil.rmtree(pr_dir, ignore_errors=True) return results # ────────────────────────────────────────────────────────────────────────── # Main # ────────────────────────────────────────────────────────────────────────── def main(): H = Hyperparams() parser = argparse.ArgumentParser( description="Fewshot RS-IMLE: noise-resilience evaluation", formatter_class=argparse.RawDescriptionHelpFormatter, epilog=__doc__, ) parser = add_imle_arguments(parser) parser.add_argument('--noise_scales', type=str, default='1.0,1.5,2.0,2.5,3.0', help='Comma-separated noise scale factors for latent z ' '(1.0 = standard Gaussian, >1 = heavier tails)') parser.add_argument('--test_H_cycles', type=int, default=None) parser.add_argument('--test_L_cycles', type=int, default=None) parser.add_argument('--test_refinement_steps', type=int, default=None) # --num_fid_samples is defined in add_imle_arguments (hps.py) parser.add_argument('--num_pr_samples', type=int, default=1000) parser.add_argument('--num_grid_samples', type=int, default=40) parser.add_argument('--grid_nrow', type=int, default=8) parser.add_argument('--output_dir', type=str, default=None) parser.add_argument('--use_ema', default=False, type=lambda x: bool(strtobool(x))) parser.add_argument('--test_batch_size', type=int, default=16) parser.add_argument('--grid_only', action='store_true', help='Skip FID and Precision/Recall; only write sample grids. ' 'Sets num_fid_samples and num_pr_samples to 0 and ' 'noise_scales to 1.0.') parser.add_argument('--grid_mode', type=str, default='iid', choices=['iid', 'slerp'], help='iid: independent z samples (default). ' 'slerp: latent SLERP strips as in train.py / ' 'supplementary few-shot figure (needs num_grid_samples ' '= K * grid_nrow for K chains of length grid_nrow).') parse_args_and_update_hparams(H, parser) if H.grid_only: H.num_fid_samples = 0 H.num_pr_samples = 0 H.noise_scales = '1.0' noise_scales = [float(s.strip()) for s in H.noise_scales.split(',')] if H.output_dir: output_dir = H.output_dir else: output_dir = os.path.join(H.save_dir, 'noise_resilience') os.makedirs(output_dir, exist_ok=True) H, data_train, data_valid, preprocess_fn = set_up_data(H) print("=" * 70) if H.grid_only: print(" Fewshot RS-IMLE — Quality grid only (no FID / P/R)") print(f" grid_mode: {getattr(H, 'grid_mode', 'iid')}") else: print(" Fewshot RS-IMLE — Noise Resilience Evaluation") print("=" * 70) print(f" Dataset: {H.dataset}") print(f" Data root: {H.data_root}") print(f" Checkpoint: {H.restore_path}") print(f" EMA checkpoint: {getattr(H, 'restore_ema_path', None)}") print(f" use_ema: {H.use_ema}") print(f" use_rtm: {getattr(H, 'use_rtm', False)}") print(f" latent_dim: {H.latent_dim}") print(f" H/L/refinement: {H.H_cycles}/{H.L_cycles}/{H.refinement_steps}") if H.test_H_cycles is not None or H.test_L_cycles is not None: print(f" test override H/L/refinement: {H.test_H_cycles}/{H.test_L_cycles}/{H.test_refinement_steps}") print(f" Noise scales: {noise_scales}") print(f" FID samples: {H.num_fid_samples}") print(f" P/R samples: {H.num_pr_samples}") print(f" Output dir: {output_dir}") print("=" * 70) model, ema_model = build_and_load_model(H) if H.use_ema and ema_model is not None: eval_model = ema_model print(" Using EMA model for evaluation") else: eval_model = model if H.use_ema and ema_model is None: print(" WARNING: --use_ema True but no EMA loaded, using main model") print(" Using main model for evaluation") override_model_cycles( eval_model, test_H_cycles=H.test_H_cycles, test_L_cycles=H.test_L_cycles, test_refinement_steps=H.test_refinement_steps, ) eval_model.eval() sampler = None if H.grid_only and getattr(H, 'grid_mode', 'iid').lower() == 'slerp': sampler = Sampler(H, len(data_train), preprocess_fn) print(f" Sampler dataset len: {len(data_train)} (for IMLE buffers)") max_samples = max( H.num_fid_samples, H.num_pr_samples, H.num_grid_samples if H.num_grid_samples > 0 else 0, ) if max_samples <= 0: max_samples = max(H.test_batch_size, 16) batch_size = min(H.test_batch_size, max_samples) all_results = [] for scale in noise_scales: print(f"\n{'─' * 70}") print(f" Noise scale = {scale:.2f} (z ~ N(0, {scale:.2f}²))") print(f"{'─' * 70}") res = evaluate_at_scale( eval_model, H, scale, output_dir, batch_size, H.num_fid_samples, H.num_pr_samples, H.num_grid_samples, H.grid_nrow, sampler=sampler, grid_mode=getattr(H, 'grid_mode', 'iid'), ) all_results.append(res) # Build config block raw = eval_model.module if hasattr(eval_model, 'module') else eval_model mapper = raw.decoder.mapping_network inner = getattr(mapper, 'trm', None) output = { 'config': { 'restore_path': H.restore_path, 'restore_ema_path': getattr(H, 'restore_ema_path', None), 'use_ema': H.use_ema, 'dataset': H.dataset, 'data_root': H.data_root, 'use_rtm': bool(getattr(H, 'use_rtm', False)), 'latent_dim': H.latent_dim, 'H_cycles_trained': H.H_cycles, 'L_cycles_trained': H.L_cycles, 'refinement_steps_trained': H.refinement_steps, 'H_cycles_eval': inner.H_cycles if inner else H.H_cycles, 'L_cycles_eval': inner.L_cycles if inner else H.L_cycles, 'refinement_steps_eval': mapper.refinement_steps if hasattr(mapper, 'refinement_steps') else H.refinement_steps, 'num_fid_samples': H.num_fid_samples, 'num_pr_samples': H.num_pr_samples, 'grid_mode': getattr(H, 'grid_mode', 'iid'), }, 'results': all_results, } json_path = os.path.join(output_dir, "noise_resilience.json") with open(json_path, "w") as f: json.dump(output, f, indent=2) # Human-readable summary summary_lines = [] if H.grid_only: header = f"{'Scale':>7s} {'Grid':>50s}" sep = "─" * len(header) summary_lines.append(sep) summary_lines.append(header) summary_lines.append(sep) for r in all_results: s = r['noise_scale'] g = r.get('grid_path', 'N/A') summary_lines.append(f"{s:>7.2f} {g}") summary_lines.append(sep) else: header = (f"{'Scale':>7s} {'FID':>10s} " f"{'Precision':>10s} {'Recall':>10s}") sep = "─" * len(header) summary_lines.append(sep) summary_lines.append(header) summary_lines.append(sep) for r in all_results: s = r['noise_scale'] f_val = f"{r['fid']:.4f}" if 'fid' in r else "N/A" p_val = f"{r['precision']:.4f}" if 'precision' in r else "N/A" r_val = f"{r['recall']:.4f}" if 'recall' in r else "N/A" summary_lines.append( f"{s:>7.2f} {f_val:>10s} {p_val:>10s} {r_val:>10s}") summary_lines.append(sep) summary_text = "\n".join(summary_lines) summary_path = os.path.join(output_dir, "summary.txt") with open(summary_path, "w") as f: f.write(summary_text + "\n") print(f"\n{'=' * 70}") if H.grid_only: print(" QUALITY GRID OUTPUT") else: print(" NOISE RESILIENCE RESULTS") print(f"{'=' * 70}") print(summary_text) print(f"\n Full results: {json_path}") print(f" Summary: {summary_path}") print(f"{'=' * 70}") if __name__ == "__main__": main()