File size: 7,463 Bytes
d9ae24d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
"""Scale-aware masked autoencoder for paired remote-sensing resolutions."""
import math
import torch
from torch import nn
from torch.nn import functional as F


def sincos_2d(height, width, dim, device, dtype, scale=1.0):
    if dim % 4:
        raise ValueError("embedding dimension must be divisible by four")
    y, x = torch.meshgrid(torch.arange(height, device=device, dtype=dtype),
                          torch.arange(width, device=device, dtype=dtype), indexing="ij")
    quarter = dim // 4
    omega = torch.exp(-math.log(10000.0) * torch.arange(quarter, device=device, dtype=dtype) /
                      max(quarter - 1, 1))
    x, y = x.flatten()[:, None] * scale, y.flatten()[:, None] * scale
    return torch.cat((torch.sin(x * omega), torch.cos(x * omega),
                      torch.sin(y * omega), torch.cos(y * omega)), dim=1)


class ScaleMAE(nn.Module):
    def __init__(self, input_size=32, target_size=64, patch_size=4, in_channels=3,
                 embed_dim=64, encoder_depth=2, encoder_heads=4, decoder_dim=48,
                 decoder_depth=1, decoder_heads=4, mask_ratio=0.75, reference_gsd=1.0,
                 blur_kernel=5, band_config=None, **legacy):
        super().__init__()
        input_size = legacy.get("image_size", input_size)
        self.input_size, self.target_size, self.patch_size = input_size, target_size, patch_size
        self.in_channels, self.embed_dim, self.mask_ratio = in_channels, embed_dim, mask_ratio
        self.reference_gsd, self.blur_kernel = reference_gsd, blur_kernel
        if input_size % patch_size or target_size % patch_size:
            raise ValueError("input_size and target_size must be divisible by patch_size")
        self.input_grid = input_size // patch_size
        self.target_grid = target_size // patch_size
        self.num_patches = self.input_grid ** 2
        patch_dim = in_channels * patch_size ** 2
        self.patch_embed = nn.Linear(patch_dim, embed_dim)
        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
        self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
        enc = nn.TransformerEncoderLayer(embed_dim, encoder_heads, embed_dim * 4,
                                         dropout=0, activation="gelu", batch_first=True, norm_first=True)
        self.encoder, self.encoder_norm = nn.TransformerEncoder(enc, encoder_depth), nn.LayerNorm(embed_dim)
        self.decoder_input = nn.Linear(embed_dim, decoder_dim)
        dec = nn.TransformerEncoderLayer(decoder_dim, decoder_heads, decoder_dim * 4,
                                         dropout=0, activation="gelu", batch_first=True, norm_first=True)
        self.decoder, self.decoder_norm = nn.TransformerEncoder(dec, decoder_depth), nn.LayerNorm(decoder_dim)
        self.fpn = nn.Sequential(nn.Conv2d(decoder_dim, decoder_dim, 3, padding=1), nn.GELU(),
                                 nn.Conv2d(decoder_dim, decoder_dim, 3, padding=1), nn.GELU())
        self.low_head = nn.Conv2d(decoder_dim, in_channels, 1)
        self.high_head = nn.Conv2d(decoder_dim, in_channels, 1)
        self.band_config = band_config or {"low": {"kernel": blur_kernel}, "high": {"residual": True}}
        nn.init.normal_(self.cls_token, std=.02); nn.init.normal_(self.mask_token, std=.02)

    def patchify(self, images):
        b, c, h, w = images.shape; p = self.patch_size
        if (c, h, w) != (self.in_channels, self.input_size, self.input_size):
            raise ValueError(f"expected input {(self.in_channels, self.input_size, self.input_size)}, got {(c,h,w)}")
        return images.reshape(b, c, h//p, p, w//p, p).permute(0,2,4,1,3,5).reshape(b, -1, c*p*p)

    def bandpass_targets(self, target):
        k = int(self.band_config.get("low", {}).get("kernel", self.blur_kernel))
        if k < 3 or k % 2 == 0: raise ValueError("low-frequency kernel must be odd and at least three")
        low = F.avg_pool2d(target, k, 1, k//2, count_include_pad=False)
        return low, target - low

    def _positions(self, gsd, grid, dim):
        scales = gsd.to(dtype=self.cls_token.dtype).flatten() / self.reference_gsd
        base = sincos_2d(grid, grid, dim, gsd.device, self.cls_token.dtype)
        return base.unsqueeze(0) * scales[:, None, None]

    def forward(self, images, gsd, target=None, mask_ratio=None, return_features=True):
        if target is None: target = F.interpolate(images, (self.target_size, self.target_size), mode="bilinear", align_corners=False)
        if target.shape[-2:] != (self.target_size, self.target_size):
            raise ValueError("target resolution does not match target_size")
        if target.shape[:2] != (images.shape[0], self.in_channels):
            raise ValueError("target batch or channel dimensions do not match images")
        b = images.shape[0]; patches = self.patchify(images); n = patches.shape[1]
        ratio = self.mask_ratio if mask_ratio is None else float(mask_ratio)
        if not 0 <= ratio < 1: raise ValueError("mask_ratio must be in [0, 1)")
        noise = torch.rand(b, n, device=images.device)
        ids_shuffle = noise.argsort(dim=1); ids_restore = ids_shuffle.argsort(dim=1)
        keep = max(1, int(n * (1 - ratio))); ids_keep = ids_shuffle[:, :keep]
        visible = torch.gather(self.patch_embed(patches), 1, ids_keep[..., None].expand(-1, -1, self.embed_dim))
        visible = visible + torch.gather(self._positions(gsd, self.input_grid, self.embed_dim), 1,
                                         ids_keep[..., None].expand(-1, -1, self.embed_dim))
        encoded = self.encoder_norm(self.encoder(torch.cat((self.cls_token.expand(b,-1,-1), visible), 1)))
        cls, visible_encoded = encoded[:, :1], encoded[:, 1:]
        decoded = self.decoder_input(visible_encoded)
        full = self.mask_token.to(decoded.dtype).expand(b, n, -1).clone(); full.scatter_(1, ids_keep[..., None].expand(-1,-1,decoded.shape[-1]), decoded)
        full = full + self._positions(gsd, self.input_grid, full.shape[-1])
        full = self.decoder_norm(self.decoder(full))
        fmap = full.transpose(1, 2).reshape(b, -1, self.input_grid, self.input_grid)
        fmap = F.interpolate(fmap, (self.target_grid, self.target_grid), mode="bilinear", align_corners=False)
        fmap = self.fpn(fmap)
        low = F.interpolate(self.low_head(fmap), (self.target_size, self.target_size), mode="bilinear", align_corners=False)
        high = F.interpolate(self.high_head(fmap), (self.target_size, self.target_size), mode="bilinear", align_corners=False)
        low_target, high_target = self.bandpass_targets(target)
        mask = torch.ones(b, n, dtype=torch.bool, device=images.device); mask.scatter_(1, ids_keep, False)
        pixel_mask = F.interpolate(mask.view(b, 1, self.input_grid, self.input_grid).float(),
                                   (self.target_size, self.target_size), mode="nearest")
        masked = pixel_mask.sum().clamp_min(1.0) * self.in_channels
        low_loss = ((low-low_target).square() * pixel_mask).sum() / masked
        high_loss = ((high-high_target).square() * pixel_mask).sum() / masked
        return {"loss": low_loss + high_loss, "low_loss": low_loss, "high_loss": high_loss,
                "reconstruction": low + high, "low_reconstruction": low, "high_reconstruction": high,
                "low_target": low_target, "high_target": high_target, "mask": mask,
                "ids_restore": ids_restore, "features": cls[:, 0] if return_features else None}