Download model/satmae.py from OneScience-Group/SatMAE: direct link, hf CLI and curl.
- Browser
- Download file 20.9 kB
-
https://huggingface.co/OneScience-Group/SatMAE/resolve/main/model/satmae.py
- Command line
-
hf download hf://OneScience-Group/SatMAE/model/satmae.py
-
curl -L -o satmae.py https://huggingface.co/OneScience-Group/SatMAE/resolve/main/model/satmae.py
20.9 kB
| """Paper-aligned SatMAE model components. | |
| This is an original implementation of the architecture described in SatMAE. | |
| The upstream repository was used only as a behavioral reference; no upstream | |
| source text is incorporated here. | |
| """ | |
| import math | |
| from functools import partial | |
| import torch | |
| from torch import nn | |
| def _sincos_1d(values, dim): | |
| """Return a fixed sine-cosine embedding for arbitrary scalar positions.""" | |
| if dim <= 0: | |
| return values.new_zeros((*values.shape, 0)) | |
| pairs = (dim + 1) // 2 | |
| omega = torch.arange(pairs, device=values.device, dtype=torch.float32) | |
| omega = torch.exp(-math.log(10000.0) * omega / max(pairs - 1, 1)) | |
| phase = values.to(torch.float32).unsqueeze(-1) * omega | |
| return torch.cat((phase.sin(), phase.cos()), dim=-1)[..., :dim] | |
| def _sincos_2d(grid_size, dim): | |
| """Return a fixed row-major 2D sine-cosine position embedding.""" | |
| rows, cols = torch.meshgrid( | |
| torch.arange(grid_size, dtype=torch.float32), | |
| torch.arange(grid_size, dtype=torch.float32), | |
| indexing="ij", | |
| ) | |
| row_dim = dim // 2 | |
| return torch.cat( | |
| (_sincos_1d(rows.reshape(-1), row_dim), | |
| _sincos_1d(cols.reshape(-1), dim - row_dim)), | |
| dim=-1, | |
| ) | |
| def _timestamp_embedding(timestamps, dim): | |
| """Encode either scalar times or fMoW ``[year, month, hour]`` tuples.""" | |
| if timestamps.ndim == 2: | |
| return _sincos_1d(timestamps, dim) | |
| if timestamps.ndim != 3 or timestamps.shape[-1] != 3: | |
| raise ValueError("timestamps must have shape [B, T] or [B, T, 3]") | |
| field_dims = [dim // 3] * 3 | |
| for index in range(dim % 3): | |
| field_dims[index] += 1 | |
| return torch.cat( | |
| [_sincos_1d(timestamps[..., index], field_dim) | |
| for index, field_dim in enumerate(field_dims)], | |
| dim=-1, | |
| ) | |
| class PatchEmbed(nn.Module): | |
| def __init__(self, image_size, patch_size, in_channels, embed_dim): | |
| super().__init__() | |
| self.image_size = image_size | |
| self.patch_size = patch_size | |
| self.num_patches = (image_size // patch_size) ** 2 | |
| self.proj = nn.Conv2d( | |
| in_channels, embed_dim, kernel_size=patch_size, stride=patch_size | |
| ) | |
| def forward(self, images): | |
| if images.shape[-2:] != (self.image_size, self.image_size): | |
| raise ValueError( | |
| f"expected {self.image_size}x{self.image_size} images, " | |
| f"got {tuple(images.shape[-2:])}" | |
| ) | |
| return self.proj(images).flatten(2).transpose(1, 2) | |
| class TransformerBlock(nn.Module): | |
| def __init__(self, dim, num_heads, mlp_ratio=4.0, norm_layer=nn.LayerNorm): | |
| super().__init__() | |
| self.norm1 = norm_layer(dim) | |
| self.attention = nn.MultiheadAttention( | |
| dim, num_heads, dropout=0.0, bias=True, batch_first=True | |
| ) | |
| self.norm2 = norm_layer(dim) | |
| hidden_dim = int(dim * mlp_ratio) | |
| self.mlp = nn.Sequential( | |
| nn.Linear(dim, hidden_dim), nn.GELU(), nn.Linear(hidden_dim, dim) | |
| ) | |
| def forward(self, tokens): | |
| normalized = self.norm1(tokens) | |
| tokens = tokens + self.attention( | |
| normalized, normalized, normalized, need_weights=False | |
| )[0] | |
| return tokens + self.mlp(self.norm2(tokens)) | |
| class SatMAE(nn.Module): | |
| """Masked autoencoder for temporal or grouped multispectral imagery. | |
| Temporal inputs use shape ``[B, T, C, H, W]`` and optional timestamps | |
| ``[B, T]``. Multispectral inputs use shape ``[B, C, H, W]``. | |
| """ | |
| def __init__( | |
| self, | |
| image_size=224, | |
| patch_size=16, | |
| in_channels=3, | |
| frames=3, | |
| embed_dim=1024, | |
| encoder_depth=24, | |
| encoder_heads=16, | |
| decoder_dim=512, | |
| decoder_depth=8, | |
| decoder_heads=16, | |
| mlp_ratio=4.0, | |
| mode="temporal", | |
| spectral_groups=None, | |
| mask_ratio=0.75, | |
| norm_pix_loss=False, | |
| same_mask=False, | |
| spatial_mask=False, | |
| temporal_embed_dim=None, | |
| decoder_temporal_embed_dim=None, | |
| channel_embed_dim=None, | |
| decoder_channel_embed_dim=None, | |
| norm_layer=None, | |
| ): | |
| super().__init__() | |
| if image_size % patch_size: | |
| raise ValueError("image_size must be divisible by patch_size") | |
| if not 0.0 <= mask_ratio < 1.0: | |
| raise ValueError("mask_ratio must be in [0, 1)") | |
| if mode not in {"temporal", "multispectral"}: | |
| raise ValueError("mode must be temporal or multispectral") | |
| if embed_dim % encoder_heads or decoder_dim % decoder_heads: | |
| raise ValueError("embedding dimensions must be divisible by head counts") | |
| norm_layer = norm_layer or partial(nn.LayerNorm, eps=1e-6) | |
| self.image_size = image_size | |
| self.patch_size = patch_size | |
| self.in_channels = in_channels | |
| self.frames = frames | |
| self.embed_dim = embed_dim | |
| self.decoder_dim = decoder_dim | |
| self.mode = mode | |
| self.mask_ratio = mask_ratio | |
| self.norm_pix_loss = norm_pix_loss | |
| self.same_mask = same_mask | |
| self.spatial_mask = spatial_mask | |
| self.grid_size = image_size // patch_size | |
| self.num_patches = self.grid_size ** 2 | |
| if mode == "temporal": | |
| self.spectral_groups = None | |
| self.patch_embed = PatchEmbed( | |
| image_size, patch_size, in_channels, embed_dim | |
| ) | |
| self.token_groups = frames | |
| semantic_dim = temporal_embed_dim | |
| if semantic_dim is None: | |
| semantic_dim = min(128, max(2, embed_dim // 4)) | |
| decoder_semantic_dim = decoder_temporal_embed_dim | |
| if decoder_semantic_dim is None: | |
| decoder_semantic_dim = min(64, max(2, decoder_dim // 4)) | |
| prediction_dims = [patch_size ** 2 * in_channels] | |
| else: | |
| groups = spectral_groups or [list(range(in_channels))] | |
| flattened = [channel for group in groups for channel in group] | |
| if sorted(flattened) != list(range(in_channels)): | |
| raise ValueError("spectral_groups must partition all input channels") | |
| self.spectral_groups = tuple(tuple(group) for group in groups) | |
| self.patch_embed = nn.ModuleList( | |
| PatchEmbed(image_size, patch_size, len(group), embed_dim) | |
| for group in self.spectral_groups | |
| ) | |
| self.token_groups = len(self.spectral_groups) | |
| semantic_dim = channel_embed_dim | |
| if semantic_dim is None: | |
| semantic_dim = min(256, max(2, embed_dim // 4)) | |
| decoder_semantic_dim = decoder_channel_embed_dim | |
| if decoder_semantic_dim is None: | |
| decoder_semantic_dim = min(128, max(2, decoder_dim // 4)) | |
| prediction_dims = [patch_size ** 2 * len(g) for g in self.spectral_groups] | |
| if not 0 < semantic_dim < embed_dim: | |
| raise ValueError("encoder semantic embedding dimension is invalid") | |
| if not 0 < decoder_semantic_dim < decoder_dim: | |
| raise ValueError("decoder semantic embedding dimension is invalid") | |
| self.semantic_dim = semantic_dim | |
| self.decoder_semantic_dim = decoder_semantic_dim | |
| self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) | |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, decoder_dim)) | |
| self.register_buffer( | |
| "spatial_pos_embed", | |
| _sincos_2d(self.grid_size, embed_dim - semantic_dim), | |
| persistent=True, | |
| ) | |
| self.register_buffer( | |
| "decoder_spatial_pos_embed", | |
| _sincos_2d(self.grid_size, decoder_dim - decoder_semantic_dim), | |
| persistent=True, | |
| ) | |
| if mode == "multispectral": | |
| group_ids = torch.arange(self.token_groups, dtype=torch.float32) | |
| self.register_buffer( | |
| "group_embed", _sincos_1d(group_ids, semantic_dim), persistent=True | |
| ) | |
| self.register_buffer( | |
| "decoder_group_embed", | |
| _sincos_1d(group_ids, decoder_semantic_dim), | |
| persistent=True, | |
| ) | |
| self.blocks = nn.ModuleList( | |
| TransformerBlock(embed_dim, encoder_heads, mlp_ratio, norm_layer) | |
| for _ in range(encoder_depth) | |
| ) | |
| self.norm = norm_layer(embed_dim) | |
| self.decoder_embed = nn.Linear(embed_dim, decoder_dim) | |
| self.decoder_blocks = nn.ModuleList( | |
| TransformerBlock(decoder_dim, decoder_heads, mlp_ratio, norm_layer) | |
| for _ in range(decoder_depth) | |
| ) | |
| self.decoder_norm = norm_layer(decoder_dim) | |
| self.decoder_pred = nn.ModuleList( | |
| nn.Linear(decoder_dim, output_dim) for output_dim in prediction_dims | |
| ) | |
| self.initialize_weights() | |
| def initialize_weights(self): | |
| patch_embeds = ( | |
| [self.patch_embed] | |
| if isinstance(self.patch_embed, PatchEmbed) | |
| else self.patch_embed | |
| ) | |
| for patch_embed in patch_embeds: | |
| nn.init.xavier_uniform_(patch_embed.proj.weight.flatten(1)) | |
| if patch_embed.proj.bias is not None: | |
| nn.init.zeros_(patch_embed.proj.bias) | |
| 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.init.xavier_uniform_(module.weight) | |
| 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): | |
| if images.ndim != 4: | |
| raise ValueError("patchify expects [B, C, H, W]") | |
| batch, channels, height, width = images.shape | |
| patch = self.patch_size | |
| if height != width or height != self.image_size: | |
| raise ValueError(f"expected square images of size {self.image_size}") | |
| patches = images.reshape( | |
| batch, channels, height // patch, patch, width // patch, patch | |
| ) | |
| patches = patches.permute(0, 2, 4, 1, 3, 5) | |
| return patches.reshape(batch, self.num_patches, channels * patch ** 2) | |
| def unpatchify(self, patches, channels=None): | |
| channels = channels or self.in_channels | |
| batch = patches.shape[0] | |
| patch = self.patch_size | |
| expected = channels * patch ** 2 | |
| if patches.shape[1:] != (self.num_patches, expected): | |
| raise ValueError("patch tensor has incompatible shape") | |
| images = patches.reshape( | |
| batch, self.grid_size, self.grid_size, channels, patch, patch | |
| ) | |
| images = images.permute(0, 3, 1, 4, 2, 5) | |
| return images.reshape(batch, channels, self.image_size, self.image_size) | |
| def _random_masking(self, tokens, mask_ratio, share_spatial_mask): | |
| batch, length, dim = tokens.shape | |
| if share_spatial_mask: | |
| units = self.num_patches | |
| len_keep_units = int(units * (1.0 - mask_ratio)) | |
| noise = torch.rand(batch, units, device=tokens.device) | |
| spatial_order = noise.argsort(dim=1) | |
| kept = [spatial_order[:, :len_keep_units] + g * units | |
| for g in range(self.token_groups)] | |
| removed = [spatial_order[:, len_keep_units:] + g * units | |
| for g in range(self.token_groups)] | |
| ids_shuffle = torch.cat(kept + removed, dim=1) | |
| len_keep = len_keep_units * self.token_groups | |
| else: | |
| len_keep = int(length * (1.0 - mask_ratio)) | |
| ids_shuffle = torch.rand(batch, length, device=tokens.device).argsort(dim=1) | |
| ids_restore = ids_shuffle.argsort(dim=1) | |
| ids_keep = ids_shuffle[:, :len_keep] | |
| visible = torch.gather(tokens, 1, ids_keep.unsqueeze(-1).expand(-1, -1, dim)) | |
| mask = torch.ones(batch, length, device=tokens.device) | |
| mask[:, :len_keep] = 0 | |
| mask = torch.gather(mask, 1, ids_restore) | |
| return visible, mask, ids_restore | |
| def _temporal_tokens(self, images, timestamps): | |
| if images.ndim != 5: | |
| raise ValueError("temporal mode expects images shaped [B, T, C, H, W]") | |
| batch, frames, channels, _, _ = images.shape | |
| if frames != self.frames or channels != self.in_channels: | |
| raise ValueError( | |
| f"expected T={self.frames}, C={self.in_channels}; got T={frames}, C={channels}" | |
| ) | |
| if timestamps is None: | |
| timestamps = torch.arange(frames, device=images.device).expand(batch, -1) | |
| if timestamps.shape[:2] != (batch, frames): | |
| raise ValueError( | |
| f"timestamps must start with shape {(batch, frames)}, " | |
| f"got {tuple(timestamps.shape)}" | |
| ) | |
| spatial = self.spatial_pos_embed.to(dtype=images.dtype) | |
| time = _timestamp_embedding(timestamps, self.semantic_dim).to(dtype=images.dtype) | |
| position = torch.cat( | |
| (spatial.view(1, 1, self.num_patches, -1).expand(batch, frames, -1, -1), | |
| time.unsqueeze(2).expand(-1, -1, self.num_patches, -1)), | |
| dim=-1, | |
| ).reshape(batch, frames * self.num_patches, self.embed_dim) | |
| tokens = torch.stack( | |
| [self.patch_embed(images[:, frame]) for frame in range(frames)], dim=1 | |
| ).reshape(batch, frames * self.num_patches, self.embed_dim) | |
| return tokens + position, timestamps | |
| def _multispectral_tokens(self, images): | |
| if images.ndim != 4 or images.shape[1] != self.in_channels: | |
| raise ValueError( | |
| f"multispectral mode expects images shaped [B, {self.in_channels}, H, W]" | |
| ) | |
| spatial = self.spatial_pos_embed.to(dtype=images.dtype) | |
| group = self.group_embed.to(dtype=images.dtype) | |
| positions = torch.cat( | |
| (spatial.view(1, self.num_patches, -1).expand(self.token_groups, -1, -1), | |
| group.view(self.token_groups, 1, -1).expand(-1, self.num_patches, -1)), | |
| dim=-1, | |
| ).reshape(1, self.token_groups * self.num_patches, self.embed_dim) | |
| tokens = torch.cat( | |
| [embed(images[:, channels]) | |
| for embed, channels in zip(self.patch_embed, self.spectral_groups)], | |
| dim=1, | |
| ) | |
| return tokens + positions | |
| def forward_encoder(self, images, timestamps=None, mask_ratio=None): | |
| ratio = self.mask_ratio if mask_ratio is None else mask_ratio | |
| if not 0.0 <= ratio < 1.0: | |
| raise ValueError("mask_ratio must be in [0, 1)") | |
| if self.mode == "temporal": | |
| tokens, timestamps = self._temporal_tokens(images, timestamps) | |
| shared = self.same_mask | |
| else: | |
| tokens = self._multispectral_tokens(images) | |
| shared = self.spatial_mask | |
| tokens, mask, ids_restore = self._random_masking(tokens, ratio, shared) | |
| cls = self.cls_token.expand(tokens.shape[0], -1, -1) | |
| tokens = torch.cat((cls, tokens), dim=1) | |
| for block in self.blocks: | |
| tokens = block(tokens) | |
| return self.norm(tokens), mask, ids_restore, timestamps | |
| def _decoder_positions(self, batch, timestamps, dtype, device): | |
| spatial = self.decoder_spatial_pos_embed.to(device=device, dtype=dtype) | |
| if self.mode == "temporal": | |
| semantic = _timestamp_embedding(timestamps, self.decoder_semantic_dim).to(dtype=dtype) | |
| else: | |
| semantic = self.decoder_group_embed.to(device=device, dtype=dtype) | |
| semantic = semantic.unsqueeze(0).expand(batch, -1, -1) | |
| position = torch.cat( | |
| (spatial.view(1, 1, self.num_patches, -1).expand(batch, self.token_groups, -1, -1), | |
| semantic.unsqueeze(2).expand(-1, -1, self.num_patches, -1)), | |
| dim=-1, | |
| ) | |
| return position.reshape(batch, self.token_groups * self.num_patches, self.decoder_dim) | |
| def forward_decoder(self, latent, ids_restore, timestamps=None): | |
| tokens = self.decoder_embed(latent) | |
| mask_tokens = self.mask_token.expand( | |
| tokens.shape[0], ids_restore.shape[1] + 1 - tokens.shape[1], -1 | |
| ) | |
| restored = torch.cat((tokens[:, 1:], mask_tokens), dim=1) | |
| restored = torch.gather( | |
| restored, 1, ids_restore.unsqueeze(-1).expand(-1, -1, self.decoder_dim) | |
| ) | |
| positions = self._decoder_positions( | |
| tokens.shape[0], timestamps, tokens.dtype, tokens.device | |
| ) | |
| tokens = torch.cat((tokens[:, :1], restored + positions), dim=1) | |
| for block in self.decoder_blocks: | |
| tokens = block(tokens) | |
| decoded = self.decoder_norm(tokens)[:, 1:] | |
| if self.mode == "temporal": | |
| return [self.decoder_pred[0](decoded)] | |
| decoded = decoded.reshape( | |
| decoded.shape[0], self.token_groups, self.num_patches, self.decoder_dim | |
| ) | |
| return [head(decoded[:, index]) for index, head in enumerate(self.decoder_pred)] | |
| def _targets(self, images): | |
| if self.mode == "temporal": | |
| return [torch.cat( | |
| [self.patchify(images[:, frame]) for frame in range(self.frames)], dim=1 | |
| )] | |
| return [self.patchify(images[:, group]) for group in self.spectral_groups] | |
| def forward_loss(self, targets, predictions, mask): | |
| losses = [] | |
| if self.mode == "temporal": | |
| pairs = [(targets[0], predictions[0], mask)] | |
| else: | |
| group_mask = mask.reshape(mask.shape[0], self.token_groups, self.num_patches) | |
| pairs = [ | |
| (target, prediction, group_mask[:, index]) | |
| for index, (target, prediction) in enumerate(zip(targets, predictions)) | |
| ] | |
| removed = mask.new_zeros(()) | |
| total = mask.new_zeros(()) | |
| for target, prediction, patch_mask in pairs: | |
| patch_loss = (prediction - target).square().mean(dim=-1) | |
| total = total + (patch_loss * patch_mask).sum() | |
| removed = removed + patch_mask.sum() | |
| losses.append(patch_loss) | |
| return total / removed.clamp_min(1), losses | |
| def _normalize_targets(self, targets): | |
| if not self.norm_pix_loss: | |
| return targets | |
| normalized = [] | |
| for target in targets: | |
| mean = target.mean(dim=-1, keepdim=True) | |
| variance = target.var(dim=-1, keepdim=True, unbiased=False) | |
| normalized.append((target - mean) / torch.sqrt(variance + 1e-6)) | |
| return normalized | |
| def _padded_outputs(self, tensors): | |
| if self.mode == "temporal": | |
| return tensors[0] | |
| width = max(tensor.shape[-1] for tensor in tensors) | |
| padded = [] | |
| for tensor in tensors: | |
| if tensor.shape[-1] < width: | |
| tensor = torch.nn.functional.pad(tensor, (0, width - tensor.shape[-1])) | |
| padded.append(tensor) | |
| return torch.cat(padded, dim=1) | |
| def forward(self, images, timestamps=None, mask_ratio=None): | |
| latent, mask, ids_restore, timestamps = self.forward_encoder( | |
| images, timestamps, mask_ratio | |
| ) | |
| predictions = self.forward_decoder(latent, ids_restore, timestamps) | |
| targets = self._normalize_targets(self._targets(images)) | |
| loss, patch_losses = self.forward_loss(targets, predictions, mask) | |
| return { | |
| "loss": loss, | |
| "prediction": self._padded_outputs(predictions), | |
| "target": self._padded_outputs(targets), | |
| "mask": mask.bool(), | |
| "features": latent, | |
| "ids_restore": ids_restore, | |
| "group_predictions": predictions, | |
| "group_targets": targets, | |
| "patch_losses": patch_losses, | |
| } | |
| def satmae_vit_base_patch16(**kwargs): | |
| return SatMAE( | |
| patch_size=16, embed_dim=768, encoder_depth=12, encoder_heads=12, | |
| decoder_dim=512, decoder_depth=8, decoder_heads=16, | |
| temporal_embed_dim=128, decoder_temporal_embed_dim=64, | |
| channel_embed_dim=256, decoder_channel_embed_dim=128, **kwargs | |
| ) | |
| def satmae_vit_large_patch16(**kwargs): | |
| return SatMAE( | |
| patch_size=16, embed_dim=1024, encoder_depth=24, encoder_heads=16, | |
| decoder_dim=512, decoder_depth=8, decoder_heads=16, | |
| temporal_embed_dim=128, decoder_temporal_embed_dim=64, | |
| channel_embed_dim=256, decoder_channel_embed_dim=128, **kwargs | |
| ) | |
| def satmae_vit_huge_patch14(**kwargs): | |
| return SatMAE( | |
| patch_size=14, embed_dim=1280, encoder_depth=32, encoder_heads=16, | |
| decoder_dim=512, decoder_depth=8, decoder_heads=16, | |
| temporal_embed_dim=128, decoder_temporal_embed_dim=64, | |
| channel_embed_dim=256, decoder_channel_embed_dim=128, **kwargs | |
| ) | |