"""SentinelSharp verification head: SUPIR candidate + Sentinel-2 frames + DINOv3-SAT -> faithful 2.5 m. SUPIR gives a crisp candidate that is right about *what kind* of place it is and often wrong about *where exactly* things are, and sometimes invents objects. The head keeps SUPIR's crisp detail where the measurement supports it and rewrites it where it does not. It sees, at inference, only things that exist at inference: supir (B,3,256,256) SUPIR candidate, reflectance lr (B,N,4,64,64) Sentinel-2 R,G,B,NIR, every usable date, + validity mask + sub-pixel phases DINOv3-SAT tokens of the SUPIR candidate -> what the generator CLAIMS is there DINOv3-SAT tokens of the Sentinel-2 ref -> what the measurement SUPPORTS (scene semantics) and outputs a residual on top of the candidate (R,G,B from SUPIR, NIR from the upsampled reference frame) plus a per-pixel log-scale `s` (Laplace uncertainty; exp(s) = expected |error|). Layout: a U-Net at 256/128/64/32 px. The multi-frame Sentinel-2 encoder (masked pooling over dates, phase-conditioned) enters at 64 px -- its native grid. DINOv3 enters at 32 px: both token maps, their difference (the "generator disagrees with the evidence" signal) and a self-attention block that lets the disagreement propagate over the whole crop. """ from __future__ import annotations import torch import torch.nn as nn import torch.nn.functional as F def gn(c): return nn.GroupNorm(min(32, c // 4), c) class Res(nn.Module): def __init__(self, c, cout=None): super().__init__() cout = cout or c self.b = nn.Sequential(gn(c), nn.SiLU(), nn.Conv2d(c, cout, 3, 1, 1), gn(cout), nn.SiLU(), nn.Conv2d(cout, cout, 3, 1, 1)) self.skip = nn.Identity() if c == cout else nn.Conv2d(c, cout, 1) def forward(self, x): return self.skip(x) + self.b(x) class Attn(nn.Module): """Global self-attention over a (small) feature map.""" def __init__(self, c, heads=8): super().__init__() self.n = gn(c); self.a = nn.MultiheadAttention(c, heads, batch_first=True) self.ff = nn.Sequential(nn.LayerNorm(c), nn.Linear(c, 4 * c), nn.GELU(), nn.Linear(4 * c, c)) def forward(self, x): B, C, H, W = x.shape h = self.n(x).flatten(2).transpose(1, 2) t = h + self.a(h, h, h, need_weights=False)[0] t = t + self.ff(t) return x + (t - h).transpose(1, 2).reshape(B, C, H, W) # only the update is added to x class FrameEncoder(nn.Module): """Every Sentinel-2 date through shared convs, conditioned on its sub-pixel phase, then pooled over valid dates (masked mean + masked max): order-free, any number of dates.""" def __init__(self, cin=4, c=96): super().__init__() self.inp = nn.Conv2d(cin + 1, c, 3, 1, 1) # + a validity channel self.film = nn.Sequential(nn.Linear(2, c), nn.SiLU(), nn.Linear(c, 2 * c)) self.body = nn.Sequential(Res(c), Res(c)) self.out = nn.Conv2d(2 * c, c, 1) def forward(self, lr, valid, phases): B, N, C, h, w = lr.shape v = valid.float().view(B, N, 1, 1, 1) x = torch.cat([lr, v.expand(B, N, 1, h, w)], 2).flatten(0, 1) f = self.inp(x) g, b = self.film(phases.flatten(0, 1).float()).chunk(2, -1) f = self.body(f * (1 + g[..., None, None]) + b[..., None, None]).view(B, N, -1, h, w) m = v.expand_as(f[:, :, :1]) mean = (f * m).sum(1) / m.sum(1).clamp(min=1) mx = f.masked_fill(m.expand_as(f) == 0, -1e4).amax(1) return self.out(torch.cat([mean, mx], 1)) class VerificationHead(nn.Module): def __init__(self, dino_dim=1024, c=(64, 128, 192, 256), n_attn=2): super().__init__() c0, c1, c2, c3 = c self.frames = FrameEncoder(4, 96) self.stem = nn.Conv2d(3 + 4, c0, 3, 1, 1) # SUPIR RGB + bicubic S2 ref (RGBN) self.e0 = nn.Sequential(Res(c0), Res(c0)) self.d01 = nn.Conv2d(c0, c1, 4, 2, 1); self.e1 = nn.Sequential(Res(c1), Res(c1)) self.d12 = nn.Conv2d(c1, c2, 4, 2, 1); self.fuse64 = nn.Conv2d(c2 + 96, c2, 1); self.e2 = nn.Sequential(Res(c2), Res(c2)) self.d23 = nn.Conv2d(c2, c3, 4, 2, 1) self.dproj = nn.ModuleList([nn.Sequential(nn.LayerNorm(dino_dim), nn.Linear(dino_dim, c3)) for _ in range(2)]) self.fuse32 = nn.Conv2d(4 * c3, c3, 1) # enc, dino(supir), dino(s2), difference self.mid = nn.Sequential(Res(c3), *[Attn(c3) for _ in range(n_attn)], Res(c3)) # Bilinear x2 + 3x3 conv, NOT ConvTranspose2d: transposed convs produced checkerboard # patches (runs 1-3) that any texture-rewarding loss then amplified. up = lambda ci, co: nn.Sequential(nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False), nn.Conv2d(ci, co, 3, 1, 1)) self.u32 = up(c3, c2); self.g2 = nn.Sequential(Res(2 * c2, c2), Res(c2)) self.u21 = up(c2, c1); self.g1 = nn.Sequential(Res(2 * c1, c1), Res(c1)) self.u10 = up(c1, c0); self.g0 = nn.Sequential(Res(2 * c0, c0), Res(c0)) self.out = nn.Conv2d(c0, 4 + 1, 3, 1, 1) nn.init.zeros_(self.out.weight); nn.init.zeros_(self.out.bias) # starts as "return SUPIR" def _tok(self, t, i, hw): B, N, D = t.shape return self.dproj[i](t).transpose(1, 2).reshape(B, -1, *hw) def forward(self, supir, lr, valid, phases, ref_up, tok_supir, tok_s2, hw): """supir (B,3,256,256), lr (B,N,4,64,64), ref_up (B,4,256,256) bicubic reference frame, tok_* (B,T,1024) DINOv3 patch tokens on an hw grid. Returns (out (B,4,..), log-scale s).""" base = torch.cat([supir, ref_up[:, 3:4]], 1) x0 = self.e0(self.stem(torch.cat([supir, ref_up], 1))) x1 = self.e1(self.d01(x0)) x2 = self.e2(self.fuse64(torch.cat([self.d12(x1), self.frames(lr, valid, phases)], 1))) x3 = self.d23(x2) ds = F.interpolate(self._tok(tok_supir, 0, hw), size=x3.shape[-2:], mode="bilinear", align_corners=False) dm = F.interpolate(self._tok(tok_s2, 1, hw), size=x3.shape[-2:], mode="bilinear", align_corners=False) x3 = self.mid(self.fuse32(torch.cat([x3, ds, dm, ds - dm], 1))) y = self.g2(torch.cat([self.u32(x3), x2], 1)) y = self.g1(torch.cat([self.u21(y), x1], 1)) y = self.g0(torch.cat([self.u10(y), x0], 1)) o = self.out(y) return base + o[:, :4], o[:, 4:5] - 4.0 # s starts at log(0.018): typical |error| in reflectance class PatchDisc(nn.Module): """Spectral-norm PatchGAN on RGB: 'does this 2.5 m texture look like real SPOT imagery'.""" def __init__(self, c=64): super().__init__() sn = nn.utils.spectral_norm self.net = nn.Sequential( sn(nn.Conv2d(3, c, 4, 2, 1)), nn.LeakyReLU(0.2, True), sn(nn.Conv2d(c, 2 * c, 4, 2, 1)), nn.LeakyReLU(0.2, True), sn(nn.Conv2d(2 * c, 4 * c, 4, 2, 1)), nn.LeakyReLU(0.2, True), sn(nn.Conv2d(4 * c, 4 * c, 3, 1, 1)), nn.LeakyReLU(0.2, True), sn(nn.Conv2d(4 * c, 1, 3, 1, 1))) def forward(self, x): return self.net(x) class MultiScaleDisc(nn.Module): """Two PatchGAN critics: full resolution (fine texture, dots, blobs) and 2x downsampled (structure, smears). Returns one logit map per scale.""" def __init__(self, c=96): super().__init__() self.d = nn.ModuleList([PatchDisc(c), PatchDisc(c)]) def forward(self, x): return [self.d[0](x), self.d[1](F.avg_pool2d(x, 2))] class Dino(nn.Module): """Frozen DINOv3-SAT ViT-L/16. Preprocessing as verified in scripts/train_pmrf.py (DinoCond): reflectance -> per-channel affine onto the SAT-493M product range -> official SAT mean/std.""" GAIN = (1.494, 1.196, 1.098); BIAS = (0.2093, 0.2556, 0.1868) def __init__(self, name="facebook/dinov3-vitl16-pretrain-sat493m"): super().__init__() from transformers import AutoModel, AutoImageProcessor self.net = AutoModel.from_pretrained(name).eval() for q in self.net.parameters(): q.requires_grad_(False) pr = AutoImageProcessor.from_pretrained(name) self.nreg = int(self.net.config.num_register_tokens); self.ps = int(self.net.config.patch_size) self.register_buffer("mean", torch.tensor(pr.image_mean).view(1, 3, 1, 1)) self.register_buffer("std", torch.tensor(pr.image_std).view(1, 3, 1, 1)) self.register_buffer("gain", torch.tensor(self.GAIN).view(1, 3, 1, 1)) self.register_buffer("bias", torch.tensor(self.BIAS).view(1, 3, 1, 1)) def tokens(self, rgb, size): """rgb (B,3,H,W) reflectance -> patch tokens (B, T, D), grid (h, w). Gradient flows to rgb.""" x = (rgb * self.gain + self.bias).clamp(0, 1) x = (x - self.mean) / self.std x = F.interpolate(x, size=(size, size), mode="bilinear", align_corners=False) t = self.net(pixel_values=x).last_hidden_state[:, 1 + self.nreg:] return t, (size // self.ps, size // self.ps)