Spaces:
Sleeping
Sleeping
| """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() | |