Download model/prithvi_eo.py from OneScience-Group/PrithviEO: direct link, hf CLI and curl.
- Browser
- Download file 8.5 kB
-
https://huggingface.co/OneScience-Group/PrithviEO/resolve/main/model/prithvi_eo.py
- Command line
-
hf download hf://OneScience-Group/PrithviEO/model/prithvi_eo.py
-
curl -L -o prithvi_eo.py https://huggingface.co/OneScience-Group/PrithviEO/resolve/main/model/prithvi_eo.py
8.5 kB
| """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, | |
| } | |