File size: 17,983 Bytes
71d64bb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
"""
T-SRDA model β€” AD-STCLN backbone with PDAViT-style Temporal
Spatial-Reduction Dual-Attention (Zhou et al., Neurocomputing 2026)
replacing the DaViT/Swin temporal encoder.

Pipeline (identical I/O contract to the det experiment):

  UTAE_ASPP                encoder:  (B,T,C,H,W), (B,T) days
                              -> out (B,H,W,T,D), x2 (B*T,D,H,W)
  UTAEPrediction           pretrain wrapper: NDVI-aware masking + MSE recon
                              -> (recon, recon2, target, mask)  [1=visible]
  UTAEClassificationDual   finetune wrapper: STA temporal collapse +
                              dual decoder (semantic + boundary) + gated
                              refinement -> dict(sem_logits, bnd_logits,
                              refined_logits)
  compute_boundary_target  morphological-gradient boundary GT (B,H,W)

Temporal encoder change (the ONLY architectural difference vs det/DaViT):
  TemporalDaViTEncoder -> TemporalSRDAEncoder (temporal_srda.py)
    KV path: Linear d->d/R, unfold R steps -> (T/R, d), first SA (K-V
    self-checking); second SA: full-T queries x reduced KV -> global
    temporal receptive field per layer.

Checkpointing: chunk-level only (TEMPORAL_CHUNK sequences per checkpointed
chunk). The temporal encoder has NO internal checkpointing β€” layer-level +
chunk-level double-checkpointing wasted ~30% backward compute in the DaViT
version.
"""
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint

import config as C
from temporal_srda import TemporalSRDAEncoder

TEMPORAL_CHUNK = 1024      # per-pixel sequences per checkpointed chunk (train)
EVAL_TEMPORAL_CHUNK = 4096   # larger in eval β€” no checkpointing to pay for
SPATIAL_CHUNK = 32           # frames per spatial-encoder chunk in eval
#
# Eval-mode chunking exists for the official test protocol: test_STCLN.py feeds
# WHOLE 128x128 patches, so one batch-4 sample is 4*61 = 244 frames and
# 4*128*128 = 65,536 per-pixel temporal sequences. Unchunked, the ASPP concat
# alone (1024 ch at 128x128) would need ~16 GB. Training is unaffected: at
# 32x32 crops the chunk paths below are never entered.


# ═══════════════════════════════════════════════════════════════════════════
# 1. Positional encoding over real acquisition days
# ═══════════════════════════════════════════════════════════════════════════
class PositionalEncoder(nn.Module):
    """Sinusoidal PE over day offsets. days: (..., T) -> (..., T, d)."""

    def __init__(self, d_model, T=1000):
        super().__init__()
        self.d_model = d_model
        denom = torch.pow(
            T, 2.0 * (torch.arange(d_model).div(2, rounding_mode="floor"))
               / d_model)
        self.register_buffer("denom", denom, persistent=False)

    def forward(self, days):
        pe = days[..., None] / self.denom            # (..., T, d)
        pe = torch.stack([torch.sin(pe[..., 0::2]),
                          torch.cos(pe[..., 1::2])], dim=-1)
        return pe.flatten(-2)                        # interleave sin/cos


