Download model/satmaepp.py from OneScience-Group/SatMAE-PP: direct link, hf CLI and curl.
- Browser
- Download file 17.1 kB
-
https://huggingface.co/OneScience-Group/SatMAE-PP/resolve/main/model/satmaepp.py
- Command line
-
hf download hf://OneScience-Group/SatMAE-PP/model/satmaepp.py
-
curl -L -o satmaepp.py https://huggingface.co/OneScience-Group/SatMAE-PP/resolve/main/model/satmaepp.py
17.1 kB
| """SatMAE++ masked autoencoder with convolutional multiscale decoding. | |
| Adapted for a self-contained package from the Apache-2.0 SatMAE++ reference | |
| implementation at commit bf02548ab2bf5123761cf059491ac5a631fbb428. | |
| """ | |
| import math | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| def sincos_1d(values, dim): | |
| half = (dim + 1) // 2 | |
| omega = torch.exp( | |
| -math.log(10000.0) * torch.arange(half, dtype=torch.float32) | |
| / max(half - 1, 1) | |
| ) | |
| phase = values.float().unsqueeze(-1) * omega | |
| return torch.cat((phase.sin(), phase.cos()), dim=-1)[..., :dim] | |
| def sincos_2d(side, dim): | |
| y, x = torch.meshgrid( | |
| torch.arange(side, dtype=torch.float32), | |
| torch.arange(side, dtype=torch.float32), indexing="ij" | |
| ) | |
| split = dim // 2 | |
| return torch.cat((sincos_1d(y.flatten(), split), | |
| sincos_1d(x.flatten(), dim - split)), dim=-1) | |
| class TransformerBlock(nn.Module): | |
| def __init__(self, dim, heads, mlp_ratio=4.0): | |
| super().__init__() | |
| self.norm1 = nn.LayerNorm(dim, eps=1e-6) | |
| self.attn = nn.MultiheadAttention(dim, heads, batch_first=True) | |
| self.norm2 = nn.LayerNorm(dim, eps=1e-6) | |
| hidden = int(dim * mlp_ratio) | |
| self.mlp = nn.Sequential(nn.Linear(dim, hidden), nn.GELU(), nn.Linear(hidden, dim)) | |
| def forward(self, x): | |
| value = self.norm1(x) | |
| x = x + self.attn(value, value, value, need_weights=False)[0] | |
| return x + self.mlp(self.norm2(x)) | |
| class ChannelFirstNorm(nn.Module): | |
| def __init__(self, channels): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(channels)) | |
| self.bias = nn.Parameter(torch.zeros(channels)) | |
| def forward(self, x): | |
| mean = x.mean(1, keepdim=True) | |
| variance = (x - mean).square().mean(1, keepdim=True) | |
| x = (x - mean) / torch.sqrt(variance + 1e-6) | |
| return x * self.weight[:, None, None] + self.bias[:, None, None] | |
| class ResidualBlock(nn.Module): | |
| def __init__(self, channels): | |
| super().__init__() | |
| self.conv1 = nn.Conv2d(channels, channels, 3, padding=1) | |
| self.conv2 = nn.Conv2d(channels, channels, 3, padding=1) | |
| def forward(self, x): | |
| return x + 0.5 * self.conv2(F.relu(self.conv1(x))) | |
| class UpsampleBlock(nn.Module): | |
| def __init__(self, hidden_channels, output_channels): | |
| super().__init__() | |
| self.up = nn.ConvTranspose2d(hidden_channels, hidden_channels, 4, 2, 1) | |
| self.up_norm = ChannelFirstNorm(hidden_channels) | |
| self.residual = ResidualBlock(hidden_channels) | |
| self.residual_norm = ChannelFirstNorm(hidden_channels) | |
| self.output = nn.Conv2d(hidden_channels, output_channels, 3, padding=1) | |
| def forward(self, x): | |
| hidden = F.leaky_relu(self.up_norm(self.up(x))) | |
| hidden = self.residual_norm(self.residual(hidden)) | |
| return hidden, self.output(hidden) | |
| class SatMAEPP(nn.Module): | |
| def __init__(self, image_size=224, patch_size=16, in_channels=3, | |
| embed_dim=1024, encoder_depth=24, encoder_heads=16, | |
| decoder_dim=512, decoder_depth=8, decoder_heads=16, | |
| mask_ratio=0.75, scales=None, mode="rgb", spectral_groups=None, | |
| spatial_mask=False, norm_pix_loss=False, proj_ratio=4, | |
| channel_embed_dim=None, decoder_channel_embed_dim=None): | |
| super().__init__() | |
| if image_size <= 0 or patch_size <= 0 or image_size % patch_size: | |
| raise ValueError("image_size must be divisible by patch_size") | |
| if in_channels <= 0 or embed_dim <= 0 or decoder_dim <= 0: | |
| raise ValueError("channel and embedding dimensions must be positive") | |
| if encoder_depth <= 0 or decoder_depth <= 0 or encoder_heads <= 0 or decoder_heads <= 0: | |
| raise ValueError("transformer depths and head counts must be positive") | |
| if embed_dim % encoder_heads or decoder_dim % decoder_heads: | |
| raise ValueError("embedding dimensions must be divisible by head counts") | |
| if not 0 <= mask_ratio < 1: | |
| raise ValueError("mask_ratio must be in [0, 1)") | |
| if mode not in {"rgb", "multispectral"}: | |
| raise ValueError("mode must be rgb or multispectral") | |
| self.image_size = image_size | |
| self.patch_size = patch_size | |
| self.in_channels = in_channels | |
| self.mask_ratio = mask_ratio | |
| self.mode = mode | |
| self.spatial_mask = spatial_mask | |
| self.norm_pix_loss = norm_pix_loss | |
| self.grid = image_size // patch_size | |
| self.num_patches = self.grid ** 2 | |
| self.scales = tuple(scales or ([1, 2] if mode == "rgb" else [1, 2, 4])) | |
| expected_scales = (1, 2) if mode == "rgb" else (1, 2, 4) | |
| if self.scales != expected_scales: | |
| raise ValueError("supported scales are [1, 2] or [1, 2, 4]") | |
| if mode == "rgb": | |
| self.groups = (tuple(range(in_channels)),) | |
| channel_embed_dim = 0 | |
| decoder_channel_embed_dim = 0 | |
| else: | |
| groups = spectral_groups or [[0, 1, 2, 6], [3, 4, 5, 7], [8, 9]] | |
| if sorted(channel for group in groups for channel in group) != list(range(in_channels)): | |
| raise ValueError("spectral_groups must partition all channels") | |
| self.groups = tuple(tuple(group) for group in groups) | |
| channel_embed_dim = channel_embed_dim or min(256, embed_dim // 4) | |
| decoder_channel_embed_dim = decoder_channel_embed_dim or min(128, decoder_dim // 4) | |
| self.group_count = len(self.groups) | |
| self.channel_embed_dim = channel_embed_dim | |
| self.decoder_channel_embed_dim = decoder_channel_embed_dim | |
| self.patch_embeds = nn.ModuleList([ | |
| nn.Conv2d(len(group), embed_dim, patch_size, patch_size) for group in self.groups | |
| ]) | |
| spatial_dim = embed_dim - channel_embed_dim | |
| decoder_spatial_dim = decoder_dim - decoder_channel_embed_dim | |
| self.register_buffer("position", sincos_2d(self.grid, spatial_dim)) | |
| self.register_buffer("decoder_position", sincos_2d(self.grid, decoder_spatial_dim)) | |
| if self.group_count > 1: | |
| ids = torch.arange(self.group_count, dtype=torch.float32) | |
| self.register_buffer("group_position", sincos_1d(ids, channel_embed_dim)) | |
| self.register_buffer("decoder_group_position", sincos_1d(ids, decoder_channel_embed_dim)) | |
| self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) | |
| self.blocks = nn.ModuleList([ | |
| TransformerBlock(embed_dim, encoder_heads) for _ in range(encoder_depth) | |
| ]) | |
| self.norm = nn.LayerNorm(embed_dim, eps=1e-6) | |
| self.decoder_embed = nn.Linear(embed_dim, decoder_dim) | |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim)) | |
| self.decoder_blocks = nn.ModuleList([ | |
| TransformerBlock(decoder_dim, decoder_heads) for _ in range(decoder_depth) | |
| ]) | |
| self.decoder_norm = nn.LayerNorm(decoder_dim, eps=1e-6) | |
| self.decoder_heads = nn.ModuleList([ | |
| nn.Linear(decoder_dim, len(group) * patch_size ** 2) for group in self.groups | |
| ]) | |
| hidden_channels = in_channels * proj_ratio | |
| self.multiscale_projection = nn.Conv2d(in_channels, hidden_channels, 1) | |
| self.multiscale_norm = ChannelFirstNorm(hidden_channels) | |
| self.upsample_blocks = nn.ModuleList([ | |
| UpsampleBlock(hidden_channels, in_channels) for _ in self.scales[1:] | |
| ]) | |
| self.initialize_weights() | |
| def initialize_weights(self): | |
| nn.init.normal_(self.cls_token, std=0.02) | |
| nn.init.normal_(self.mask_token, std=0.02) | |
| for module in self.modules(): | |
| if isinstance(module, (nn.Linear, nn.Conv2d, nn.ConvTranspose2d)): | |
| nn.init.xavier_uniform_(module.weight.flatten(1)) | |
| if module.bias is not None: | |
| nn.init.zeros_(module.bias) | |
| elif isinstance(module, nn.LayerNorm): | |
| nn.init.ones_(module.weight) | |
| nn.init.zeros_(module.bias) | |
| def patchify(self, images): | |
| batch, channels, height, width = images.shape | |
| p = self.patch_size | |
| if channels != self.in_channels or height != self.image_size or width != self.image_size: | |
| raise ValueError("images must have shape [B, in_channels, image_size, image_size]") | |
| x = images.reshape(batch, channels, height // p, p, width // p, p) | |
| return x.permute(0, 2, 4, 1, 3, 5).reshape(batch, -1, channels * p ** 2) | |
| def unpatchify(self, patches): | |
| batch = patches.shape[0] | |
| p = self.patch_size | |
| x = patches.reshape(batch, self.grid, self.grid, self.in_channels, p, p) | |
| return x.permute(0, 3, 1, 4, 2, 5).reshape( | |
| batch, self.in_channels, self.image_size, self.image_size | |
| ) | |
| def _positions(self, decoder=False): | |
| spatial = self.decoder_position if decoder else self.position | |
| if self.group_count == 1: | |
| return spatial.unsqueeze(0) | |
| group = self.decoder_group_position if decoder else self.group_position | |
| value = torch.cat(( | |
| spatial.unsqueeze(0).expand(self.group_count, -1, -1), | |
| group.unsqueeze(1).expand(-1, self.num_patches, -1), | |
| ), dim=-1) | |
| return value.reshape(1, self.group_count * self.num_patches, -1) | |
| def _mask(self, tokens, ratio): | |
| batch, length, dim = tokens.shape | |
| shared = self.group_count > 1 and self.spatial_mask | |
| if shared: | |
| kept_spatial = int(self.num_patches * (1 - ratio)) | |
| order = torch.rand(batch, self.num_patches, device=tokens.device).argsort(1) | |
| kept = [order[:, :kept_spatial] + index * self.num_patches | |
| for index in range(self.group_count)] | |
| removed = [order[:, kept_spatial:] + index * self.num_patches | |
| for index in range(self.group_count)] | |
| shuffle = torch.cat(kept + removed, dim=1) | |
| keep = kept_spatial * self.group_count | |
| else: | |
| keep = int(length * (1 - ratio)) | |
| shuffle = torch.rand(batch, length, device=tokens.device).argsort(1) | |
| restore = shuffle.argsort(1) | |
| visible = torch.gather(tokens, 1, shuffle[:, :keep, None].expand(-1, -1, dim)) | |
| mask = torch.ones(batch, length, device=tokens.device) | |
| mask[:, :keep] = 0 | |
| return visible, torch.gather(mask, 1, restore), restore | |
| def forward_encoder(self, images, ratio): | |
| pieces = [embed(images[:, group]).flatten(2).transpose(1, 2) | |
| for embed, group in zip(self.patch_embeds, self.groups)] | |
| tokens = torch.cat(pieces, dim=1) + self._positions(False) | |
| tokens, mask, restore = self._mask(tokens, ratio) | |
| tokens = torch.cat((self.cls_token.expand(images.shape[0], -1, -1), tokens), 1) | |
| for block in self.blocks: | |
| tokens = block(tokens) | |
| return self.norm(tokens), mask, restore | |
| def forward_decoder(self, latent, restore): | |
| tokens = self.decoder_embed(latent) | |
| missing = restore.shape[1] + 1 - tokens.shape[1] | |
| restored = torch.cat((tokens[:, 1:], self.mask_token.expand(tokens.shape[0], missing, -1)), 1) | |
| restored = torch.gather(restored, 1, restore[:, :, None].expand(-1, -1, tokens.shape[-1])) | |
| tokens = torch.cat((tokens[:, :1], restored + self._positions(True)), 1) | |
| for block in self.decoder_blocks: | |
| tokens = block(tokens) | |
| decoded = self.decoder_norm(tokens)[:, 1:].reshape( | |
| tokens.shape[0], self.group_count, self.num_patches, -1 | |
| ) | |
| group_predictions = [head(decoded[:, index]) for index, head in enumerate(self.decoder_heads)] | |
| patch_channels = [] | |
| for prediction, group in zip(group_predictions, self.groups): | |
| patch_channels.append(prediction.reshape( | |
| prediction.shape[0], self.num_patches, len(group), self.patch_size ** 2 | |
| )) | |
| patches = torch.empty( | |
| tokens.shape[0], self.num_patches, self.in_channels, self.patch_size ** 2, | |
| device=tokens.device, dtype=group_predictions[0].dtype | |
| ) | |
| for values, group in zip(patch_channels, self.groups): | |
| patches[:, :, list(group)] = values | |
| return patches.flatten(2), group_predictions | |
| def forward_multiscale(self, reconstruction): | |
| hidden = self.multiscale_norm(F.gelu(self.multiscale_projection(reconstruction))) | |
| predictions = {"1": reconstruction} | |
| for scale, block in zip(self.scales[1:], self.upsample_blocks): | |
| hidden, predictions[str(scale)] = block(hidden) | |
| return predictions | |
| def forward(self, images, high_resolution_targets=None, mask_ratio=None): | |
| ratio = self.mask_ratio if mask_ratio is None else mask_ratio | |
| if images.ndim != 4 or images.shape[1:] != (self.in_channels, self.image_size, self.image_size): | |
| raise ValueError("images must have shape [B, in_channels, image_size, image_size]") | |
| if not 0 <= ratio < 1: | |
| raise ValueError("mask_ratio must be in [0, 1)") | |
| high_resolution_targets = high_resolution_targets or {} | |
| latent, mask, restore = self.forward_encoder(images, ratio) | |
| patch_prediction, group_predictions = self.forward_decoder(latent, restore) | |
| target_patches = self.patchify(images) | |
| loss_target = target_patches | |
| if self.norm_pix_loss: | |
| mean = loss_target.mean(-1, keepdim=True) | |
| variance = loss_target.var(-1, keepdim=True, unbiased=False) | |
| loss_target = (loss_target - mean) / torch.sqrt(variance + 1e-6) | |
| grouped_mask = mask.reshape(mask.shape[0], self.group_count, self.num_patches) | |
| prediction_channels = patch_prediction.reshape( | |
| patch_prediction.shape[0], self.num_patches, | |
| self.in_channels, self.patch_size ** 2 | |
| ) | |
| target_channels = loss_target.reshape( | |
| loss_target.shape[0], self.num_patches, | |
| self.in_channels, self.patch_size ** 2 | |
| ) | |
| base_total = mask.new_zeros(()) | |
| removed = mask.new_zeros(()) | |
| for index, group in enumerate(self.groups): | |
| group_error = ( | |
| prediction_channels[:, :, list(group)] | |
| - target_channels[:, :, list(group)] | |
| ).square().mean(dim=(-1, -2)) | |
| base_total = base_total + (group_error * grouped_mask[:, index]).sum() | |
| removed = removed + grouped_mask[:, index].sum() | |
| base_mse = base_total / removed.clamp_min(1) | |
| base_l1_total = mask.new_zeros(()) | |
| for index, group in enumerate(self.groups): | |
| group_l1 = (prediction_channels[:, :, list(group)] - | |
| target_channels[:, :, list(group)]).abs().mean(dim=(-1, -2)) | |
| base_l1_total = base_l1_total + (group_l1 * grouped_mask[:, index]).sum() | |
| base_l1 = base_l1_total / removed.clamp_min(1) | |
| base_loss = base_mse + base_l1 | |
| reconstruction = self.unpatchify(patch_prediction) | |
| predictions = self.forward_multiscale(reconstruction) | |
| targets = {"1": images} | |
| multiscale_losses = {} | |
| for scale in self.scales[1:]: | |
| key = str(scale) | |
| target = None if high_resolution_targets is None else high_resolution_targets.get(key) | |
| if target is None: | |
| raise ValueError(f"native target for scale x{scale} is required") | |
| expected_shape = (images.shape[0], self.in_channels, | |
| self.image_size * scale, self.image_size * scale) | |
| if tuple(target.shape) != expected_shape: | |
| raise ValueError(f"target x{scale} must have shape {expected_shape}") | |
| targets[key] = target | |
| multiscale_losses[key] = F.mse_loss(predictions[key], target) + F.l1_loss( | |
| predictions[key], target | |
| ) | |
| multiscale_loss = (sum(multiscale_losses.values()) / len(multiscale_losses) | |
| if multiscale_losses else base_loss.new_zeros(())) | |
| return { | |
| "loss": base_loss + multiscale_loss, | |
| "reconstruction_loss": base_loss, | |
| "multiscale_loss": multiscale_loss, | |
| "reconstruction": reconstruction, | |
| "patch_prediction": patch_prediction, | |
| "target_patches": loss_target, | |
| "predictions": predictions, | |
| "targets": targets, | |
| "mask": mask.bool(), | |
| "features": latent, | |
| "ids_restore": restore, | |
| "group_predictions": group_predictions, | |
| "scale_losses": multiscale_losses, | |
| } | |
| def satmaepp_vit_base(**kwargs): | |
| return SatMAEPP(embed_dim=768, encoder_depth=12, encoder_heads=12, | |
| decoder_dim=512, decoder_depth=8, decoder_heads=16, **kwargs) | |
| def satmaepp_vit_large(**kwargs): | |
| return SatMAEPP(embed_dim=1024, encoder_depth=24, encoder_heads=16, | |
| decoder_dim=512, decoder_depth=8, decoder_heads=16, **kwargs) | |