| 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, |
| } |
|
|