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