sentinel-sharp / head.py
Harp404's picture
SentinelSharp: verification head and all pipeline weights
78e667d
Raw History Blame Contribute Delete
9.03 kB
"""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)