Spaces:
Running on Zero
Running on Zero
Download tests/test_engine.py from akashch1512/SingleViewHeigthEstimation: direct link, hf CLI and curl.
- Browser
- Download file 4.77 kB
-
https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/tests/test_engine.py
- Command line
-
hf download hf://spaces/akashch1512/SingleViewHeigthEstimation/tests/test_engine.py
-
curl -L -o test_engine.py https://huggingface.co/spaces/akashch1512/SingleViewHeigthEstimation/resolve/main/tests/test_engine.py
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) | |
| 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 | |
| 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) | |
| 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) | |