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