File size: 3,785 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
66
67
68
69
70
71
72
73
74
75
76
77
78
"""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