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