double-exposure / synth /refine_bench.py
Eddie Faillace
double-exposure app deploy snapshot 2026-07-21 (WP-24 calibration pass)
7dff04f
Raw
History Blame Contribute Delete
6.04 kB
"""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()