noaa-goes-visualization/suvi/fillers.py

526 lines
22 KiB
Python
Raw Normal View History

"""Methods for reconstructing a missing or rejected frame.
Every filler takes the same :class:`FillContext` and returns a replacement array, so
the bench can swap them without knowing which one it is holding. All are pure: they
read the context and return an array, nothing else.
The methods span a deliberate range of physical sophistication, from "repeat the last
good frame" (what the pipeline does today) to a differential-rotation warp that models
how the Sun actually moves. The bench exists to say which of them is worth the cost
at which gap length.
"""
from dataclasses import dataclass, field
import cv2 as cv
import numpy as np
# Shared with the torch-side alignment; defined there because suvi.align must import
# inside the ROCm container, which has no OpenCV.
from .align import (EARTH_ORBIT_DEG_PER_DAY, SECONDS_PER_DAY, SNODGRASS_A,
SNODGRASS_B, SNODGRASS_C, gain_fit)
#: Nominal solar radius in metres (IAU 2015).
R_SUN = 6.957e8
@dataclass
class FillContext:
"""Everything a filler may draw on to reconstruct one frame."""
#: Nearest good frame before the gap, and how many seconds back it sits.
before: np.ndarray | None = None
dt_before: float = 0.0
#: Nearest good frame after the gap, and how many seconds forward.
after: np.ndarray | None = None
dt_after: float = 0.0
#: The other satellite's view of this same instant, if it has one.
counterpart: np.ndarray | None = None
#: A (counterpart, this satellite) pair from the nearest slot where *both* were
#: good, used to calibrate one instrument against the other. The two spacecraft
#: observe simultaneously, so such a pair isolates the instrument difference with
#: no solar evolution mixed in -- which a bracketing frame from this satellite
#: alone cannot do. See :func:`crosssat`.
calibration: tuple | None = None
#: Header of the frame being reconstructed, for the WCS a rotation warp needs.
header: dict = field(default_factory=dict)
#: The full stack, for fillers that fuse more than two frames. One entry per input
#: frame: ``{"image", "state", "dt", "same_satellite", "scores", "verdict"}``, where
#: `state` is 'available', 'missing' or 'suspect'.
#:
#: The hand-written fillers ignore this and read only the fields above, which is why
#: adding it changes none of them. A **suspect** entry is the reason it exists:
#: every method in this module discards a flagged frame outright, but 14 of the 20
#: modes in :mod:`suvi.corruptions` leave one substantially usable, and the learned
#: filler is built to exploit exactly that.
stack: list = field(default_factory=list)
@property
def alpha(self):
"""Position within the gap: 0 at `before`, 1 at `after`."""
span = self.dt_before + self.dt_after
if span <= 0:
return 0.0
return self.dt_before / span
@property
def gap_frames(self):
"""Gap width in 4-minute slots, for reporting quality against gap length."""
return int(round((self.dt_before + self.dt_after) / 240.0))
def _finite(image):
return np.nan_to_num(np.asarray(image, dtype=np.float32), nan=0.0, posinf=0.0, neginf=0.0)
# ---------------------------------------------------------------- simple baselines
def hold_last(context):
"""Repeat the last good frame.
What ``merger_FITS.py`` does today (up to ``max_time_gap`` slots). Cheap, never
invents structure, but freezes the Sun and then jumps -- the visible stutter in
the current videos. The baseline every other method must beat.
"""
if context.before is not None:
return _finite(context.before)
if context.after is not None:
return _finite(context.after)
return None
def linear_blend(context):
"""Cross-fade between the frames bracketing the gap.
Removes the jump that ``hold_last`` leaves, at the cost of ghosting: moving
features appear twice, faintly, rather than moving.
"""
if context.before is None:
return hold_last(context)
if context.after is None or context.before.shape != context.after.shape:
# Frames of differing size cannot be mixed; the nearer one is the best
# available answer. This is also the fallback the other fillers unwind to.
return _finite(context.before)
alpha = context.alpha
return ((1.0 - alpha) * _finite(context.before) + alpha * _finite(context.after)).astype(
np.float32
)
# ------------------------------------------------------------------- optical flow
def _for_flow(image):
"""Compress radiance into a range optical flow can work with.
Radiance is heavy-tailed -- a flare can be 1000x the quiet corona -- so raw
values make flow chase the brightest pixels only. log1p plus a percentile
stretch keeps faint structure in play.
"""
scaled = np.log1p(np.clip(_finite(image), 0.0, None))
high = np.percentile(scaled, 99.5)
if high <= 0:
return np.zeros(scaled.shape, dtype=np.uint8)
return np.clip(scaled / high * 255.0, 0, 255).astype(np.uint8)
def optical_flow(context, use_dis=True):
"""Motion-compensated interpolation between the bracketing frames.
Estimates dense flow both ways and warps each bracket forward to the target
instant, then blends. Unlike ``linear_blend`` this moves features instead of
dissolving between them, which is what the eye reads as smooth motion.
"""
if context.before is None or context.after is None:
return linear_blend(context)
before, after = _finite(context.before), _finite(context.after)
if before.shape != after.shape:
return linear_blend(context)
first, second = _for_flow(before), _for_flow(after)
if use_dis:
engine = cv.DISOpticalFlow_create(cv.DISOPTICAL_FLOW_PRESET_MEDIUM)
forward = engine.calc(first, second, None)
backward = engine.calc(second, first, None)
else:
forward = cv.calcOpticalFlowFarneback(
first, second, None, 0.5, 3, 15, 3, 5, 1.2, 0
)
backward = cv.calcOpticalFlowFarneback(
second, first, None, 0.5, 3, 15, 3, 5, 1.2, 0
)
alpha = context.alpha
height, width = before.shape
grid_x, grid_y = np.meshgrid(
np.arange(width, dtype=np.float32), np.arange(height, dtype=np.float32)
)
# cv.remap samples the source *at* the map coordinates, so to place a feature
# where it should be at time alpha we read from where it was: a feature at x in
# `before` sits at x + forward(x) in `after`, hence at p - alpha*forward(p) when
# looking back from the interpolated frame. Adding the flow instead of
# subtracting it moves every feature the wrong way, which is worse than not
# compensating at all.
warped_before = cv.remap(
before,
grid_x - forward[..., 0] * alpha,
grid_y - forward[..., 1] * alpha,
cv.INTER_LINEAR,
borderMode=cv.BORDER_REPLICATE,
)
warped_after = cv.remap(
after,
grid_x - backward[..., 0] * (1.0 - alpha),
grid_y - backward[..., 1] * (1.0 - alpha),
cv.INTER_LINEAR,
borderMode=cv.BORDER_REPLICATE,
)
return ((1.0 - alpha) * warped_before + alpha * warped_after).astype(np.float32)
# ------------------------------------------------------------ cross-satellite fill
def gain_match(source, reference):
"""Put `source` on `reference`'s radiance scale by least squares.
GOES-16 and GOES-18 carry different SUVI flight models, so their radiances
differ by a roughly affine factor even when both are healthy.
"""
gain, offset = gain_fit(source, reference)
return (np.asarray(source, dtype=np.float32) * gain + offset).astype(np.float32)
def _shift(image, dx, dy):
matrix = np.array([[1.0, 0.0, dx], [0.0, 1.0, dy]], dtype=np.float32)
return cv.warpAffine(
np.asarray(image, dtype=np.float32),
matrix,
(image.shape[1], image.shape[0]),
flags=cv.INTER_LINEAR,
borderMode=cv.BORDER_REPLICATE,
)
def _register(source, reference):
"""Sub-pixel translation carrying `source` onto `reference`."""
window = cv.createHanningWindow((source.shape[1], source.shape[0]), cv.CV_64F)
(dx, dy), _ = cv.phaseCorrelate(
source.astype(np.float64), reference.astype(np.float64), window
)
return dx, dy
def crosssat(context, align=True):
"""Substitute the other satellite's view of the same instant.
The two spacecraft see the same Sun at the same moment, so the substitute is a
real observation rather than an interpolation, and its quality does not decay
with gap length the way every temporal method does. Measured against truth at
195A it scores 46.45 dB whether the gap is one slot or three hundred.
What does decay is the *calibration*. The instruments differ by a band-dependent
gain -- 0.81 at 195A, up to 1.49 at 94A, drifting 12% within a single week -- so
the substitute has to be put on this satellite's scale before it is usable, and
that gain has to be estimated from somewhere.
Estimating it from ``context.before`` is what this used to do, and it is wrong:
at a 300-slot gap that frame is twenty hours old, so the fit absorbs the Sun's
own evolution into what is supposed to be an instrument constant. The result
collapsed from 46.45 dB at gap 1 to 18.75 dB at gap 300 -- which is why the bench
reported this as the worst filler at every length, and why that finding was an
artifact of the estimator rather than a property of the method.
`context.calibration` instead supplies a *simultaneous* pair from the nearest
slot where both satellites were good. Because the two frames in that pair are of
the same Sun at the same instant, their ratio is the instrument difference and
nothing else, however far away the pair sits in time. Alignment is measured on
the same pair for the same reason.
"""
if context.counterpart is None:
return None
counterpart = _finite(context.counterpart)
if context.calibration is not None:
pair_counterpart, pair_local = context.calibration
if pair_counterpart is not None and pair_local is not None:
pair_counterpart = _finite(pair_counterpart)
pair_local = _finite(pair_local)
if pair_counterpart.shape == pair_local.shape == counterpart.shape:
if align:
dx, dy = _register(pair_counterpart, pair_local)
counterpart = _shift(counterpart, dx, dy)
pair_counterpart = _shift(pair_counterpart, dx, dy)
gain, offset = gain_fit(pair_counterpart, pair_local)
return (counterpart * gain + offset).astype(np.float32)
# No simultaneous pair anywhere in the series. Fall back to the bracketing
# frame, which is sound at short gaps and degrades as the bracket recedes.
reference = context.before if context.before is not None else context.after
if reference is None:
return counterpart
reference = _finite(reference)
if counterpart.shape != reference.shape:
return counterpart
if align:
counterpart = _shift(counterpart, *_register(counterpart, reference))
return gain_match(counterpart, reference)
# --------------------------------------------------------- solar rotation warping
def rotation_rate(latitude_rad, synodic=True):
"""Snodgrass differential rotation in degrees per day at a given latitude."""
sin2 = np.sin(latitude_rad) ** 2
rate = SNODGRASS_A + SNODGRASS_B * sin2 + SNODGRASS_C * sin2**2
return rate - EARTH_ORBIT_DEG_PER_DAY if synodic else rate
def _disc_radius_pixels(header, shape):
"""Solar radius in pixels, from the header if possible."""
diameter = header.get("diam_sun")
if diameter:
return float(diameter) / 2.0
distance, scale = header.get("dsun_obs"), header.get("cdelt1")
if distance and scale:
return float(np.degrees(np.arcsin(R_SUN / distance)) * 3600.0 / scale)
return shape[0] * 0.3 # falls back to the archive's typical disc fraction
def _rotation_map(shape, header, delta_seconds, synodic=True):
"""Inverse map: for each output pixel, where in the input it came from.
Works in heliographic coordinates -- de-project each pixel onto the sphere, undo
the rotation that happened over `delta_seconds`, re-project. Returns
(map_x, map_y, on_disc) with NaN where the source point is not visible.
"""
height, width = shape
radius = _disc_radius_pixels(header, shape)
crpix1 = float(header.get("crpix1", (width + 1) / 2.0)) - 1.0
crpix2 = float(header.get("crpix2", (height + 1) / 2.0)) - 1.0
b0 = np.radians(float(header.get("solar_b0", 0.0)))
grid_x, grid_y = np.meshgrid(np.arange(width), np.arange(height))
x = (grid_x - crpix1) / radius
y = (grid_y - crpix2) / radius
rho2 = x**2 + y**2
on_disc = rho2 < 1.0
z = np.sqrt(np.clip(1.0 - rho2, 0.0, None))
# Plane-of-sky -> heliographic, undoing the observer's B0 tilt.
sin_lat = y * np.cos(b0) + z * np.sin(b0)
sin_lat = np.clip(sin_lat, -1.0, 1.0)
latitude = np.arcsin(sin_lat)
longitude = np.arctan2(x, z * np.cos(b0) - y * np.sin(b0))
# Step the longitude back to where this material was `delta_seconds` ago.
days = delta_seconds / SECONDS_PER_DAY
source_longitude = longitude - np.radians(rotation_rate(latitude, synodic)) * days
# Heliographic -> plane-of-sky.
cos_lat = np.cos(latitude)
sx = cos_lat * np.sin(source_longitude)
sy = sin_lat * np.cos(b0) - cos_lat * np.cos(source_longitude) * np.sin(b0)
sz = sin_lat * np.sin(b0) + cos_lat * np.cos(source_longitude) * np.cos(b0)
visible = on_disc & (sz > 0)
map_x = (sx * radius + crpix1).astype(np.float32)
map_y = (sy * radius + crpix2).astype(np.float32)
return map_x, map_y, visible
def solar_rotation(context, synodic=True):
"""Warp the bracketing frames by differential solar rotation, then blend.
The Sun is not a rigid body: the equator turns in about 25 days, the poles in
about 35. Over a short gap that is a sub-pixel effect, but across a multi-hour
outage it is the difference between features landing where they belong and
smearing. This is the only method here that uses a physical model of the scene.
Applies on-disc only. The corona above the limb does not co-rotate with the
photosphere, so off-disc pixels fall back to a plain cross-fade.
"""
if context.before is None and context.after is None:
return None
if context.before is None or context.after is None:
source = context.before if context.before is not None else context.after
delta = context.dt_before if context.before is not None else -context.dt_after
warped, visible = _warp(source, context.header, delta, synodic)
blended = np.where(visible, warped, _finite(source))
return blended.astype(np.float32)
before, after = _finite(context.before), _finite(context.after)
if before.shape != after.shape:
return linear_blend(context)
# Roll `before` forward to the target instant and `after` backward to it.
warped_before, visible_before = _warp(before, context.header, context.dt_before, synodic)
warped_after, visible_after = _warp(after, context.header, -context.dt_after, synodic)
alpha = context.alpha
rotated = (1.0 - alpha) * warped_before + alpha * warped_after
faded = (1.0 - alpha) * before + alpha * after
visible = visible_before & visible_after
return np.where(visible, rotated, faded).astype(np.float32)
def _warp(image, header, delta_seconds, synodic):
image = _finite(image)
map_x, map_y, visible = _rotation_map(image.shape, header, delta_seconds, synodic)
warped = cv.remap(
image, map_x, map_y, cv.INTER_LINEAR, borderMode=cv.BORDER_CONSTANT, borderValue=0.0
)
return warped, visible
# ------------------------------------------------------------------ learned filler
#: Loaded checkpoint, kept between calls. The bench fills thousands of slots one at a
#: time, and reloading 14M parameters per slot would dominate the run.
_LEARNED = {}
#: Where to find the checkpoint, overridable so a bench run can name a specific one.
LEARNED_CHECKPOINT_ENV = "SUVI_MODEL"
def load_learned(path=None, device=None):
"""Load the trained filler, once. Returns (net, torch, device) or None."""
import os
path = path or os.environ.get(LEARNED_CHECKPOINT_ENV)
if not path:
return None
key = (path, device)
if key in _LEARNED:
return _LEARNED[key]
import torch
from . import model as model_module
state = torch.load(path, map_location="cpu", weights_only=False)
settings = state.get("args", {})
net = model_module.build(base=settings.get("base", 32), depth=settings.get("depth", 3))
try:
net.load_state_dict(state["model"])
except RuntimeError as error:
# Say which checkpoint and what changed. Torch's own message names tensor
# shapes and nothing else, which is unhelpful when several runs are on disk and
# only some predate an architecture change.
raise SystemExit(
f"{path} does not match the current model.\n{error}\n"
"This checkpoint was trained against a different architecture; retrain or "
"point SUVI_MODEL at a newer run."
) from error
resolved = torch.device(device or ("cuda" if torch.cuda.is_available() else "cpu"))
net.to(resolved).eval()
_LEARNED[key] = (net, torch, resolved)
return _LEARNED[key]
def assemble_stack(context, torch, device):
"""Turn a FillContext's stack into the model's aligned inputs.
Shared by the trained filler and by diagnostics (ceiling.py), so there is
exactly one implementation of the entry-to-tensor path. Returns
``(aligned, condition)`` -- the (1, S, 6, H, W) aligned candidates and their
(1, S, COND_DIM) conditioning -- or None when nothing in the stack carries
pixels.
"""
from . import align as align_module
from . import model as model_module
from . import samples
entries = [entry for entry in context.stack if entry.get("image") is not None
or entry.get("state") == "missing"]
shape = next((entry["image"].shape for entry in entries
if entry.get("image") is not None), None)
if shape is None:
return None
transfer = None
if context.calibration is not None:
pair_counterpart, pair_local = context.calibration
if pair_counterpart is not None and pair_local is not None:
transfer = align_module.fit_photometry(pair_counterpart, pair_local)
frames, conditions, dts, gains, offsets = [], [], [], [], []
for entry in entries:
image = entry.get("image")
state = entry.get("state", "available") if image is not None else "missing"
same = bool(entry.get("same_satellite", True))
frames.append(np.zeros((6, *shape[-2:]), dtype=np.float32) if image is None
else samples.encode_for_model(image))
conditions.append(model_module.frame_conditioning(
state, same, float(entry.get("dt", 0.0))
))
dts.append(float(entry.get("dt", 0.0)))
cross = transfer is not None and not same and image is not None
gains.append(transfer[0] if cross else np.ones(6, dtype=np.float32))
offsets.append(transfer[1] if cross else np.zeros(6, dtype=np.float32))
first = entries[0]
target_time = float(first["slot"][1]) - float(first.get("dt", 0.0))
b0, radius = align_module.solar_ephemeris(target_time)
with torch.no_grad():
stack = torch.from_numpy(np.stack(frames))[None].to(device)
condition = torch.stack(conditions)[None].to(device)
valid = (condition[..., 0] + condition[..., 2]).clamp(0, 1)
if float(valid.sum()) == 0:
return None
aligned = align_module.align_stack(
stack, torch.tensor(dts, device=device)[None], valid,
torch.from_numpy(np.stack(gains))[None].to(device),
torch.from_numpy(np.stack(offsets))[None].to(device),
torch.tensor([b0], device=device),
torch.tensor([radius], device=device),
)
return aligned, condition
def learned(context, path=None, device=None):
"""Reconstruct by fusing the whole stack, using the trained model.
This is the only filler that sees more than three frames, and the only one that
looks at a frame the detector flagged rather than discarding it. The stack is
first aligned deterministically -- rotation-warped to the target instant, and
cross-satellite frames put on this instrument's scale using the simultaneous
pair in ``context.calibration`` (all six bands) -- and the model then chooses,
per pixel and band, which aligned observation to trust.
Returns None when there is no checkpoint, no stack, or nothing in the stack
carries pixels; the bench reports that as "not applicable" rather than scoring
a fabricated frame.
"""
if not context.stack:
return None
loaded = load_learned(path, device)
if loaded is None:
return None
net, torch, device = loaded
from . import samples
assembled = assemble_stack(context, torch, device)
if assembled is None:
return None
aligned, condition = assembled
with torch.no_grad():
prediction = net(aligned, condition)
return samples.decode_from_model(prediction[0].float().cpu().numpy())
FILLERS = {
"hold_last": hold_last,
"linear_blend": linear_blend,
"optical_flow": optical_flow,
"crosssat": crosssat,
"solar_rotation": solar_rotation,
"learned": learned,
}