noaa-goes-visualization/tests/test_model.py

337 lines
14 KiB
Python
Raw Normal View History

import numpy as np
import pytest
torch = pytest.importorskip("torch")
from suvi import model # noqa: E402
BANDS = model.BANDS
def conditioning(states, satellites=None, dts=None):
"""Build a (1, S, COND_DIM) conditioning tensor."""
count = len(states)
satellites = satellites or [True] * count
dts = dts if dts is not None else [(i - count // 2) * 240.0 for i in range(count)]
rows = [
model.frame_conditioning(states[i], satellites[i], dts[i])
for i in range(count)
]
return torch.stack(rows)[None]
def stack(count, size=64, seed=0):
generator = torch.Generator().manual_seed(seed)
return torch.randn(1, count, BANDS, size, size, generator=generator)
@pytest.fixture(scope="module")
def net():
torch.manual_seed(0)
return model.build(base=16, depth=3).eval()
# ------------------------------------------------------------------- conditioning
def test_conditioning_is_the_declared_width():
assert model.frame_conditioning("available", True, 0.0).shape == (model.COND_DIM,)
def test_conditioning_separates_the_three_states():
vectors = [model.frame_conditioning(s, True, 240.0)[:3]
for s in ("available", "missing", "suspect")]
for index, vector in enumerate(vectors):
assert vector[index] == 1.0 and vector.sum() == 1.0
def test_conditioning_rejects_an_unknown_state():
with pytest.raises(KeyError):
model.frame_conditioning("probably-fine", True, 0.0)
def test_conditioning_distinguishes_direction_and_distance():
back = model.frame_conditioning("available", True, -7200.0)
forward = model.frame_conditioning("available", True, 7200.0)
far = model.frame_conditioning("available", True, 72000.0)
assert back[4] == -1.0 and forward[4] == 1.0
assert back[5] == pytest.approx(forward[5]) # magnitude, not direction
assert far[5] > forward[5]
def test_conditioning_leaves_agreement_to_the_model():
"""Columns 6 and 7 are computed from the stack inside forward(), identically at
training and inference -- the fix for detector scores that were always zero in
training and populated on the bench."""
vector = model.frame_conditioning("available", True, 240.0)
assert vector[6] == 0.0 and vector[7] == 0.0
# ------------------------------------------------------------------- blend prior
def test_prior_favours_the_nearest_frame_in_time():
condition = conditioning(["available"] * 3, satellites=[True] * 3,
dts=[-240.0, -3840.0, -240.0 * 300])
prior = model.blend_prior(condition)[0]
assert prior[0] > prior[1] > prior[2]
def test_prior_no_longer_penalises_the_counterpart():
"""The old prior down-weighted cross-satellite frames because their gain was
wrong at initialisation. Candidates now arrive gain-matched by suvi.align, so
a simultaneous counterpart is as trustworthy as a local frame."""
same = conditioning(["available"], satellites=[True], dts=[0.0])
cross = conditioning(["available"], satellites=[False], dts=[0.0])
assert model.blend_prior(same)[0, 0] == model.blend_prior(cross)[0, 0]
def test_prior_downweights_suspect_frames_without_excluding_them():
clean = conditioning(["available"], satellites=[True], dts=[-240.0])
suspect = conditioning(["suspect"], satellites=[True], dts=[-240.0])
assert model.blend_prior(clean)[0, 0] > model.blend_prior(suspect)[0, 0]
assert torch.isfinite(model.blend_prior(suspect)).all()
# ------------------------------------------------------------------------- shapes
def test_forward_returns_the_source_shape(net):
out = net(stack(5), conditioning(["available"] * 5))
assert out.shape == (1, BANDS, 64, 64)
assert torch.isfinite(out).all()
@pytest.mark.parametrize("size", [160, 320, 640])
def test_one_model_serves_every_source_resolution(net, size):
"""Blend and polish are fields: they upsample, so 1280 needs no retraining."""
sources = torch.randn(1, 3, BANDS, size, size)
out = net(sources, conditioning(["available"] * 3))
assert out.shape == (1, BANDS, size, size)
assert torch.isfinite(out).all()
@pytest.mark.parametrize("count", [1, 2, 7, 23])
def test_stack_depth_is_not_fixed(net, count):
"""Availability varies per slot, so the stack cannot be a compile-time constant."""
out = net(stack(count), conditioning(["available"] * count))
assert out.shape == (1, BANDS, 64, 64)
# --------------------------------------------------------------- identity behaviour
def expected_blend(sources, condition):
"""What an untrained model must output: the prior-weighted blend of its sources.
Both heads are zero-initialised, so the polish gain is exactly one and the blend
logits are exactly the prior."""
valid = (condition[..., 0] + condition[..., 2]).clamp(0, 1)
logits = model.blend_prior(condition).masked_fill(valid < 0.5, float("-inf"))
weights = torch.nan_to_num(torch.softmax(logits, dim=1), nan=0.0)
return (sources * weights[..., None, None, None]).sum(dim=1)
def test_untrained_model_is_the_prior_weighted_blend(net):
"""The wiring gate: any scaling or orientation bug in the composition shows up
here, and is invisible once training starts moving the loss."""
frames = stack(5)
condition = conditioning(["available"] * 5)
out = net(frames, condition)
assert torch.allclose(out, expected_blend(frames, condition), atol=1e-4)
def test_identity_holds_when_fields_are_upsampled(net):
"""A 160-resolution field applied to larger sources must still be the identity."""
sources = torch.randn(1, 3, BANDS, 320, 320)
condition = conditioning(["available"] * 3)
out = net(sources, condition)
assert torch.allclose(out, expected_blend(sources, condition), atol=1e-3)
def test_missing_frames_take_no_weight(net):
"""A frame with no pixels must not contribute, however confident the head is."""
frames = stack(4)
frames[:, 0] = 999.0 # would dominate any blend that included it
condition = conditioning(["missing", "available", "available", "available"])
out = net(frames, condition)
assert torch.allclose(out, expected_blend(frames, condition), atol=1e-4)
assert out.abs().max() < 100, "the masked frame leaked into the blend"
def test_suspect_frames_do_take_weight(net):
"""The design decision under test: a flagged frame is data, not a hole."""
frames = stack(3)
excluded = conditioning(["missing", "available", "available"])
included = conditioning(["suspect", "available", "available"])
a = net(frames, excluded)
b = net(frames, included)
assert not torch.allclose(a, b, atol=1e-3)
assert torch.allclose(b, expected_blend(frames, included), atol=1e-4)
def test_a_stack_with_nothing_usable_returns_zeros_not_nan(net):
"""The sampler declines such targets; the network itself must still not NaN."""
out = net(stack(4), conditioning(["missing"] * 4))
assert torch.isfinite(out).all()
assert out.abs().max() == 0.0
def test_heads_start_at_zero():
"""Training must start exactly at the hand-written policy; every step afterwards
is a departure the data paid for."""
torch.manual_seed(0)
fresh = model.build(base=8, depth=2)
for head in (fresh.weight, fresh.polish):
with torch.no_grad():
assert float(head.weight.abs().sum()) == 0.0
assert float(head.bias.abs().sum()) == 0.0
# --------------------------------------------------------- structural non-collapse
def test_output_is_a_bounded_combination_whatever_the_weights():
"""The invariant that makes the old failure mode unreachable: with *arbitrary*
parameters, the output of a single-frame stack stays within the polish bound of
that frame. No parameter setting can emit an image that is not built from the
inputs -- which is exactly what the collapsed run did."""
torch.manual_seed(3)
wild = model.build(base=8, depth=2)
with torch.no_grad():
for parameter in wild.parameters():
parameter.normal_(0.0, 2.0)
frames = torch.full((1, 1, BANDS, 48, 48), 0.5)
with torch.no_grad():
out = wild(frames, conditioning(["available"], dts=[240.0]))
bound = 0.5 * float(np.exp(model.POLISH_RANGE))
lower = 0.5 * float(np.exp(-model.POLISH_RANGE))
assert float(out.max()) <= bound + 1e-4
assert float(out.min()) >= lower - 1e-4
def test_polish_gain_saturates_at_its_bound():
torch.manual_seed(0)
fresh = model.build(base=8, depth=2).eval()
with torch.no_grad():
fresh.polish.bias.fill_(1000.0)
frames = torch.ones(1, 1, BANDS, 48, 48)
out = fresh(frames, conditioning(["available"], dts=[240.0]))
assert float(out.mean()) == pytest.approx(np.exp(model.POLISH_RANGE), rel=1e-3)
# --------------------------------------------------------------------- agreement
def test_agreement_flags_the_frame_that_disagrees():
"""A frozen or torn frame announces itself as deviation from the stack median --
the deterministic replacement for both cross-attention and detector scores."""
torch.manual_seed(0)
fresh = model.build(base=8, depth=2)
frames = torch.zeros(1, 4, BANDS, 32, 32)
frames[:, 3] = 5.0 # the odd one out
condition = conditioning(["available"] * 4)
valid = (condition[..., 0] + condition[..., 2]).clamp(0, 1)
got = fresh._with_agreement(frames.reshape(4, BANDS, 32, 32), condition, valid,
batch=1, stack=4)
assert float(got[0, 3, 6]) > float(got[0, 0, 6]) + 0.5
assert float(got[0, 0, 6]) < 0.1
def test_agreement_is_zero_for_a_unanimous_stack():
torch.manual_seed(0)
fresh = model.build(base=8, depth=2)
frames = torch.full((1, 3, BANDS, 32, 32), 0.25)
condition = conditioning(["available"] * 3)
valid = (condition[..., 0] + condition[..., 2]).clamp(0, 1)
got = fresh._with_agreement(frames.reshape(3, BANDS, 32, 32), condition, valid,
batch=1, stack=3)
assert float(got[..., 6:].abs().max()) < 1e-6
def test_agreement_ignores_missing_frames():
"""A zeroed missing frame must not drag the median toward black."""
torch.manual_seed(0)
fresh = model.build(base=8, depth=2)
frames = torch.full((1, 3, BANDS, 32, 32), 0.5)
frames[:, 0] = 0.0 # missing, zero-filled
condition = conditioning(["missing", "available", "available"])
valid = (condition[..., 0] + condition[..., 2]).clamp(0, 1)
got = fresh._with_agreement(frames.reshape(3, BANDS, 32, 32), condition, valid,
batch=1, stack=3)
assert float(got[0, 1, 6]) < 1e-6
assert float(got[0, 2, 6]) < 1e-6
# ------------------------------------------------------------------------ training
def test_gradients_reach_every_head():
"""A head with no gradient path never learns, and zero-init hides it. Checked
after three optimiser steps, because the zero initialisation is nested two deep:
at step 1 only the output heads move, at step 2 the FiLM output layers behind
them, and only then does gradient reach everything. Asserting earlier would
fail on a model that is perfectly healthy."""
torch.manual_seed(0)
fresh = model.build(base=8, depth=2)
frames = stack(3, size=48)
condition = conditioning(["available", "suspect", "available"])
optimiser = torch.optim.Adam(fresh.parameters(), lr=1e-3)
for _ in range(3):
optimiser.zero_grad()
out = fresh(frames, condition)
out.square().mean().backward()
optimiser.step()
for name, parameter in fresh.named_parameters():
assert parameter.grad is not None, f"{name} has no gradient"
assert parameter.grad.abs().sum() > 0, f"{name} gradient is identically zero"
def smooth_stack(count, size=64, seed=1):
"""Spatially correlated frames; white noise tests nothing the real task needs."""
generator = torch.Generator().manual_seed(seed)
coarse = torch.randn(1, count * BANDS, size // 8, size // 8, generator=generator)
smooth = torch.nn.functional.interpolate(coarse, size=(size, size), mode="bicubic",
align_corners=False)
return smooth.view(1, count, BANDS, size, size)
def train_a_little(net, frames, condition, target, steps=120, lr=3e-3):
optimiser = torch.optim.Adam(net.parameters(), lr=lr)
first = None
for step in range(steps):
optimiser.zero_grad()
out = net(frames, condition)
loss = (out - target).abs().mean()
loss.backward()
optimiser.step()
if step == 0:
first = loss.item()
return first, loss.item()
def test_model_learns_to_select_one_frame_from_the_stack():
"""The weight head's core job: when one candidate is the answer, put the weight
there. Fails in seconds if the softmax, masking or composition is miswired."""
torch.manual_seed(0)
fresh = model.build(base=16, depth=2)
frames = smooth_stack(3)
first, last = train_a_little(
fresh, frames, conditioning(["available"] * 3), frames[:, 1].clone()
)
assert last < first * 0.2, f"loss barely moved: {first:.4f} -> {last:.4f}"
def test_model_learns_a_per_band_gain_within_the_polish_bound():
"""The polish head's core job: the residual drift the daily fit leaves behind."""
torch.manual_seed(0)
fresh = model.build(base=16, depth=2)
frames = smooth_stack(1)
per_band = torch.tensor([1.2, 0.85, 1.1, 0.9, 1.25, 0.8]).view(1, BANDS, 1, 1)
target = frames[:, 0] * per_band
first, last = train_a_little(fresh, frames, conditioning(["available"],
dts=[240.0]), target)
assert last < first * 0.3, f"loss barely moved: {first:.4f} -> {last:.4f}"