Download model/alphaearthfoundations.py from OneScience-Group/AlphaEarthFoundations: direct link, hf CLI and curl.
- Browser
- Download file 12.4 kB
-
https://huggingface.co/OneScience-Group/AlphaEarthFoundations/resolve/main/model/alphaearthfoundations.py
- Command line
-
hf download hf://OneScience-Group/AlphaEarthFoundations/model/alphaearthfoundations.py
-
curl -L -o alphaearthfoundations.py https://huggingface.co/OneScience-Group/AlphaEarthFoundations/resolve/main/model/alphaearthfoundations.py
12.4 kB
| """Engineering reproduction of AlphaEarth Foundations from the paper specification.""" | |
| import math | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| def sinusoidal_timecode(timestamps, dim, origin=None, scale=365.25 * 24 * 3600 * 1000): | |
| if origin is None: | |
| origin = timestamps.amin(dim=1, keepdim=True) | |
| values = (timestamps.double() - origin.double()) / scale | |
| values = values.float() | |
| frequencies = torch.exp( | |
| torch.arange(0, dim, 2, device=timestamps.device) * (-math.log(10000.0) / dim) | |
| ) | |
| angles = values.unsqueeze(-1) * frequencies | |
| code = torch.zeros(*timestamps.shape, dim, device=timestamps.device) | |
| code[..., 0::2] = angles.sin() | |
| code[..., 1::2] = angles.cos() | |
| return code | |
| class STPBlock(nn.Module): | |
| """Parallel precision, time and space operators with learned pyramid exchange.""" | |
| def __init__(self, precision_dim, time_dim, space_dim, num_heads): | |
| super().__init__() | |
| self.precision = nn.Sequential( | |
| nn.GroupNorm(1, precision_dim), | |
| nn.Conv2d(precision_dim, precision_dim, 3, padding=1), | |
| nn.GELU(), | |
| nn.Conv2d(precision_dim, precision_dim, 3, padding=1), | |
| ) | |
| self.time_norm = nn.LayerNorm(time_dim) | |
| self.time_attention = nn.MultiheadAttention(time_dim, num_heads, batch_first=True) | |
| self.space_norm = nn.LayerNorm(space_dim) | |
| self.space_attention = nn.MultiheadAttention(space_dim, num_heads, batch_first=True) | |
| self.to_precision = nn.ModuleList([nn.Conv2d(time_dim, precision_dim, 1), nn.Conv2d(space_dim, precision_dim, 1)]) | |
| self.to_time = nn.Conv2d(precision_dim, time_dim, 1) | |
| self.to_space = nn.Conv2d(precision_dim, space_dim, 1) | |
| def forward(self, precision, time, space, frame_available): | |
| batch, frames = precision.shape[:2] | |
| p_size, t_size, s_size = precision.shape[-2:], time.shape[-2:], space.shape[-2:] | |
| p = precision.flatten(0, 1) | |
| p = p + self.precision(p) | |
| sequence = time.permute(0, 3, 4, 1, 2).reshape(-1, frames, time.shape[2]) | |
| normalized = self.time_norm(sequence) | |
| time_mask = (~frame_available.bool())[:, None, None, :].expand(batch, *t_size, frames).reshape(-1, frames) | |
| sequence = sequence + self.time_attention( | |
| normalized, normalized, normalized, key_padding_mask=time_mask, need_weights=False | |
| )[0] | |
| time = sequence.reshape(batch, *t_size, frames, -1).permute(0, 3, 4, 1, 2) | |
| available = frame_available[:, :, None, None, None].to(space.dtype) | |
| spatial = (space * available).sum(dim=1) / available.sum(dim=1).clamp_min(1) | |
| spatial = spatial.flatten(2).transpose(1, 2) | |
| normalized = self.space_norm(spatial) | |
| spatial = spatial + self.space_attention(normalized, normalized, normalized, need_weights=False)[0] | |
| spatial = spatial.transpose(1, 2).reshape(batch, -1, *s_size) | |
| space = space + spatial[:, None] | |
| t_flat, s_flat = time.flatten(0, 1), space.flatten(0, 1) | |
| precision = p + self.to_precision[0](F.interpolate(t_flat, p_size, mode="bilinear", align_corners=False)) | |
| precision = precision + self.to_precision[1](F.interpolate(s_flat, p_size, mode="bilinear", align_corners=False)) | |
| time = time + self.to_time(F.interpolate(p, t_size, mode="bilinear", align_corners=False)).unflatten(0, (batch, frames)) | |
| space = space + self.to_space(F.interpolate(p, s_size, mode="bilinear", align_corners=False)).unflatten(0, (batch, frames)) | |
| return precision.unflatten(0, (batch, frames)), time, space | |
| class ConditionalDecoder(nn.Module): | |
| def __init__(self, embedding_dim, condition_dim, hidden_dim, output_dim): | |
| super().__init__() | |
| self.condition = nn.Linear(condition_dim, hidden_dim) | |
| self.network = nn.Sequential( | |
| nn.Conv2d(embedding_dim + hidden_dim, hidden_dim, 1), | |
| nn.GELU(), | |
| nn.Conv2d(hidden_dim, hidden_dim, 1), | |
| nn.GELU(), | |
| nn.Conv2d(hidden_dim, output_dim, 1), | |
| ) | |
| def forward(self, embedding, condition): | |
| context = self.condition(condition)[:, :, None, None].expand(-1, -1, *embedding.shape[-2:]) | |
| return self.network(torch.cat([embedding, context], dim=1)) | |
| class AlphaEarthFoundations(nn.Module): | |
| def __init__(self, input_sources, target_sources, config): | |
| super().__init__() | |
| p_dim, t_dim, s_dim = config["precision_dim"], config["time_dim"], config["space_dim"] | |
| self.input_names = list(input_sources) | |
| self.target_sources = target_sources | |
| self.embedding_dim = config["embedding_dim"] | |
| self.vmf_kappa = float(config["vmf_kappa"]) | |
| self.projectors = nn.ModuleDict({ | |
| name: nn.Sequential(nn.Conv2d(spec["channels"], p_dim, 3, stride=2, padding=1), nn.GELU()) | |
| for name, spec in input_sources.items() | |
| }) | |
| self.time_projector = nn.Conv2d(p_dim, t_dim, 3, stride=4, padding=1) | |
| self.space_projector = nn.Conv2d(p_dim, s_dim, 3, stride=8, padding=1) | |
| self.time_context = nn.Linear(t_dim, t_dim) | |
| self.blocks = nn.ModuleList([ | |
| STPBlock(p_dim, t_dim, s_dim, config["num_heads"]) for _ in range(config["num_blocks"]) | |
| ]) | |
| self.summary_query = nn.Linear(t_dim * 2, p_dim) | |
| self.embedding_head = nn.Conv2d(p_dim, self.embedding_dim, 1) | |
| self.embedding_upsample = nn.ConvTranspose2d(p_dim, p_dim, 4, stride=2, padding=1) | |
| condition_dim = t_dim + config["max_geometry_dim"] | |
| self.decoders = nn.ModuleDict({ | |
| name: ConditionalDecoder(self.embedding_dim, condition_dim, config["decoder_hidden_dim"], spec["channels"]) | |
| for name, spec in target_sources.items() | |
| }) | |
| def _summarize(self, precision, availability, period, origin): | |
| period_codes = sinusoidal_timecode(period, self.time_context.in_features, origin) | |
| query = self.summary_query(period_codes.flatten(1)) | |
| scores = (precision * query[:, None, :, None, None]).sum(dim=2).mean(dim=(-1, -2)) | |
| scores = scores.masked_fill(~availability.bool(), torch.finfo(scores.dtype).min) | |
| summary = (precision * scores.softmax(dim=1)[:, :, None, None, None]).sum(dim=1) | |
| return F.normalize(self.embedding_head(self.embedding_upsample(summary)), dim=1) | |
| def forward(self, sources, timestamps, valid_period, frame_available, target_times=None, | |
| target_geometry=None, target_periods=None): | |
| precision_parts, code_parts = [], [] | |
| origin = torch.cat(list(timestamps.values()), dim=1).amin(dim=1, keepdim=True) | |
| for name in self.input_names: | |
| values = sources[name] | |
| batch, frames = values.shape[:2] | |
| projected = self.projectors[name](values.flatten(0, 1)).unflatten(0, (batch, frames)) | |
| precision_parts.append(projected) | |
| code_parts.append(sinusoidal_timecode(timestamps[name], self.time_context.in_features, origin)) | |
| availability = torch.cat([frame_available[name] for name in self.input_names], dim=1) | |
| precision = torch.cat(precision_parts, dim=1) | |
| codes = torch.cat(code_parts, dim=1) | |
| time = self.time_projector(precision.flatten(0, 1)).unflatten(0, precision.shape[:2]) | |
| time = time + self.time_context(codes)[:, :, :, None, None] | |
| space = self.space_projector(precision.flatten(0, 1)).unflatten(0, precision.shape[:2]) | |
| for block in self.blocks: | |
| precision, time, space = block(precision, time, space, availability) | |
| embedding = self._summarize(precision, availability, valid_period, origin) | |
| outputs = {"embedding": embedding} | |
| if target_times is not None: | |
| outputs["reconstructions"] = {} | |
| for name in self.target_sources: | |
| source_embedding = self._summarize(precision, availability, target_periods[name], origin) | |
| if self.training: | |
| source_embedding = F.normalize( | |
| source_embedding + torch.randn_like(source_embedding) / math.sqrt(self.vmf_kappa), dim=1 | |
| ) | |
| relative_time = ( | |
| (target_times[name] - target_periods[name][:, 0]).float() | |
| / (target_periods[name][:, 1] - target_periods[name][:, 0]).float().clamp_min(1) | |
| ) | |
| time_code = sinusoidal_timecode( | |
| relative_time[:, None], self.time_context.in_features, | |
| torch.zeros_like(relative_time[:, None]), scale=1.0 | |
| )[:, 0] | |
| geometry = target_geometry[name] | |
| outputs["reconstructions"][name] = self.decoders[name](source_embedding, torch.cat([time_code, geometry], dim=1)) | |
| return outputs | |
| def _pool_continuous(values, grid_m): | |
| if grid_m == 10: | |
| return values | |
| size = max(1, round(values.shape[-1] * 10 / grid_m)) | |
| return F.adaptive_avg_pool2d(values, (size, size)) | |
| def _shift_invariant_l1(prediction, target, mask, radius): | |
| losses = [] | |
| for dy in range(-radius, radius + 1): | |
| for dx in range(-radius, radius + 1): | |
| shifted = torch.roll(prediction, (dy, dx), dims=(-2, -1)) | |
| valid = mask.clone() | |
| if dy > 0: valid[..., :dy, :] = 0 | |
| if dy < 0: valid[..., dy:, :] = 0 | |
| if dx > 0: valid[..., :, :dx] = 0 | |
| if dx < 0: valid[..., :, dx:] = 0 | |
| losses.append((torch.abs(shifted - target) * valid).sum() / valid.sum().clamp_min(1)) | |
| return torch.stack(losses).amin() | |
| def compute_losses(teacher, student, targets, masks, text_target, target_sources, weights): | |
| reconstruction = teacher["embedding"].new_zeros(()) | |
| components = {} | |
| for name, spec in target_sources.items(): | |
| prediction, target, mask = teacher["reconstructions"][name], targets[name], masks[name] | |
| grid_m = int(spec["loss_grid_m"]) | |
| if spec["type"] == "categorical": | |
| size = max(1, round(prediction.shape[-1] * 10 / grid_m)) | |
| prediction = F.adaptive_avg_pool2d(prediction, (size, size)) | |
| one_hot = F.one_hot(target.long(), num_classes=prediction.shape[1]).permute(0, 3, 1, 2).float() | |
| target = F.adaptive_avg_pool2d(one_hot, (size, size)).argmax(dim=1) | |
| mask = F.adaptive_avg_pool2d(mask, (size, size)) | |
| value = F.cross_entropy(prediction, target, reduction="none") | |
| value = (value * mask[:, 0]).sum() / mask[:, 0].sum().clamp_min(1) | |
| else: | |
| if spec.get("shift_pixels", 0): | |
| value = _shift_invariant_l1(prediction, target, mask, int(spec["shift_pixels"])) | |
| else: | |
| prediction, target, mask = (_pool_continuous(item, grid_m) for item in (prediction, target, mask)) | |
| value = (torch.abs(prediction - target) * mask).sum() / mask.sum().clamp_min(1) | |
| components[f"reconstruction_{name}"] = value | |
| reconstruction = reconstruction + float(spec["weight"]) * value | |
| flat = teacher["embedding"].permute(0, 2, 3, 1).reshape(-1, teacher["embedding"].shape[1]) | |
| rotated = torch.roll(flat, max(1, flat.shape[0] // 2), dims=0) | |
| uniformity = (flat * rotated).sum(dim=1).abs().mean() | |
| consistency = 1.0 - (teacher["embedding"] * student["embedding"]).sum(dim=1).mean() | |
| pooled = F.normalize(teacher["embedding"].mean(dim=(2, 3)), dim=1) | |
| normalized_text = F.normalize(text_target, dim=1) | |
| logits = pooled @ normalized_text.transpose(0, 1) | |
| labels = torch.arange(len(logits), device=logits.device) | |
| text_alignment = 0.5 * (F.cross_entropy(logits, labels) + F.cross_entropy(logits.transpose(0, 1), labels)) | |
| total = (weights["reconstruction"] * reconstruction + weights["uniformity"] * uniformity | |
| + weights["consistency"] * consistency + weights["text"] * text_alignment) | |
| components.update(reconstruction=reconstruction, uniformity=uniformity, | |
| consistency=consistency, text_alignment=text_alignment, total=total) | |
| return total, components | |
| def quantize_embeddings(embedding, power=2, scale=127.5): | |
| transformed = embedding.abs().pow(1.0 / power) * embedding.sign() | |
| return torch.round(transformed * scale).clamp(-127, 127).to(torch.int8) | |
| def dequantize_embeddings(quantized, power=2, scale=127.5): | |
| values = quantized.float() / scale | |
| return values.abs().pow(power) * values.sign() | |