"""`infer/engine.py` tiling, batching and parallel TTA against a model whose answer is known. The stand-in model is pointwise (each output pixel depends only on the same input pixel), so it is exactly D4-equivariant and translation-invariant: the Hann-blended mosaic of overlapping tiles, at any batch size, with or without TTA, must reproduce the model applied to the whole scene in one go. Runs on the CPU in a second; no checkpoint needed. """ from __future__ import annotations import numpy as np import pytest torch = pytest.importorskip("torch") from config import N_SEG_CLASSES # noqa: E402 from dwdata.preprocess import PreprocSpec # noqa: E402 from infer.engine import count_tiles, predict_canonical # noqa: E402 SPEC = PreprocSpec(tile_size=64, patch=16) CPU = torch.device("cpu") class Pointwise: """height = sum of channels; class = 1 where channel 0 > 0 else 0.""" def __init__(self, oom_above: int | None = None): self.oom_above = oom_above self.batches: list[int] = [] def __call__(self, x): if self.oom_above is not None and x.shape[0] > self.oom_above: raise torch.cuda.OutOfMemoryError("simulated") self.batches.append(x.shape[0]) fused = x.sum(1, keepdim=True) seg = torch.zeros((x.shape[0], N_SEG_CLASSES, *x.shape[-2:])) seg[:, 1] = (x[:, 0] > 0).float() seg[:, 0] = 0.5 return {"fused": fused, "seg": seg} def scene(h, w, seed=0): return np.random.default_rng(seed).integers(0, 256, (h, w, 3), dtype=np.uint8) def reference(rgb): x = SPEC.normalise(rgb) return x.sum(0), (x[0] > 0).astype(np.int16) @pytest.mark.parametrize("batch_tiles", [0, 1, 3, 50]) def test_tiled_mosaic_matches_whole_scene(batch_tiles): rgb = scene(150, 211) height, seg = predict_canonical( Pointwise(), rgb, SPEC, CPU, overlap=0.5, batch_tiles=batch_tiles, want_seg=True, ) ref_h, ref_s = reference(rgb) assert height.shape == (150, 211) and seg.shape == (150, 211) np.testing.assert_allclose(height, ref_h, atol=1e-4) np.testing.assert_array_equal(seg, ref_s) def test_parallel_tta_stacks_all_views_in_one_batch(): rgb = scene(100, 130, seed=1) model = Pointwise() height, seg = predict_canonical( model, rgb, SPEC, CPU, tta=True, overlap=0.25, batch_tiles=2, want_seg=True, ) ref_h, ref_s = reference(rgb) np.testing.assert_allclose(height, ref_h, atol=1e-4) np.testing.assert_array_equal(seg, ref_s) # 2 tiles x 8 views per forward, not 8 forwards of 2 assert max(model.batches) == 16 def test_oom_halves_the_batch_and_finishes(): rgb = scene(150, 211, seed=2) model = Pointwise(oom_above=3) height, _ = predict_canonical(model, rgb, SPEC, CPU, overlap=0.5, batch_tiles=8) np.testing.assert_allclose(height, reference(rgb)[0], atol=1e-4) assert max(model.batches) <= 3 def test_scene_smaller_than_a_tile_is_padded_and_cropped(): rgb = scene(40, 30, seed=3) height, seg = predict_canonical(Pointwise(), rgb, SPEC, CPU, want_seg=True) assert height.shape == (40, 30) and seg.shape == (40, 30) np.testing.assert_allclose(height, reference(rgb)[0], atol=1e-4) def test_count_tiles_matches_the_grid(): # 150 x 211 at the canonical GSD, 64 px tiles at 50 % overlap -> 4 x 6 assert count_tiles(150, 211, SPEC.canonical_gsd_m, SPEC, 0.5) == 4 * 6 @pytest.mark.parametrize("tta", [False, True]) def test_scaled_inference_preserves_metres_labels_and_uncertainty(tta): class Constant: def __init__(self): self.shapes = [] def __call__(self, x): self.shapes.append(tuple(x.shape)) height = torch.full((x.shape[0], 1, *x.shape[-2:]), 7.0) seg = torch.zeros((x.shape[0], N_SEG_CLASSES, *x.shape[-2:])) seg[:, 3] = 1.0 return {"fused": height, "b_std": torch.full_like(height, 2.0), "seg": seg} model = Constant() height, seg, std = predict_canonical( model, scene(80, 110), SPEC, CPU, tta=tta, tta_scales=(1.5,), batch_tiles=2, want_seg=True, want_std=True, ) assert height.shape == seg.shape == std.shape == (80, 110) np.testing.assert_allclose(height, 7.0, atol=1e-4) np.testing.assert_allclose(std, 2.0, atol=1e-4) np.testing.assert_array_equal(seg, 3) assert all(shape[-2:] == (96, 96) for shape in model.shapes) assert max(shape[0] for shape in model.shapes) == (16 if tta else 2) @pytest.mark.parametrize("scales", [(), (0,), (-1,), (float("nan"),), (float("inf"),)]) def test_invalid_inference_scale_is_rejected(scales): with pytest.raises(ValueError, match="finite and positive"): predict_canonical(Pointwise(), scene(64, 64), SPEC, CPU, tta_scales=scales)