"""WP-7 refine bench: rank demo pool, refine best candidate, report LPIPS before/after on the K=2 fixtures. Run: python -m synth.refine_bench [--steps N] [--fixtures-dir DIR] [--limit N] The verbatim table (with --steps value) is pasted into MASTERPLAN WP-7 Result. """ from __future__ import annotations import argparse import time from typing import List, Tuple import numpy as np from PIL import Image from film_physics import get_film_curve from hybrid_loss import HybridFilmLoss from latent_optimizer import LatentSpaceOptimizer from synth.evaluation import score_pair from synth.generate import load_fixtures, select_k2_cases from app.preprocessing import preprocess_negative from app.api_client import generate_demo_candidates from app.scoring import rank_candidates def run_bench(steps: int = 60, fixtures_dir: str | None = None, limit: int = 6) -> None: """Run refinement on best-of-demo for each K=2 fixture (up to `limit`); print before/after table. Uses pixel fallback (no VAE). Reports per-case LPIPS (perm-invar), improved flag, and degeneracy indicator. Fixture-level smoke gate (6-case default): >=4/6 improved, 0 degenerate; the ≥70% accept item is judged on the 50-case run (WP-1.2). """ dir_note = f" --fixtures-dir {fixtures_dir}" if fixtures_dir else "" print(f"refine_bench --steps {steps}{dir_note} (limit {limit}): K=2 fixtures (demo pool -> best -> refine pixel fallback)") cases = load_fixtures(fixtures_dir) k2_cases = select_k2_cases(cases, limit) print(f"Using {len(k2_cases)} K=2 fixtures") rows: List[Tuple] = [] improved_count = 0 degen_count = 0 for idx, case in enumerate(k2_cases): scan = case["scan"] pil = Image.fromarray((scan * 255.0).clip(0, 255).astype(np.uint8)) pre = preprocess_negative(pil) if pre.density is None or pre.confidence_mask is None: print(f" case {idx}: skipped (no density)") continue stock = case.get("stock", "Generic") curve = get_film_curve(stock) loss_fn = HybridFilmLoss( film_curve=curve, physics_weight=1.0, perceptual_weight=0.5, ) # Demo pool + rank to pick best (same path as UI) demos = generate_demo_candidates(pre.rgb, num_candidates=6) ranked = rank_candidates( demos, observed_log_exposure=pre.log_exposure, observed_rgb=pre.rgb, film_curve=curve, physics_weight=1.0, perceptual_weight=0.5, density=pre.density, confidence_mask=pre.confidence_mask, ) if not ranked: continue best = ranked[0] before_a = best.separation.image_a before_b = best.separation.image_b # Refine (pixel fallback enforced) t0 = time.time() opt = LatentSpaceOptimizer( hybrid_loss=loss_fn, steps=steps, lr=0.05, max_side=256, vae_id="nonexistent-pixel-fallback", ) result = opt.refine( recon_a=before_a, recon_b=before_b, observed_log_exposure=pre.log_exposure, observed_rgb=pre.rgb, observed_density=pre.density, confidence_mask=pre.confidence_mask, ) dt = time.time() - t0 after_a = result.refined_a after_b = result.refined_b sc_before = score_pair(case["gt_a"], case["gt_b"], before_a, before_b) sc_after = score_pair(case["gt_a"], case["gt_b"], after_a, after_b) lp_before = float(sc_before.get("lpips", float("nan"))) lp_after = float(sc_after.get("lpips", float("nan"))) improved = lp_after < lp_before - 1e-6 if improved: improved_count += 1 # degeneracy_indicator: low value (~0) means one layer near-black (I.6) degen_ind = float(sc_after.get("degeneracy_indicator", 0.5)) degen = degen_ind < 0.05 if degen: degen_count += 1 rows.append((idx, case["seed"], lp_before, lp_after, improved, result.steps_run, dt, degen)) print( f"case{idx} seed{case['seed']}: before={lp_before:.4f} after={lp_after:.4f} " f"improved={improved} steps={result.steps_run} t={dt:.1f}s degen={degen}" ) print("\nPer-case table (demo-best before vs refined after):") print("| case | seed | LPIPS_before | LPIPS_after | improved | steps_run | runtime_s | degen |") print("|------|------|--------------|-------------|----------|-----------|-----------|-------|") for r in rows: imp = "yes" if r[4] else "no" dg = "yes" if r[7] else "no" print(f"| {r[0]} | {r[1]} | {r[2]:.4f} | {r[3]:.4f} | {imp} | {r[5]} | {r[6]:.1f} | {dg} |") print("Note: degen = (degeneracy_indicator < 0.05) per synth/evaluation.") n = len(rows) if n: m_before = float(np.mean([r[2] for r in rows])) m_after = float(np.mean([r[3] for r in rows])) print(f"\nMean LPIPS: before={m_before:.4f} after={m_after:.4f} (improved {improved_count}/{n})") print(f"Degenerate outputs: {degen_count}") print("Bench complete.") def main() -> None: p = argparse.ArgumentParser(description="WP-7 refine bench (latent optimizer hardening)") p.add_argument("--steps", type=int, default=60, help="refine steps per case (default 60 per spec)") p.add_argument("--fixtures-dir", type=str, default=None, dest="fixtures_dir", help="Directory with case_*.npz fixtures (default: synth/fixtures; for 256px re-measure)") p.add_argument("--limit", type=int, default=6, help="Max K=2 cases to process (default 6; for 50-case bench)") p.add_argument("--bench", action="store_true", help="run the fixture bench (default if no arg)") args = p.parse_args() if args.bench or True: run_bench(steps=args.steps, fixtures_dir=args.fixtures_dir, limit=args.limit) if __name__ == "__main__": main()