"""Deterministic alignment of a frame stack onto its target instant. The learned filler's job is *selection*, not physics. Everything about this archive that has a closed form is applied here, before any network sees a pixel: * **Rotation.** The Sun's differential rotation is known (Snodgrass 1983), so a frame at dt != 0 is warped forward or back to the target instant -- the same field :func:`suvi.fillers.solar_rotation` uses, ported to torch so one implementation serves 640-pixel training shards and 1280-pixel bench frames, on CPU or GPU. * **Photometry.** GOES-16 and GOES-18 SUVI differ by a band-dependent affine transfer (gain 0.81-1.49, drifting ~12%/week). A cross-satellite frame is put on the target instrument's scale using a (gain, offset) fitted from a simultaneous clean pair -- the corrected-`crosssat` estimator, verified at +16.5 dB over the stale-bracket fit. After this stage every stack entry is *an estimate of the target frame*, and the two views' geometry needs nothing more: the spacecraft sit 0.11 px of parallax apart on identical WCS grids, so there is no disparity to solve. Shards do not store WCS headers, so the geometry (solar B0 angle and apparent disc radius) comes from the analytic ephemeris below. Both are smooth annual functions of the date; the accuracy required is loose -- a 0.1 deg B0 error moves a 20-hour warp by under a tenth of a pixel. """ import math import numpy as np from . import dataset #: Snodgrass (1983) sidereal differential rotation, degrees per day, by latitude. #: Defined here rather than in :mod:`suvi.fillers` because this module must import #: inside the ROCm training container, which has torch but not OpenCV; `fillers` #: re-exports them so its callers see no change. SNODGRASS_A = 14.713 SNODGRASS_B = -2.396 SNODGRASS_C = -1.787 #: Earth's mean orbital motion, subtracted to get the rotation an Earth-orbiting #: observer actually sees. EARTH_ORBIT_DEG_PER_DAY = 0.9856 SECONDS_PER_DAY = 86400.0 #: Mean apparent solar radius at 1 AU, arcseconds (IAU 2015 nominal radius). RADIUS_ARCSEC_1AU = 959.63 #: SUVI L2 plate scale, arcsec/pixel. Constant across the archive: every frame is #: reprojected onto the same grid (CDELT1 = CDELT2 = 2.5, CROTA = 0). PLATE_SCALE = 2.5 #: Native SUVI L2 frame width, pixels. Geometry is expressed as a *fraction* of the #: frame so the same numbers serve 640-pixel shards and 1280-pixel archive frames. NATIVE_SIZE = 1280 #: Inclination of the solar equator to the ecliptic, degrees (Carrington). SOLAR_INCLINATION = 7.25 #: Unix time of the J2000.0 epoch. J2000_UNIX = 946728000.0 def solar_ephemeris(t_unix): """(b0 radians, disc radius as a fraction of frame width) for a unix time. Low-precision solar position (Meeus, Astronomical Algorithms ch. 25) -- good to ~0.1 deg in B0 and ~0.1% in distance, far inside what the rotation warp needs. B0 is the heliographic latitude of the disc centre: the Earth rides 7.25 deg above and below the solar equator over the year, and ignoring that tilts every latitude the differential-rotation profile is evaluated at. """ n = (t_unix - J2000_UNIX) / SECONDS_PER_DAY mean_longitude = math.radians((280.460 + 0.9856474 * n) % 360.0) mean_anomaly = math.radians((357.528 + 0.9856003 * n) % 360.0) ecliptic_longitude = mean_longitude + math.radians( 1.915 * math.sin(mean_anomaly) + 0.020 * math.sin(2 * mean_anomaly) ) distance_au = 1.00014 - 0.01671 * math.cos(mean_anomaly) \ - 0.00014 * math.cos(2 * mean_anomaly) # Ascending node of the solar equator on the ecliptic, precessing slowly. node = math.radians(73.6667 + 1.395833 * (n / 36525.0 + 1.5)) b0 = math.asin( math.sin(ecliptic_longitude - node) * math.sin(math.radians(SOLAR_INCLINATION)) ) radius_arcsec = RADIUS_ARCSEC_1AU / distance_au radius_fraction = radius_arcsec / PLATE_SCALE / NATIVE_SIZE return b0, radius_fraction # ------------------------------------------------------------------ rotation warp def rotation_grid(size, dt_seconds, b0, radius_fraction, synodic=True): """Sampling grids that undo `dt_seconds` of differential rotation. Torch port of :func:`suvi.fillers._rotation_map`, batched: `dt_seconds`, `b0` and `radius_fraction` are 1-D tensors of N frames, and the result is an (N, H, W, 2) grid in the normalised align_corners=False convention `grid_sample` expects. Where the source point is off-disc or behind the limb the grid holds the pixel's *own* centre, so sampling returns the unwarped value there -- the corona above the limb does not co-rotate with the photosphere, matching `solar_rotation`'s behaviour exactly. """ import torch height = width = int(size) dt = dt_seconds.reshape(-1, 1, 1).to(torch.float32) b0 = b0.reshape(-1, 1, 1).to(torch.float32) radius = (radius_fraction.reshape(-1, 1, 1) * width).to(torch.float32) device = dt.device centre_x = (width - 1) / 2.0 centre_y = (height - 1) / 2.0 grid_y, grid_x = torch.meshgrid( torch.arange(height, device=device, dtype=torch.float32), torch.arange(width, device=device, dtype=torch.float32), indexing="ij", ) x = (grid_x - centre_x) / radius y = (grid_y - centre_y) / radius rho2 = x**2 + y**2 on_disc = rho2 < 1.0 z = torch.sqrt((1.0 - rho2).clamp(min=0.0)) sin_b0, cos_b0 = torch.sin(b0), torch.cos(b0) sin_lat = (y * cos_b0 + z * sin_b0).clamp(-1.0, 1.0) latitude = torch.asin(sin_lat) longitude = torch.atan2(x, z * cos_b0 - y * sin_b0) sin2 = sin_lat**2 rate = SNODGRASS_A + SNODGRASS_B * sin2 + SNODGRASS_C * sin2**2 if synodic: rate = rate - EARTH_ORBIT_DEG_PER_DAY source_longitude = longitude - torch.deg2rad(rate) * (dt / SECONDS_PER_DAY) cos_lat = torch.cos(latitude) source_x = cos_lat * torch.sin(source_longitude) source_y = sin_lat * cos_b0 - cos_lat * torch.cos(source_longitude) * sin_b0 source_z = sin_lat * sin_b0 + cos_lat * torch.cos(source_longitude) * cos_b0 visible = on_disc & (source_z > 0) map_x = torch.where(visible, source_x * radius + centre_x, grid_x) map_y = torch.where(visible, source_y * radius + centre_y, grid_y) # Pixel-centre normalisation: (2p + 1)/n - 1 is the align_corners=False # convention; linspace(-1, 1) would shift everything by half a pixel and blur. grid = torch.stack( [(2 * map_x + 1) / width - 1, (2 * map_y + 1) / height - 1], dim=-1 ) return grid def rotate(frames, dt_seconds, b0, radius_fraction, synodic=True): """Warp (N, C, H, W) frames by their per-frame dt. Zero dt is the identity.""" import torch.nn.functional as F grid = rotation_grid(frames.shape[-1], dt_seconds, b0, radius_fraction, synodic) return F.grid_sample(frames, grid.to(frames.dtype), mode="bilinear", padding_mode="border", align_corners=False) # -------------------------------------------------------------------- photometry def gain_fit(source, reference): """Least-squares (gain, offset) putting `source` on `reference`'s scale.""" x = np.nan_to_num(np.asarray(source, dtype=np.float64), nan=0.0, posinf=0.0, neginf=0.0).ravel() y = np.nan_to_num(np.asarray(reference, dtype=np.float64), nan=0.0, posinf=0.0, neginf=0.0).ravel() variance = float(((x - x.mean()) ** 2).sum()) if variance <= 0: return 1.0, 0.0 gain = float(((x - x.mean()) * (y - y.mean())).sum() / variance) offset = float(y.mean() - gain * x.mean()) return gain, offset def apply_photometry(coded, gain, offset): """Apply a per-band radiance-space affine transfer to asinh-coded frames. `coded` is (N, BANDS, H, W) in the model's asinh space; `gain`/`offset` are (N, BANDS) in radiance units. The transfer is defined on radiance -- a gain multiplies a physical quantity -- so this computes ``asinh(sinh(x * R) * g + o / S) / R`` rather than scaling the coded values, which would model a different (and wrong) transform. """ import torch scaled = torch.sinh(coded.clamp(-1.0, 1.0) * dataset.ASINH_RANGE) moved = scaled * gain[..., None, None] + offset[..., None, None] / dataset.ASINH_SCALE return (torch.asinh(moved) / dataset.ASINH_RANGE).clamp(-1.0, 1.0) def fit_photometry(counterpart, local): """Per-band (gain, offset) putting `counterpart` on `local`'s radiance scale. Least squares on a simultaneous pair of (BANDS, H, W) radiance arrays -- the two spacecraft observe the same Sun at the same instant, so the fit isolates the instrument difference with no solar evolution mixed in. numpy, because it runs in the CPU data path (sampler and bench), once per day rather than per frame. """ pairs = [gain_fit(c, l) for c, l in zip(counterpart, local)] gains = np.array([g for g, _ in pairs], dtype=np.float32) offsets = np.array([o for _, o in pairs], dtype=np.float32) return gains, offsets # ------------------------------------------------------------------- whole stacks def align_stack(stack, dts, valid, gains, offsets, b0, radius_fraction, synodic=True): """Align every frame of a batch of stacks onto its target instant. `stack` is (B, S, BANDS, H, W) in asinh space; `dts` (B, S) seconds; `valid` (B, S) with 1 where a frame carries pixels; `gains`/`offsets` (B, S, BANDS) in radiance units (identity rows for same-satellite frames); `b0`/`radius_fraction` (B,) from :func:`solar_ephemeris` at the target instant. Photometry first (calibrate the instrument), then rotation (account for time). Invalid frames are forced to the identity transfer -- an offset applied to a frame of zeros would manufacture a constant image out of nothing. """ batch, stack_depth, bands = stack.shape[:3] flat = stack.reshape(batch * stack_depth, bands, *stack.shape[-2:]) keep = valid.reshape(-1, 1).to(flat.dtype) gain = gains.reshape(-1, bands) * keep + (1.0 - keep) offset = offsets.reshape(-1, bands) * keep needs_transfer = ((gain != 1.0) | (offset != 0.0)).any(dim=1) if bool(needs_transfer.any()): moved = apply_photometry(flat[needs_transfer], gain[needs_transfer], offset[needs_transfer]) flat = flat.clone() flat[needs_transfer] = moved.to(flat.dtype) expand = lambda values: values.reshape(batch, 1).expand(batch, stack_depth).reshape(-1) # `dts` is (frame time - target time); the warp must advance each frame by the # *negation* of that, (target - frame time), to land on the target instant. # Passing dts unnegated rotates every candidate AWAY from the target, doubling # the displacement instead of cancelling it -- sub-pixel at short gaps, which is # how it slipped past the identity and parity tests, and ~7 dB of candidate # quality at a 300-slot gap, which is how it was caught: the aligned anchors # scored far below plain solar_rotation on the same frames. aligned = rotate(flat, -dts.reshape(-1), expand(b0), expand(radius_fraction), synodic) return aligned.view_as(stack)