Download code/model_v3.py from fnruha0921/knps-change-detection-tmp: direct link, hf CLI and curl.
- Browser
- Download file 3.26 kB
-
https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/model_v3.py
- Command line
-
hf download hf://fnruha0921/knps-change-detection-tmp/code/model_v3.py
-
curl -L -o model_v3.py https://huggingface.co/fnruha0921/knps-change-detection-tmp/resolve/main/code/model_v3.py
3.26 kB
| """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 | |