"""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()