""" 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)