Download main_method/code/model.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 18 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/model.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/model.py
-
curl -L -o model.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/model.py
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) | |