| """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, |
| } |
|
|