akashch1512's picture
fix: match v5 checkpoint architecture and benchmark scaled inference
5878871 verified
Raw History Blame Contribute Delete
4.77 kB
"""`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)