Spaces:
Running on Zero
Running on Zero
| """WP-0: Tests for LatentSpaceOptimizer — pixel fallback path and VAE path (slow).""" | |
| from __future__ import annotations | |
| import numpy as np | |
| import pytest | |
| import torch | |
| from film_physics import get_film_curve | |
| from hybrid_loss import HybridFilmLoss | |
| from latent_optimizer import LatentSpaceOptimizer, RefinementResult | |
| def _make_small_pair(size: int = 32) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: | |
| """Return (img_a, img_b, obs_rgb) all as float32 (H, W, 3) in [0.2, 0.8].""" | |
| rng = np.random.default_rng(42) | |
| img_a = rng.uniform(0.2, 0.8, (size, size, 3)).astype(np.float32) | |
| img_b = rng.uniform(0.2, 0.8, (size, size, 3)).astype(np.float32) | |
| obs_rgb = rng.uniform(0.2, 0.8, (size, size, 3)).astype(np.float32) | |
| return img_a, img_b, obs_rgb | |
| def _build_obs_log_h(size: int = 32) -> torch.Tensor: | |
| return torch.full((1, 1, size, size), -0.3) | |
| def _build_loss_fn() -> HybridFilmLoss: | |
| curve = get_film_curve("Generic") | |
| return HybridFilmLoss( | |
| film_curve=curve, | |
| physics_weight=1.0, | |
| perceptual_weight=0.0, # skip LPIPS so pixel-fallback tests stay fast | |
| ) | |
| class TestPixelFallback: | |
| """Pixel-space optimization runs offline without a VAE.""" | |
| def test_returns_refinement_result(self, monkeypatch): | |
| monkeypatch.setattr( | |
| LatentSpaceOptimizer, "_try_load_vae", lambda self: None | |
| ) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=3, lr=0.01) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert isinstance(result, RefinementResult) | |
| assert not result.used_vae | |
| assert result.steps_run == 3 | |
| def test_output_shape_matches_input(self, monkeypatch): | |
| monkeypatch.setattr( | |
| LatentSpaceOptimizer, "_try_load_vae", lambda self: None | |
| ) | |
| img_a, img_b, obs_rgb = _make_small_pair(size=48) | |
| obs_log_h = _build_obs_log_h(size=48) | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=2, lr=0.01) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert result.refined_a.shape == img_a.shape | |
| assert result.refined_b.shape == img_b.shape | |
| def test_output_values_in_range(self, monkeypatch): | |
| monkeypatch.setattr( | |
| LatentSpaceOptimizer, "_try_load_vae", lambda self: None | |
| ) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=3, lr=0.01) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert float(result.refined_a.min()) >= 0.0 - 1e-5 | |
| assert float(result.refined_a.max()) <= 1.0 + 1e-5 | |
| assert float(result.refined_b.min()) >= 0.0 - 1e-5 | |
| assert float(result.refined_b.max()) <= 1.0 + 1e-5 | |
| def test_improved_flag(self, monkeypatch): | |
| monkeypatch.setattr( | |
| LatentSpaceOptimizer, "_try_load_vae", lambda self: None | |
| ) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=10, lr=0.05) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| # The 'improved' property is just final < initial — verify it's a bool | |
| assert isinstance(result.improved, bool) | |
| def test_progress_callback_called(self, monkeypatch): | |
| monkeypatch.setattr( | |
| LatentSpaceOptimizer, "_try_load_vae", lambda self: None | |
| ) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| calls = [] | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=4, lr=0.01) | |
| result = opt.refine( | |
| img_a, img_b, obs_log_h, obs_rgb, | |
| progress_callback=lambda step, total, loss: calls.append((step, total, loss)), | |
| ) | |
| assert len(calls) == 4 | |
| assert all(isinstance(l, float) for _, _, l in calls) | |
| def test_refine_large_nonmultiple_with_density_no_crash(self, monkeypatch): | |
| """Fix 1 regression: 700x900 (>max_side, not %8) with density+mask must not crash.""" | |
| monkeypatch.setattr( | |
| LatentSpaceOptimizer, "_try_load_vae", lambda self: None | |
| ) | |
| h, w = 700, 900 | |
| rng = np.random.default_rng(42) | |
| img_a = rng.uniform(0.2, 0.8, (h, w, 3)).astype(np.float32) | |
| img_b = rng.uniform(0.2, 0.8, (h, w, 3)).astype(np.float32) | |
| obs_rgb = rng.uniform(0.2, 0.8, (h, w, 3)).astype(np.float32) | |
| obs_log_h = torch.full((1, 1, h, w), -0.3) | |
| density = np.full((h, w), 0.5, dtype=np.float32) | |
| conf_mask = np.full((h, w), 1, dtype=np.uint8) # VALID | |
| loss_fn = _build_loss_fn() # perceptual=0 | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=1, lr=0.05, max_side=512, vae_id="nonexistent") | |
| result = opt.refine( | |
| img_a, img_b, obs_log_h, obs_rgb, | |
| observed_density=density, confidence_mask=conf_mask, | |
| ) | |
| assert isinstance(result, RefinementResult) | |
| assert np.isfinite(result.initial_loss) | |
| assert np.isfinite(result.final_loss) | |
| assert result.refined_a.shape == (h, w, 3) | |
| class TestVAEPath: | |
| """VAE path — requires downloading model weights, marked slow/offline-skip.""" | |
| def test_vae_path_runs(self): | |
| img_a, img_b, obs_rgb = _make_small_pair(size=64) | |
| obs_log_h = _build_obs_log_h(size=64) | |
| curve = get_film_curve("Generic") | |
| loss_fn = HybridFilmLoss(film_curve=curve, physics_weight=1.0, perceptual_weight=0.0) | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=2, lr=0.01) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| # If VAE was loaded, used_vae=True; if not available, falls back gracefully | |
| assert isinstance(result, RefinementResult) | |
| assert result.refined_a.shape == img_a.shape | |
| class TestLatentOptimizerWP7Hardening: | |
| """WP-7: cosine LR, best-snapshot, early-stop on physics, degeneracy guard, g_final, new flags.""" | |
| def test_new_refinement_result_fields_present(self, monkeypatch): | |
| monkeypatch.setattr(LatentSpaceOptimizer, "_try_load_vae", lambda self: None) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=2, lr=0.01) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert hasattr(result, "g_final") | |
| assert isinstance(result.g_final, float) | |
| assert hasattr(result, "early_stopped") | |
| assert hasattr(result, "degeneracy_aborted") | |
| assert result.used_density in (True, False) | |
| def test_early_stop_with_lr_zero(self, monkeypatch): | |
| """With lr=0 loss never improves => early stop after ~patience steps, flag set.""" | |
| monkeypatch.setattr(LatentSpaceOptimizer, "_try_load_vae", lambda self: None) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=200, lr=0.0) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert result.early_stopped is True | |
| assert result.steps_run <= 30 + 35 # initial + patience | |
| assert result.steps_run < 200 | |
| def test_g_final_reported_and_finite(self, monkeypatch): | |
| monkeypatch.setattr(LatentSpaceOptimizer, "_try_load_vae", lambda self: None) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=3, lr=0.01) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert np.isfinite(result.g_final) | |
| assert 0.01 < result.g_final < 100.0 # within reasonable after clamp | |
| def test_degeneracy_guard_monkey_k_and_fields(self, monkeypatch): | |
| """Monkey k_selection to K=2 high; guard fields exercised (trigger force demonstrated via teeth).""" | |
| monkeypatch.setattr(LatentSpaceOptimizer, "_try_load_vae", lambda self: None) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| class FakeBD: | |
| k_selection_score = 0.9 | |
| monkeypatch.setattr(loss_fn, "evaluate", lambda *a, **k: FakeBD()) | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=4, lr=0.05) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert result.degeneracy_aborted in (True, False) | |
| assert isinstance(result.g_final, float) | |
| def test_best_snapshot_matches_evaluated_loss(self, monkeypatch): | |
| """Regression (WP-7 review bugs 1+3): the returned state must be one whose loss | |
| was actually evaluated. With a divergent lr the first steps explode, so a correct | |
| best-snapshot returns ~the initial state (final_loss <= initial_loss). The buggy | |
| ordering snapshotted post-step latents under the pre-step loss, returning an | |
| exploded state labeled with a good loss.""" | |
| monkeypatch.setattr(LatentSpaceOptimizer, "_try_load_vae", lambda self: None) | |
| img_a, img_b, obs_rgb = _make_small_pair() | |
| obs_log_h = _build_obs_log_h() | |
| loss_fn = _build_loss_fn() | |
| opt = LatentSpaceOptimizer(hybrid_loss=loss_fn, steps=5, lr=50.0) | |
| result = opt.refine(img_a, img_b, obs_log_h, obs_rgb) | |
| assert result.final_loss <= result.initial_loss * 1.05 + 1e-6, ( | |
| f"returned state's loss {result.final_loss:.6f} is worse than initial " | |
| f"{result.initial_loss:.6f} — best-snapshot returned a state that never earned its loss" | |
| ) | |