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