TSRDA / main_method /code /model.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
18 kB
"""
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)