File size: 4,773 Bytes
17aade5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5878871
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
"""`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)