File size: 9,033 Bytes
78e667d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
"""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)