fnruha0921's picture
add code
eea5f0e verified
Raw History Blame Contribute Delete
3.79 kB
"""Siamese change detector: shared Satlas Aerial Swin-v2-B encoder, directional fusion,
full-resolution stem, asymmetric class heads, per-image semantic aux head, presence head.
Inputs: pre, post as float tensors in [0, 1], shape (B, 3, H, W), H and W divisible by 32.
Outputs: change logits (B, 2, H, W) [new_building, tree_removal], presence logits (B, 2),
semantic logits for pre and post (B, 3, H, W) [bg, building, tree].
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision
ENC_C = [128, 256, 512, 1024]
FUSE_C = [64, 96, 128, 192]
def cbr(i, o, k=3, s=1):
return nn.Sequential(nn.Conv2d(i, o, k, s, k // 2, bias=False), nn.BatchNorm2d(o), nn.ReLU(inplace=True))
class SiamCD(nn.Module):
def __init__(self, satlas_ckpt=None):
super().__init__()
sw = torchvision.models.swin_v2_b()
if satlas_ckpt:
sd = torch.load(satlas_ckpt, map_location="cpu", weights_only=False)
sd = {k.replace("backbone.backbone.", ""): v for k, v in sd.items() if k.startswith("backbone.backbone.")}
sw.load_state_dict(sd)
self.enc = sw.features
self.fuse = nn.ModuleList([nn.Sequential(cbr(4 * c, d, 1), cbr(d, d)) for c, d in zip(ENC_C, FUSE_C)])
self.dec = nn.ModuleList([cbr(FUSE_C[i + 1] + FUSE_C[i], FUSE_C[i]) for i in range(3)])
self.sem_lat = nn.ModuleList([nn.Conv2d(c, 64, 1) for c in ENC_C])
self.sem_out = nn.Sequential(cbr(64, 64), nn.Conv2d(64, 3, 1))
self.stem1 = nn.Sequential(cbr(6, 32), cbr(32, 32))
self.stem2 = cbr(32, 48, s=2)
self.up2 = cbr(FUSE_C[0] + 48, 48)
self.up1 = cbr(48 + 32, 32)
self.head_b = nn.Sequential(cbr(32 + 2, 32), nn.Conv2d(32, 1, 1))
self.head_t = nn.Sequential(cbr(32 + 2, 32), nn.Conv2d(32, 1, 1))
self.pres = nn.Sequential(nn.Linear(2 * (FUSE_C[0] + FUSE_C[3]), 128), nn.ReLU(inplace=True), nn.Linear(128, 2))
def encode(self, x):
feats = []
for i, blk in enumerate(self.enc):
x = blk(x)
if i in (1, 3, 5, 7):
feats.append(x.permute(0, 3, 1, 2).contiguous())
return feats
def forward(self, pre, post):
B, _, H, W = pre.shape
f = self.encode(torch.cat([pre, post], 0))
fp, fq = [t[:B] for t in f], [t[B:] for t in f]
# per-image semantic head (both images at once)
s = 0
for lat, t in zip(self.sem_lat, f):
s = s + F.interpolate(lat(t), size=f[0].shape[-2:], mode="bilinear", align_corners=False)
sem = F.interpolate(self.sem_out(s), size=(H, W), mode="bilinear", align_corners=False)
sem_pre, sem_post = sem[:B], sem[B:]
# directional fusion + top-down decoder
z = [fu(torch.cat([a, b, b - a, (b - a).abs()], 1)) for fu, a, b in zip(self.fuse, fp, fq)]
x = z[3]
for i in (2, 1, 0):
x = F.interpolate(x, size=z[i].shape[-2:], mode="bilinear", align_corners=False)
x = self.dec[i](torch.cat([x, z[i]], 1))
s1 = self.stem1(torch.cat([pre, post], 1))
s2 = self.stem2(s1)
y = self.up2(torch.cat([F.interpolate(x, size=s2.shape[-2:], mode="bilinear", align_corners=False), s2], 1))
y = self.up1(torch.cat([F.interpolate(y, size=(H, W), mode="bilinear", align_corners=False), s1], 1))
sp, sq = sem_pre.softmax(1), sem_post.softmax(1)
cb = self.head_b(torch.cat([y, sq[:, 1:2], sp[:, 1:2]], 1)) # new building: post-dominant
ct = self.head_t(torch.cat([y, sp[:, 2:3], sq[:, 2:3]], 1)) # tree removal: pre-dominant
g = torch.cat([x.mean((2, 3)), x.amax((2, 3)), z[3].mean((2, 3)), z[3].amax((2, 3))], 1)
return torch.cat([cb, ct], 1), self.pres(g), sem_pre, sem_post