File size: 6,707 Bytes
387a20d
 
 
 
 
 
 
 
 
 
7a2d30b
 
387a20d
 
 
 
 
 
 
 
 
 
 
 
7a2d30b
 
387a20d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7a2d30b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
387a20d
 
7a2d30b
 
 
387a20d
 
7a2d30b
 
 
387a20d
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
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
import torch
from torch import nn


class SpectralGPT(nn.Module):
    """Compact SpectralGPT masked autoencoder for 12-band spectral images."""

    def __init__(self, image_size=24, in_channels=12, patch_size=8,
                 spectral_patch_size=3, embed_dim=48, encoder_depth=2,
                 encoder_heads=4, decoder_dim=32, decoder_depth=1,
                  decoder_heads=4, mask_ratio=0.9, spectral_angle_weight=0.1,
                  spectral_gradient_weight=0.1):
        super().__init__()
        if image_size % patch_size or in_channels % spectral_patch_size:
            raise ValueError("Image and spectral dimensions must be divisible by token sizes")
        self.image_size = image_size
        self.in_channels = in_channels
        self.patch_size = patch_size
        self.spectral_patch_size = spectral_patch_size
        self.spatial_tokens = (image_size // patch_size) ** 2
        self.spectral_tokens = in_channels // spectral_patch_size
        self.num_tokens = self.spatial_tokens * self.spectral_tokens
        self.token_pixels = patch_size * patch_size * spectral_patch_size
        self.mask_ratio = mask_ratio
        self.spectral_angle_weight = spectral_angle_weight
        self.spectral_gradient_weight = spectral_gradient_weight

        self.patch_embed = nn.Conv3d(
            1, embed_dim,
            kernel_size=(spectral_patch_size, patch_size, patch_size),
            stride=(spectral_patch_size, patch_size, patch_size),
        )
        self.spatial_pos = nn.Parameter(torch.zeros(1, self.spatial_tokens, embed_dim))
        self.spectral_pos = nn.Parameter(torch.zeros(1, self.spectral_tokens, embed_dim))
        encoder_layer = nn.TransformerEncoderLayer(
            embed_dim, encoder_heads, embed_dim * 4, batch_first=True, norm_first=True
        )
        self.encoder = nn.TransformerEncoder(encoder_layer, encoder_depth)
        self.encoder_norm = nn.LayerNorm(embed_dim)
        self.decoder_embed = nn.Linear(embed_dim, decoder_dim)
        self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim))
        self.decoder_pos = nn.Linear(embed_dim, decoder_dim, bias=False)
        decoder_layer = nn.TransformerEncoderLayer(
            decoder_dim, decoder_heads, decoder_dim * 4, batch_first=True, norm_first=True
        )
        self.decoder = nn.TransformerEncoder(decoder_layer, decoder_depth)
        self.decoder_norm = nn.LayerNorm(decoder_dim)
        self.decoder_pred = nn.Linear(decoder_dim, self.token_pixels)
        nn.init.normal_(self.spatial_pos, std=0.02)
        nn.init.normal_(self.spectral_pos, std=0.02)
        nn.init.normal_(self.mask_token, std=0.02)

    def _positions(self):
        return (self.spatial_pos[:, None] + self.spectral_pos[:, :, None]).reshape(
            1, self.num_tokens, -1
        )

    def patchify(self, images):
        p, k = self.patch_size, self.spectral_patch_size
        n, c, h, w = images.shape
        if (c, h, w) != (self.in_channels, self.image_size, self.image_size):
            raise ValueError(f"Expected [N,{self.in_channels},{self.image_size},{self.image_size}]")
        x = images.reshape(n, c // k, k, h // p, p, w // p, p)
        x = x.permute(0, 1, 3, 5, 2, 4, 6)
        return x.reshape(n, self.num_tokens, self.token_pixels)

    def unpatchify(self, tokens):
        p, k = self.patch_size, self.spectral_patch_size
        n = tokens.shape[0]
        s = self.image_size // p
        x = tokens.reshape(n, self.spectral_tokens, s, s, k, p, p)
        x = x.permute(0, 1, 4, 2, 5, 3, 6)
        return x.reshape(n, self.in_channels, self.image_size, self.image_size)

    @staticmethod
    def random_masking(tokens, mask_ratio):
        n, length, dim = tokens.shape
        keep = max(1, int(length * (1.0 - mask_ratio)))
        order = torch.argsort(torch.rand(n, length, device=tokens.device), dim=1)
        restore = torch.argsort(order, dim=1)
        keep_ids = order[:, :keep]
        visible = torch.gather(tokens, 1, keep_ids.unsqueeze(-1).expand(-1, -1, dim))
        mask = torch.ones(n, length, device=tokens.device)
        mask[:, :keep] = 0
        mask = torch.gather(mask, 1, restore)
        return visible, mask, restore

    def forward(self, images, mask_ratio=None):
        ratio = self.mask_ratio if mask_ratio is None else mask_ratio
        embedded = self.patch_embed(images.unsqueeze(1)).flatten(2).transpose(1, 2)
        positions = self._positions()
        visible, mask, restore = self.random_masking(embedded + positions, ratio)
        latent = self.encoder_norm(self.encoder(visible))
        decoded_visible = self.decoder_embed(latent)
        missing = self.num_tokens - decoded_visible.shape[1]
        full = torch.cat([decoded_visible, self.mask_token.expand(images.shape[0], missing, -1)], 1)
        full = torch.gather(full, 1, restore.unsqueeze(-1).expand(-1, -1, full.shape[-1]))
        prediction = self.decoder_pred(self.decoder_norm(self.decoder(full + self.decoder_pos(positions))))
        target = self.patchify(images)
        token_error = (prediction - target).pow(2).mean(-1)
        masked_mse = (token_error * mask).sum() / mask.sum().clamp_min(1)
        mask_image = self.unpatchify(mask.unsqueeze(-1).expand(-1, -1, self.token_pixels))
        predicted_image = self.unpatchify(prediction)
        completed = images * (1.0 - mask_image) + predicted_image * mask_image
        spectral_mask = mask_image.any(dim=1)
        completed_norm = completed.norm(dim=1)
        target_norm = images.norm(dim=1)
        valid_sam = spectral_mask & (completed_norm > 1e-6) & (target_norm > 1e-6)
        cosine = (completed * images).sum(dim=1) / (completed_norm * target_norm).clamp_min(1e-6)
        angles = torch.acos(cosine.clamp(-1.0, 1.0))
        spectral_angle = (angles * valid_sam).sum() / valid_sam.sum().clamp_min(1)
        completed_gradient = completed[:, 1:] - completed[:, :-1]
        target_gradient = images[:, 1:] - images[:, :-1]
        gradient_mask = torch.maximum(mask_image[:, 1:], mask_image[:, :-1]).bool()
        spectral_gradient = ((completed_gradient - target_gradient).abs() * gradient_mask).sum() / gradient_mask.sum().clamp_min(1)
        loss = (masked_mse + self.spectral_angle_weight * spectral_angle
                + self.spectral_gradient_weight * spectral_gradient)
        return {
            "loss": loss,
            "masked_mse": masked_mse,
            "spectral_angle": spectral_angle,
            "spectral_gradient": spectral_gradient,
            "prediction": prediction,
            "mask": mask,
            "mask_image": mask_image,
            "prediction_image": predicted_image,
            "reconstruction": completed,
        }