SatMAE / model /satmae.py
zhangrenchao's picture
Upload SatMAE model package
355f250 verified
Raw
History Blame Contribute Delete
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
)