# ═══════════════════════════════════════════════════════════════════════════
# 2. Atrous spatial encoder: (B*T, C_in, H, W) -> (B*T, D, H, W)
# ═══════════════════════════════════════════════════════════════════════════
class SEBlock(nn.Module):
    def __init__(self, channels, reduction=8):
        super().__init__()
        self.fc = nn.Sequential(
            nn.AdaptiveAvgPool2d(1),
            nn.Conv2d(channels, channels // reduction, 1),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels // reduction, channels, 1),
            nn.Sigmoid(),
        )

    def forward(self, x):
        return x * self.fc(x)


class ASPP(nn.Module):
    """Atrous spatial pyramid pooling neck at full 32x32 resolution."""

    def __init__(self, in_ch, out_ch, dilations=(1, 2, 4)):
        super().__init__()
        self.branches = nn.ModuleList([
            nn.Sequential(
                nn.Conv2d(in_ch, out_ch, 3, padding=d, dilation=d, bias=False),
                nn.BatchNorm2d(out_ch),
                nn.ReLU(inplace=True))
            for d in dilations
        ])
        self.pool = nn.Sequential(
            nn.AdaptiveAvgPool2d(1),
            nn.Conv2d(in_ch, out_ch, 1, bias=False),
            nn.ReLU(inplace=True),
        )
        self.project = nn.Sequential(
            nn.Conv2d(out_ch * (len(dilations) + 1), out_ch, 1, bias=False),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        feats = [b(x) for b in self.branches]
        gp = self.pool(x)
        feats.append(F.interpolate(gp, size=x.shape[-2:], mode="nearest"))
        return self.project(torch.cat(feats, dim=1))


class AtrousSpatialEncoder(nn.Module):
    """Dilated conv stem + ASPP + SE. Keeps full spatial resolution."""

    def __init__(self, in_channels=10, d_model=256):
        super().__init__()
        mid = d_model // 4                                    # 64
        self.stem = nn.Sequential(
            nn.Conv2d(in_channels, mid, 3, padding=1, bias=False),
            nn.BatchNorm2d(mid),
            nn.ReLU(inplace=True),
            nn.Conv2d(mid, mid, 3, padding=2, dilation=2, bias=False),
            nn.BatchNorm2d(mid),
            nn.ReLU(inplace=True),
        )
        self.aspp = ASPP(mid, d_model, dilations=(1, 2, 4))
        self.se = SEBlock(d_model)

    def forward(self, x):                                     # (N, C, H, W)
        return self.se(self.aspp(self.stem(x)))               # (N, D, H, W)


# ═══════════════════════════════════════════════════════════════════════════
# 3. Encoder wrapper: spatial per-frame + T-SRDA temporal per-pixel
# ═══════════════════════════════════════════════════════════════════════════
class UTAE_ASPP(nn.Module):
    """
    forward(x, pos):
      x   : (B, T, C, H, W)  normalized S2 crops
      pos : (B, T)           day offsets (0 on padded steps)
    returns:
      out : (B, H, W, T, D)  temporally contextualized per-pixel features
      x2  : (B*T, D, H, W)   spatial features (deep-supervision recon path)
    """

    def __init__(self, n_channels=10, d_model=256, n_heads=8,
                 T_pad=48, windows=None, shifts=None, **kwargs):
        super().__init__()
        # T_pad / windows / shifts are absorbed for constructor compatibility
        # with pretrain.py / finetune.py / evaluate.py; T-SRDA ignores them
        # (the encoder self-pads T to a multiple of max(R_SCHEDULE)).
        self.d_model = d_model
        self.spatial_encoder = AtrousSpatialEncoder(n_channels, d_model)
        self.pos_encoder = PositionalEncoder(d_model)
        self.temporal_encoder = TemporalSRDAEncoder(
            d_model=d_model, n_heads=n_heads, R_schedule=C.R_SCHEDULE)

    def _run_temporal_chunk(self, x_chunk, pos_chunk):
        return self.temporal_encoder(x_chunk, pos_chunk)

    def _encode_spatial(self, frames):
        """(N, C, H, W) -> (N, D, H, W). Chunked in eval to cap the ASPP peak."""
        N = frames.shape[0]
        if self.training or N <= SPATIAL_CHUNK:
            return self.spatial_encoder(frames)
        return torch.cat([self.spatial_encoder(frames[s:s + SPATIAL_CHUNK])
                          for s in range(0, N, SPATIAL_CHUNK)], dim=0)

    def forward(self, x, pos):
        B, T, Cc, H, W = x.shape

        # ---- spatial encoding, frame by frame ------------------------------
        x2 = self._encode_spatial(x.reshape(B * T, Cc, H, W))   # (B*T,D,H,W)
        D = x2.shape[1]

        # ---- to per-pixel temporal sequences -------------------------------
        x_te = (x2.reshape(B, T, D, H, W)
                  .permute(0, 3, 4, 1, 2)                       # (B,H,W,T,D)
                  .reshape(B * H * W, T, D))

        pe = self.pos_encoder(pos)                              # (B, T, D)
        pos_exp = (pe[:, None, None, :, :]
                   .expand(B, H, W, T, D)
                   .reshape(B * H * W, T, D))

        # ---- temporal encoder, chunked + checkpointed in training ----------
        N = x_te.shape[0]
        if self.training and N > TEMPORAL_CHUNK:
            outs = []
            for s in range(0, N, TEMPORAL_CHUNK):
                e = min(s + TEMPORAL_CHUNK, N)
                outs.append(checkpoint(self._run_temporal_chunk,
                                       x_te[s:e], pos_exp[s:e],
                                       use_reentrant=False))
            out = torch.cat(outs, dim=0)
        elif not self.training and N > EVAL_TEMPORAL_CHUNK:
            outs = []
            for s in range(0, N, EVAL_TEMPORAL_CHUNK):
                e = min(s + EVAL_TEMPORAL_CHUNK, N)
                outs.append(self.temporal_encoder(x_te[s:e], pos_exp[s:e]))
            out = torch.cat(outs, dim=0)
        else:
            out = self.temporal_encoder(x_te, pos_exp)          # (B*H*W,T,D)

        out = out.reshape(B, H, W, T, D)
        return out, x2


# ═══════════════════════════════════════════════════════════════════════════
# 4. Pretrain wrapper: NDVI-aware masking + MSE reconstruction
#    (logic recovered from the det pipeline β€” mask: 1=visible, 0=masked)
# ═══════════════════════════════════════════════════════════════════════════
class UTAEPrediction(nn.Module):
    def __init__(self, encoder, n_channels=10, d_model=256, mask_ratio=0.4):
        super().__init__()
        self.utae = encoder
        self.linear = nn.Linear(d_model, n_channels)    # main reconstruction
        self.midlinear = nn.Linear(d_model, n_channels)  # deep supervision
        self.mask_ratio = mask_ratio
        self.mask_token = nn.Parameter(torch.zeros(1))

    def forward(self, x, pos):
        """Official STCLN masking β€” STCLN.py:193-202, bit-for-bit.

        The NDVI indicator is per-timestep-per-pixel, but `clusterLmean`
        averages it over H and W, so the gate is PER FRAME: a timestep whose
        frame is <= CLOUD_GATE vegetated is left ENTIRELY visible. This is
        NOT a per-pixel vegetation filter, and getting that wrong changes the
        pretext task (measured: the official gate leaves 96.2% of all values
        visible, a per-pixel version leaves 82.3%).
        """
        # x: (B, T, C, H, W)   pos: (B, T)
        target = x.clone()
        B, T, Cc, H, W = x.shape

        with torch.no_grad():
            # NDVI indicator on normalized bands: (NIR b6 - Red b2)/(Red + NIR)
            ndvi = (x[:, :, 6] - x[:, :, 2]) / (x[:, :, 2] + x[:, :, 6] + 1e-20)
            cluster = ndvi.gt(C.NDVI_THRESH).float()          # (B, T, H, W)

            # F.dropout: kept -> scaled ones, dropped -> zeros; * (1-p) undoes
            # the scaling so the mask is exactly {0, 1}. 1 = visible.
            mask = torch.ones(B, T, H, W, device=x.device)
            mask = F.dropout(mask, p=self.mask_ratio, training=True) \
                   * (1.0 - self.mask_ratio)
            mask_4d = mask[:, :, None].repeat(1, 1, Cc, 1, 1)

            # per-FRAME vegetation fraction -> exempt whole frames from masking
            frame_veg = cluster.mean(dim=[2, 3])              # (B, T)
            exempt = (frame_veg <= C.CLOUD_GATE)[:, :, None, None, None]
            mask_4d = torch.where(exempt, torch.ones_like(mask_4d), mask_4d)

        x_masked = x * mask_4d + self.mask_token * (1.0 - mask_4d)

        out, x2 = self.utae(x_masked, pos)      # (B,H,W,T,D), (B*T,D,H,W)

        recon = self.linear(out).permute(0, 3, 4, 1, 2)          # (B,T,C,H,W)
        recon2 = (self.midlinear(x2.permute(0, 2, 3, 1))
                  .permute(0, 3, 1, 2)
                  .reshape(B, T, Cc, H, W))
        return recon, recon2, target, mask_4d


# ═══════════════════════════════════════════════════════════════════════════
# 5. Finetune wrapper: STA temporal collapse + dual decoder + gated refine
# ═══════════════════════════════════════════════════════════════════════════
class STA(nn.Module):
    """Spatio-temporal aggregation: attention-pool T, refine spatially.
    (Reconstruction of base STCLN's conv1/conv2/conv3/l collapse module.)"""

    def __init__(self, d_model=256):
        super().__init__()
        self.l = nn.Linear(d_model, 1)                 # temporal attn scores
        self.conv1 = nn.Sequential(
            nn.Conv2d(d_model, d_model, 3, padding=1, bias=False),
            nn.BatchNorm2d(d_model), nn.ReLU(inplace=True))
        self.conv2 = nn.Sequential(
            nn.Conv2d(d_model, d_model, 3, padding=1, bias=False),
            nn.BatchNorm2d(d_model), nn.ReLU(inplace=True))
        self.conv3 = nn.Sequential(
            nn.Conv2d(d_model, d_model, 1, bias=False),
            nn.BatchNorm2d(d_model), nn.ReLU(inplace=True))

    def forward(self, feats):                          # (B, H, W, T, D)
        w = self.l(feats).softmax(dim=3)               # (B, H, W, T, 1)
        f = (feats * w).sum(dim=3)                     # (B, H, W, D)
        f = f.permute(0, 3, 1, 2).contiguous()         # (B, D, H, W)
        f = f + self.conv2(self.conv1(f))              # residual refine
        return self.conv3(f)


class SemanticDecoder(nn.Module):
    def __init__(self, d_model=256, num_classes=20):
        super().__init__()
        self.net = nn.Sequential(
            nn.Conv2d(d_model, d_model // 2, 3, padding=1, bias=False),
            nn.BatchNorm2d(d_model // 2), nn.ReLU(inplace=True),
            nn.Conv2d(d_model // 2, num_classes, 1))

    def forward(self, x):
        return self.net(x)


class BoundaryDecoder(nn.Module):
    def __init__(self, d_model=256):
        super().__init__()
        self.net = nn.Sequential(
            nn.Conv2d(d_model, d_model // 4, 3, padding=1, bias=False),
            nn.BatchNorm2d(d_model // 4), nn.ReLU(inplace=True),
            nn.Conv2d(d_model // 4, 1, 1))

    def forward(self, x):
        return self.net(x)                             # (B, 1, H, W)


class GatedRefinement(nn.Module):
    """Boundary-gated fusion of features and semantic logits."""

    def __init__(self, d_model=256, num_classes=20):
        super().__init__()
        self.fuse = nn.Sequential(
            nn.Conv2d(d_model + num_classes, d_model // 2, 3, padding=1,
                      bias=False),
            nn.BatchNorm2d(d_model // 2), nn.ReLU(inplace=True),
            nn.Conv2d(d_model // 2, num_classes, 1))

    def forward(self, feat, sem_logits, bnd_logits):
        gate = torch.sigmoid(bnd_logits)               # (B, 1, H, W)
        fused = torch.cat([feat * (1.0 + gate), sem_logits], dim=1)
        return sem_logits + self.fuse(fused)           # residual refinement


class UTAEClassificationDual(nn.Module):
    def __init__(self, encoder, d_model=256, num_classes=20):
        super().__init__()
        self.utae = encoder
        self.sta = STA(d_model)
        self.sem_head = SemanticDecoder(d_model, num_classes)
        self.bnd_head = BoundaryDecoder(d_model)
        self.refine = GatedRefinement(d_model, num_classes)

    def forward(self, x, pos):
        out, _ = self.utae(x, pos)                     # (B, H, W, T, D)
        feat = self.sta(out)                           # (B, D, H, W)
        sem = self.sem_head(feat)                      # (B, K, H, W)
        bnd = self.bnd_head(feat)                      # (B, 1, H, W)
        refined = self.refine(feat, sem, bnd)          # (B, K, H, W)
        return {"sem_logits": sem,
                "bnd_logits": bnd,
                "refined_logits": refined}


# ═══════════════════════════════════════════════════════════════════════════
# 6. Boundary ground truth (recovered verbatim from det pipeline)
# ═══════════════════════════════════════════════════════════════════════════
def compute_boundary_target(mask):
    """
    Morphological gradient: a pixel is a boundary iff at least one
    8-neighbour has a different class.
      mask: (B, H, W) long  ->  boundary: (B, H, W) float in {0, 1}
    (finetune.py squeezes bnd_logits to (B,H,W), so GT matches that shape.)
    """
    m = mask.unsqueeze(1).float()
    mx = F.max_pool2d(m, 3, 1, 1)
    mn = -F.max_pool2d(-m, 3, 1, 1)
    return (mx != mn).float().squeeze(1)