File size: 3,255 Bytes
eea5f0e | 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 | """Hybrid change detector: SiamCD (Satlas Swin-v2-B) + organizer-baseline UNet-R18 (6ch early fusion) as an
expert branch. The final class heads see both experts (features + probabilities) and learn whom to trust.
Inputs/outputs as in model_v2.SiamCD (pre/post in [0, 1]).
"""
import segmentation_models_pytorch as smp
import torch
import torch.nn as nn
import torch.nn.functional as F
from model_v2 import SiamCD, cbr
MEAN = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
STD = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)
class HybridCD(SiamCD):
def __init__(self, satlas_ckpt=None):
super().__init__(satlas_ckpt)
self.unet = smp.Unet(encoder_name="resnet18", encoder_weights=None, in_channels=6, classes=3)
uc = 16 # smp Unet decoder output channels
self.head_b = nn.Sequential(cbr(32 + 2 + uc + 2, 48), cbr(48, 32), nn.Conv2d(32, 1, 1))
self.head_t = nn.Sequential(cbr(32 + 2 + uc + 2, 48), cbr(48, 32), nn.Conv2d(32, 1, 1))
self.pres = nn.Sequential(nn.Linear(2 * (64 + 192) + 4, 128), nn.ReLU(inplace=True), nn.Linear(128, 2))
self.register_buffer("mean", MEAN, persistent=False)
self.register_buffer("std", STD, persistent=False)
def unet_branch(self, pre, post):
x = torch.cat([(pre - self.mean) / self.std, (post - self.mean) / self.std], 1)
feats = self.unet.encoder(x)
try:
d = self.unet.decoder(feats)
except TypeError:
d = self.unet.decoder(*feats)
logits = self.unet.segmentation_head(d)
return d, logits
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]
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:]
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))
ud, ul = self.unet_branch(pre, post)
up = ul.float().softmax(1)[:, 1:3].to(y.dtype) # baseline expert: P(new_building), P(tree_removal)
sp, sq = sem_pre.softmax(1), sem_post.softmax(1)
cb = self.head_b(torch.cat([y, sq[:, 1:2], sp[:, 1:2], ud, up], 1))
ct = self.head_t(torch.cat([y, sp[:, 2:3], sq[:, 2:3], ud, up], 1))
g = torch.cat([x.mean((2, 3)), x.amax((2, 3)), z[3].mean((2, 3)), z[3].amax((2, 3)),
up.mean((2, 3)), up.amax((2, 3))], 1)
return torch.cat([cb, ct], 1), self.pres(g), sem_pre, sem_post, ul
|