AlphaEarthFoundations / model /alphaearthfoundations.py
zhangrenchao's picture
Add engineering reproduction package
3549cf5 verified
Raw
History Blame Contribute Delete
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()