Download model/spectralgpt.py from OneScience-Group/SpectralGPT: direct link, hf CLI and curl.
- Browser
- Download file 6.71 kB
-
https://huggingface.co/OneScience-Group/SpectralGPT/resolve/main/model/spectralgpt.py
- Command line
-
hf download hf://OneScience-Group/SpectralGPT/model/spectralgpt.py
-
curl -L -o spectralgpt.py https://huggingface.co/OneScience-Group/SpectralGPT/resolve/main/model/spectralgpt.py
6.71 kB
| 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) | |
| 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, | |
| } | |