Download head.py from Harp404/sentinel-sharp: direct link, hf CLI and curl.
- Browser
- Download file 9.03 kB
-
https://huggingface.co/Harp404/sentinel-sharp/resolve/main/head.py
- Command line
-
hf download hf://Harp404/sentinel-sharp/head.py
-
curl -L -o head.py https://huggingface.co/Harp404/sentinel-sharp/resolve/main/head.py
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) | |