"""Pure-PyTorch engineering reproduction of the Prithvi-EO-2.0 TL MAE.""" import math import torch from torch import nn def sincos_1d(positions, dim): if dim % 2: raise ValueError("sine/cosine dimensions must be even") frequencies = torch.exp( torch.arange(dim // 2, device=positions.device, dtype=positions.dtype) * (-math.log(10000.0) / max(dim // 2, 1)) ) angles = positions.unsqueeze(-1) * frequencies return torch.cat((angles.sin(), angles.cos()), dim=-1) def sincos_3d(frames, height, width, dim, device, dtype): if dim % 16: raise ValueError("3D position dimension must be divisible by 16") width_dim, height_dim, time_dim = 6 * dim // 16, 6 * dim // 16, 4 * dim // 16 time, row, column = torch.meshgrid( torch.arange(frames, device=device, dtype=dtype), torch.arange(height, device=device, dtype=dtype), torch.arange(width, device=device, dtype=dtype), indexing="ij", ) return torch.cat(( sincos_1d(column.reshape(-1), width_dim), sincos_1d(row.reshape(-1), height_dim), sincos_1d(time.reshape(-1), time_dim), ), dim=-1) def patchify(values, patch_size): batch, channels, frames, height, width = values.shape pt, ph, pw = patch_size if frames % pt or height % ph or width % pw: raise ValueError("input dimensions must be divisible by patch_size") return values.reshape( batch, channels, frames // pt, pt, height // ph, ph, width // pw, pw ).permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(batch, -1, pt * ph * pw * channels) def unpatchify(patches, channels, output_size, patch_size): batch = patches.shape[0] frames, height, width = output_size pt, ph, pw = patch_size return patches.reshape( batch, frames // pt, height // ph, width // pw, pt, ph, pw, channels ).permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(batch, channels, frames, height, width) class Transformer(nn.Module): def __init__(self, dim, depth, heads, mlp_ratio): super().__init__() layer = nn.TransformerEncoderLayer( dim, heads, int(dim * mlp_ratio), activation="gelu", batch_first=True, norm_first=True ) self.blocks = nn.TransformerEncoder(layer, depth) self.norm = nn.LayerNorm(dim) def forward(self, values): return self.norm(self.blocks(values)) class CoordinateEncoder(nn.Module): def __init__(self, dim, scale=0.1): super().__init__() if dim % 4: raise ValueError("coordinate embedding dimension must be divisible by four") self.dim = dim self.scale = nn.Parameter(torch.tensor(float(scale))) def forward(self, coordinates): return self.scale * torch.cat(( sincos_1d(coordinates[..., 0], self.dim // 2), sincos_1d(coordinates[..., 1], self.dim // 2), ), dim=-1) class PrithviEO2(nn.Module): def __init__(self, config): super().__init__() self.config = dict(config) self.input_size = tuple(int(value) for value in config["input_size"]) self.patch_size = tuple(int(value) for value in config["patch_size"]) self.channels = int(config["channels"]) self.mask_ratio = float(config["mask_ratio"]) self.metadata_dropout = float(config["metadata_dropout"]) self.norm_pix_loss = bool(config.get("norm_pix_loss", False)) enc_dim, dec_dim = int(config["encoder_dim"]), int(config["decoder_dim"]) self.patch_embed = nn.Conv3d( self.channels, enc_dim, kernel_size=self.patch_size, stride=self.patch_size ) self.cls_token = nn.Parameter(torch.randn(1, 1, enc_dim) * 0.02) self.encoder = Transformer(enc_dim, int(config["encoder_depth"]), int(config["encoder_heads"]), float(config["mlp_ratio"])) self.encoder_to_decoder = nn.Linear(enc_dim, dec_dim) self.mask_token = nn.Parameter(torch.randn(1, 1, dec_dim) * 0.02) self.decoder = Transformer(dec_dim, int(config["decoder_depth"]), int(config["decoder_heads"]), float(config["mlp_ratio"])) patch_volume = math.prod(self.patch_size) * self.channels self.decoder_prediction = nn.Linear(dec_dim, patch_volume) self.time_encoder = CoordinateEncoder(enc_dim) self.location_encoder = CoordinateEncoder(enc_dim) self.decoder_time_encoder = CoordinateEncoder(dec_dim) self.decoder_location_encoder = CoordinateEncoder(dec_dim) def _grid(self, pixels): return tuple(size // patch for size, patch in zip(pixels.shape[-3:], self.patch_size)) def _metadata(self, temporal, location, grid, encoder=True): frames, height, width = grid time_encoder = self.time_encoder if encoder else self.decoder_time_encoder location_encoder = self.location_encoder if encoder else self.decoder_location_encoder temporal_embedding = time_encoder(temporal) temporal_embedding = temporal_embedding[:, :, None, :].expand(-1, -1, height * width, -1).reshape( len(temporal), frames * height * width, -1 ) location_embedding = location_encoder(location)[:, None, :].expand(-1, frames * height * width, -1) if self.training and self.metadata_dropout: time_keep = (torch.rand(len(temporal), 1, 1, device=temporal.device) >= self.metadata_dropout).to(temporal.dtype) location_keep = (torch.rand(len(location), 1, 1, device=location.device) >= self.metadata_dropout).to(location.dtype) temporal_embedding = temporal_embedding * time_keep location_embedding = location_embedding * location_keep return temporal_embedding + location_embedding def _encoded_tokens(self, pixels, temporal, location): grid = self._grid(pixels) tokens = self.patch_embed(pixels).flatten(2).transpose(1, 2) position = sincos_3d(*grid, tokens.shape[-1], tokens.device, tokens.dtype) tokens = tokens + position[None] + self._metadata(temporal, location, grid, encoder=True) return tokens, grid def encode(self, pixels, temporal, location): tokens, _ = self._encoded_tokens(pixels, temporal, location) cls = self.cls_token.expand(len(pixels), -1, -1) encoded = self.encoder(torch.cat((cls, tokens), dim=1)) return encoded[:, 0], encoded[:, 1:] def forward(self, pixels, temporal, location, mask_ratio=None): ratio = self.mask_ratio if mask_ratio is None else float(mask_ratio) tokens, grid = self._encoded_tokens(pixels, temporal, location) batch, length, dim = tokens.shape keep = max(1, int(length * (1.0 - ratio))) ordering = torch.rand(batch, length, device=pixels.device).argsort(dim=1) visible_indices, masked_indices = ordering[:, :keep], ordering[:, keep:] visible = tokens.gather(1, visible_indices[:, :, None].expand(-1, -1, dim)) encoded = self.encoder(torch.cat((self.cls_token.expand(batch, -1, -1), visible), dim=1)) embedding = encoded[:, 0] visible_decoder = self.encoder_to_decoder(encoded[:, 1:]) decoder_tokens = self.mask_token.expand(batch, length, -1).clone() decoder_tokens.scatter_(1, visible_indices[:, :, None].expand(-1, -1, visible_decoder.shape[-1]), visible_decoder) position = sincos_3d(*grid, decoder_tokens.shape[-1], decoder_tokens.device, decoder_tokens.dtype) decoder_tokens = decoder_tokens + position[None] + self._metadata(temporal, location, grid, encoder=False) predictions = self.decoder_prediction(self.decoder(decoder_tokens)) targets = patchify(pixels, self.patch_size) if self.norm_pix_loss: mean, variance = targets.mean(dim=-1, keepdim=True), targets.var(dim=-1, keepdim=True) targets = (targets - mean) / (variance + 1e-6).sqrt() mask = torch.zeros(batch, length, device=pixels.device) mask.scatter_(1, masked_indices, 1.0) patch_mse = (predictions - targets).pow(2).mean(dim=-1) loss = (patch_mse * mask).sum() / mask.sum().clamp_min(1) reconstruction = unpatchify(predictions, self.channels, pixels.shape[-3:], self.patch_size) return { "loss": loss, "embedding": embedding, "patch_embeddings": encoded[:, 1:], "reconstruction": reconstruction, "mask": mask, }