import numpy as np import pytest from conftest import solar_disc from suvi import corruptions, dataset, paths, samples SIZE = dataset.SHARD_SIZE class FakeShard: """A shard-shaped object holding frames in memory, so tests need no files.""" def __init__(self, satellite, times, size=32): self.satellite = satellite self.size = size self._frames = { t: np.stack([ solar_disc(size=size, radius=size // 3, peak=1.0 + band * 0.2 + t / 1e6) for band in range(6) ]).astype(np.float32) for t in times } def times(self): return sorted(self._frames) def frames(self, time): got = self._frames.get(time) return None if got is None else got.copy() def make_sampler(count=120, satellites=(16, 18), size=256, **kwargs): base = 1715400000 // paths.CADENCE * paths.CADENCE times = [base + i * paths.CADENCE for i in range(count)] shards = {("d", s): FakeShard(s, times, size) for s in satellites} return samples.Sampler(shards, satellites=satellites, **kwargs), times # --------------------------------------------------------------------------- layout def test_layout_excludes_the_frame_being_reconstructed(): layout = samples.stack_layout((16, 18), 16) assert (16, 0) not in layout def test_layout_includes_the_counterpart_at_the_target_instant(): """The single most valuable entry: a real observation of the right Sun, right time.""" assert (18, 0) in samples.stack_layout((16, 18), 16) assert (16, 0) in samples.stack_layout((16, 18), 18) def test_layout_is_multi_scale_in_both_directions(): layout = samples.stack_layout((16, 18), 16) for offset in samples.OFFSETS: assert (16, offset) in layout and (16, -offset) in layout assert (18, offset) in layout and (18, -offset) in layout def test_layout_is_stable(): """A frame's position in the stack must always mean the same thing.""" assert samples.stack_layout((16, 18), 16) == samples.stack_layout((16, 18), 16) assert len(set(samples.stack_layout((16, 18), 16))) == len( samples.stack_layout((16, 18), 16) ) def test_classes_cover_every_catalogued_mode_plus_clean(): assert samples.CLASSES[0] == "clean" assert set(samples.CLASSES[1:]) == set(corruptions.CATALOG) assert len(samples.CLASS_INDEX) == len(corruptions.CATALOG) + 1 # -------------------------------------------------------------------------- samples def test_build_returns_a_full_stack(): sampler, times = make_sampler() sample = sampler.build(16, times[60]) expected = len(samples.stack_layout((16, 18), 16)) assert sample["frames"].shape[0] >= expected assert sample["frames"].shape[1:] == (6, 256, 256) assert sample["target"].shape == (6, 256, 256) assert len(sample["states"]) == sample["frames"].shape[0] assert len(sample["dts"]) == sample["frames"].shape[0] assert len(sample["classes"]) == sample["frames"].shape[0] def test_build_returns_nothing_without_a_target(): sampler, times = make_sampler() assert sampler.build(16, times[-1] + 99 * paths.CADENCE) is None def test_target_is_never_in_its_own_stack(): """Leakage of the most direct kind: the answer among the inputs.""" sampler, times = make_sampler(damage_probability=0.0, drop_probability=0.0) sample = sampler.build(16, times[60]) for index, (dt, same) in enumerate(zip(sample["dts"], sample["same_satellite"])): assert not (dt == 0.0 and same), f"stack entry {index} is the target itself" def test_undamaged_sampler_marks_everything_available(): # 600 slots, so even the +/-256 rungs of the exponential ladder land on frames. sampler, times = make_sampler(count=600, size=32, damage_probability=0.0, drop_probability=0.0) sample = sampler.build(16, times[300]) assert set(sample["states"]) == {"available"} assert set(sample["classes"].tolist()) == {samples.CLASS_INDEX["clean"]} def test_missing_frames_are_zeroed_and_labelled(): """Off the end of the shard there is genuinely nothing.""" sampler, times = make_sampler(damage_probability=0.0, drop_probability=0.0) sample = sampler.build(16, times[1]) # offsets -2, -4, -16 fall off the start missing = [i for i, s in enumerate(sample["states"]) if s == "missing"] assert missing, "expected some offsets to fall outside the shard" for index in missing: assert np.all(sample["frames"][index] == 0.0) def test_anchors_reach_past_the_sampled_offsets(): """A long gap must still find something *real* to work from -- reach is measured over frames that carry pixels, not over empty ladder rungs.""" base = 1715400000 // paths.CADENCE * paths.CADENCE times = [base] + [base + i * paths.CADENCE for i in range(100, 130)] shards = {("d", 16): FakeShard(16, times, 256)} sampler = samples.Sampler(shards, satellites=(16,), damage_probability=0.0, drop_probability=0.0) sample = sampler.build(16, base) reach = max(abs(dt) for dt, state in zip(sample["dts"], sample["states"]) if state == "available") assert reach >= 100 * paths.CADENCE, "no anchor was added beyond the fixed offsets" def test_damage_produces_suspect_frames_that_keep_their_pixels(): """The design claim under test: a flagged frame is data, not a hole.""" sampler, times = make_sampler(damage_probability=1.0, drop_probability=0.0, seed=3) seen = set() for time in times[20:80]: sample = sampler.build(16, time) seen.update(sample["states"]) for index, state in enumerate(sample["states"]): if state == "suspect": assert np.any(sample["frames"][index] != 0.0) assert sample["classes"][index] != samples.CLASS_INDEX["clean"] assert "suspect" in seen def test_no_signal_modes_are_presented_as_missing(): """all_zero and friends leave nothing; calling them 'suspect' would be a lie.""" sampler, times = make_sampler(damage_probability=1.0, drop_probability=0.0, seed=5) for time in times[20:80]: sample = sampler.build(16, time) for index, state in enumerate(sample["states"]): label = samples.CLASSES[sample["classes"][index]] if label in samples.NO_SIGNAL: assert state == "missing" assert np.all(sample["frames"][index] == 0.0) def test_every_catalogued_mode_can_be_applied(): """Including the three that need a donor frame, which raise if given none.""" sampler, times = make_sampler(damage_probability=1.0, drop_probability=0.0, seed=11) applied = set() for time in times[10:110]: sample = sampler.build(16, time) applied.update(samples.CLASSES[c] for c in sample["classes"].tolist()) for mode, corruption in corruptions.CATALOG.items(): if corruption.needs_donor: assert mode in applied, f"{mode} never applied; donors may be unavailable" def test_damage_is_coherent_across_bands(): """One instrument makes all six bands, so a fault hits them together.""" radiance = np.stack([solar_disc(size=32, radius=10, peak=1.0 + i) for i in range(6)]) flipped = samples.decode_from_model( samples._damage(samples.encode_for_model(radiance), "yaw_flip", seed=1, severity=1.0) ) for band in range(6): expected, _ = corruptions.apply_array("yaw_flip", radiance[band], 1, 1.0) np.testing.assert_allclose(flipped[band], expected, rtol=2e-3, atol=1e-6) def test_damage_operates_on_radiance_not_model_space(): """A gain_shift multiplies a physical quantity; applying it to asinh values would model an entirely different fault.""" radiance = np.stack([np.full((16, 16), 2.0, np.float32) for _ in range(6)]) damaged = samples.decode_from_model( samples._damage(samples.encode_for_model(radiance), "gain_shift", seed=4, severity=1.0) ) expected, _ = corruptions.apply_array("gain_shift", radiance[0], 4, 1.0) ratio = float(damaged[0].mean() / radiance[0].mean()) np.testing.assert_allclose(damaged[0], expected, rtol=2e-3) assert abs(ratio - float(expected.mean() / 2.0)) < 1e-3 def test_sampling_is_reproducible_from_its_seed(): first, times = make_sampler(seed=7) second, _ = make_sampler(seed=7) a = first.build(16, times[60]) b = second.build(16, times[60]) assert a["states"] == b["states"] np.testing.assert_array_equal(a["classes"], b["classes"]) np.testing.assert_allclose(a["frames"], b["frames"]) def test_targets_lists_both_satellites(): sampler, times = make_sampler() found = sampler.targets() assert {s for s, _ in found} == {16, 18} assert len(found) == 2 * len(times) # ------------------------------------------------------------------------ encoding def test_encode_decode_round_trips(): values = np.array([[-3.7, -0.01, 0.0, 1e-3, 0.19, 50.0, 1231.0]], dtype=np.float32) restored = samples.decode_from_model(samples.encode_for_model(values)) np.testing.assert_allclose(restored, values, rtol=2e-3, atol=1e-7) def test_encoding_matches_the_shard_transform(): """Shards and live frames must land in the same space, or inference sees a shift.""" frames = np.stack([solar_disc(size=1280, radius=386, peak=1.0 + i) for i in range(6)]) from_shard = dataset.decode_frames(dataset.encode_frames(frames)) direct = samples.decode_from_model(samples.encode_for_model(from_shard)) np.testing.assert_allclose(direct, from_shard, rtol=2e-3, atol=1e-6) def test_encoding_stays_in_range_for_absurd_input(): coded = samples.encode_for_model(np.array([1e12, -1e12, np.nan, np.inf])) assert np.isfinite(coded).all() assert coded.max() <= 1.0 and coded.min() >= -1.0 # -------------------------------------------------------------------------- tensors def test_to_tensors_produces_what_the_model_expects(): torch = pytest.importorskip("torch") from suvi import align, model sampler, times = make_sampler(size=64) sample = sampler.build(16, times[60]) packed = samples.to_tensors(sample, torch) stack = sample["frames"].shape[0] assert packed["stack"].shape == (stack, 6, 64, 64) assert packed["condition"].shape == (stack, model.COND_DIM) assert packed["gains"].shape == (stack, 6) assert packed["offsets"].shape == (stack, 6) assert packed["target"].shape == (6, 64, 64) condition = packed["condition"][None] valid = (condition[..., 0] + condition[..., 2]).clamp(0, 1) aligned = align.align_stack( packed["stack"][None], packed["dts"][None], valid, packed["gains"][None], packed["offsets"][None], packed["b0"][None], packed["radius"][None], ) net = model.build(base=8, depth=2) out = net(aligned, condition) assert out.shape == (1, 6, 64, 64) assert torch.isfinite(out).all() # ---------------------------------------------------------------------- transfers def test_same_satellite_frames_carry_the_identity_transfer(): sampler, times = make_sampler(damage_probability=0.0, drop_probability=0.0) sample = sampler.build(16, times[60]) for index, same in enumerate(sample["same_satellite"]): if same: np.testing.assert_array_equal(sample["gains"][index], np.ones(6)) np.testing.assert_array_equal(sample["offsets"][index], np.zeros(6)) def test_cross_satellite_frames_carry_a_fitted_transfer(): """The counterpart must arrive with the day's calibration attached, so alignment can put it on the target instrument's scale before the model sees it.""" class Scaled(FakeShard): def frames(self, time): got = super().frames(time) return None if got is None else got * 0.8 + 0.01 base = 1715400000 // paths.CADENCE * paths.CADENCE times = [base + i * paths.CADENCE for i in range(40)] shards = {("d", 16): FakeShard(16, times, 32), ("d", 18): Scaled(18, times, 32)} sampler = samples.Sampler(shards, damage_probability=0.0, drop_probability=0.0) sample = sampler.build(16, times[20]) cross = [i for i, same in enumerate(sample["same_satellite"]) if not same and sample["states"][i] == "available"] assert cross, "no cross-satellite frame in the stack" for index in cross: np.testing.assert_allclose(sample["gains"][index], np.full(6, 1 / 0.8), rtol=1e-3) np.testing.assert_allclose(sample["offsets"][index], np.full(6, -0.01 / 0.8), atol=1e-4) def test_calibration_never_fits_against_the_target_itself(): """A pair at the target instant would fit the transfer against the answer -- the oracle gain the whole exercise exists to estimate honestly.""" sampler, times = make_sampler() for target in (times[0], times[60], times[-1]): sampler._transfers.clear() assert sampler._calibration("d", 16, 18, exclude=target) is not None for (_, _, _, pair_time) in sampler._transfers: assert pair_time != target def test_outages_are_contiguous_runs(): sampler, times = make_sampler() rng = np.random.default_rng(2) seen = 0 for _ in range(50): blocked = sampler._outages(rng, times[60]) for satellite, interval in blocked.items(): if not interval: continue seen += 1 ordered = sorted(interval) gaps = {b - a for a, b in zip(ordered, ordered[1:])} assert gaps <= {paths.CADENCE}, "outage is not a contiguous run" assert len(ordered) <= samples.MAX_OUTAGE_SLOTS assert seen, "no episodic outage was ever drawn" def test_dual_outages_take_down_both_satellites_sometimes(): sampler, times = make_sampler() rng = np.random.default_rng(3) dual = single = 0 for _ in range(200): blocked = sampler._outages(rng, times[60]) affected = [s for s, interval in blocked.items() if interval] if len(affected) == 2: assert blocked[16] == blocked[18], "dual outage must share one interval" dual += 1 elif len(affected) == 1: single += 1 assert dual > 10, f"dual outages too rare to train on ({dual}/200)" assert single > 10, f"single-satellite outages too rare ({single}/200)" def test_the_long_dual_outage_case_actually_occurs_in_training(): """The regime the archive says is 11% of reality: both satellites dark around the target, nearest real frame far away. With independent per-frame drops this configuration had probability ~p^14 and was never trained.""" sampler, times = make_sampler(count=600, size=32, seed=9) starved = 0 for time in times[280:380]: sample = sampler.build(16, time) if sample is None: continue near_same = [i for i, (dt, same) in enumerate(zip(sample["dts"], sample["same_satellite"])) if same and abs(dt) <= 16 * paths.CADENCE] cross = [i for i, (dt, same) in enumerate(zip(sample["dts"], sample["same_satellite"])) if not same and abs(dt) <= 16 * paths.CADENCE] if all(sample["states"][i] == "missing" for i in near_same) and \ all(sample["states"][i] == "missing" for i in cross): starved += 1 assert starved >= 2, f"long dual-outage stacks essentially absent ({starved}/100)" def test_anchors_land_outside_a_simulated_outage(): sampler, times = make_sampler(count=600, size=32, damage_probability=0.0, drop_probability=0.0) exclude = frozenset(times[300 + k] for k in range(-150, 150)) anchor = sampler._anchor(16, times[300], direction=1, exclude=exclude) assert anchor is not None and anchor not in exclude assert anchor >= times[300] + 150 * paths.CADENCE def test_build_declines_when_nothing_in_the_stack_has_pixels(): """Both satellites out for the whole window: the contract is None, not a fabricated frame.""" base = 1715400000 // paths.CADENCE * paths.CADENCE shards = {("d", 16): FakeShard(16, [base], 32)} sampler = samples.Sampler(shards, satellites=(16, 18), damage_probability=0.0, drop_probability=0.0) assert sampler.build(16, base) is None def test_build_reports_solar_geometry(): sampler, times = make_sampler() sample = sampler.build(16, times[60]) assert -0.13 < sample["b0"] < 0.13 # +/-7.25 deg in radians assert 0.28 < sample["radius"] < 0.32 # disc fraction of the frame def test_severity_is_drawn_per_sample_not_fixed_per_frame(): """Fixed rates left barely half of every stack clean, so the model never saw the easy case and learned to hedge. Real outages are episodic: some stacks should come through almost untouched and some should be wrecked.""" # A wide window, so every ladder rung exists and the clean fraction measures # injected damage rather than the shard's edges. sampler, times = make_sampler(count=600, size=32, damage_probability=0.30, drop_probability=0.40, seed=17) clean_fractions = [] for time in times[280:380]: sample = sampler.build(16, time) states = sample["states"] clean_fractions.append(sum(s == "available" for s in states) / len(states)) assert max(clean_fractions) > 0.95, "no sample came through nearly clean" assert min(clean_fractions) < 0.75, "no sample was substantially degraded" spread = max(clean_fractions) - min(clean_fractions) assert spread > 0.3, f"severity barely varied between samples (spread {spread:.2f})"