ClayFoundation / model /clayfoundation.py
zhangrenchao's picture
Add engineering reproduction package
53becf5 verified
Raw
History Blame Contribute Delete
11.2 kB
"""Pure-PyTorch engineering reproduction of the Clay v1.5 model specification."""
import math
import torch
from torch import nn
from torch.nn import functional as F
def fourier_encode(values, dim, max_frequency=10000.0):
"""Encode scalar metadata while preserving exactly ``dim`` output features."""
if dim < 1:
return values.new_zeros(*values.shape, 0)
pairs = (dim + 1) // 2
frequencies = torch.exp(
torch.linspace(0, math.log(max_frequency), pairs, device=values.device, dtype=values.dtype)
)
angles = values.unsqueeze(-1) * frequencies
return torch.cat((angles.sin(), angles.cos()), dim=-1)[..., :dim]
def position_encoding_2d(height, width, dim, gsd, device, dtype):
if dim % 4:
raise ValueError("spatial position dimension must be divisible by four")
y, x = torch.meshgrid(
torch.arange(height, device=device, dtype=dtype),
torch.arange(width, device=device, dtype=dtype),
indexing="ij",
)
scale = torch.as_tensor(gsd, device=device, dtype=dtype) / 10.0
quarter = dim // 4
frequencies = torch.exp(
torch.arange(quarter, device=device, dtype=dtype) * (-math.log(10000.0) / max(quarter, 1))
)
x_angles = x.reshape(-1, 1) * scale * frequencies
y_angles = y.reshape(-1, 1) * scale * frequencies
return torch.cat((x_angles.sin(), x_angles.cos(), y_angles.sin(), y_angles.cos()), dim=-1)
def patchify(pixels, patch_size):
batch, channels, height, width = pixels.shape
if height % patch_size or width % patch_size:
raise ValueError("image dimensions must be divisible by patch_size")
return pixels.reshape(
batch, channels, height // patch_size, patch_size, width // patch_size, patch_size
).permute(0, 2, 4, 1, 3, 5).reshape(batch, -1, channels * patch_size * patch_size)
def unpatchify(patches, channels, height, width, patch_size):
batch = patches.shape[0]
return patches.reshape(
batch, height // patch_size, width // patch_size, channels, patch_size, patch_size
).permute(0, 3, 1, 4, 2, 5).reshape(batch, channels, height, width)
class Transformer(nn.Module):
def __init__(self, dim, depth, heads, mlp_ratio=4):
super().__init__()
layer = nn.TransformerEncoderLayer(
dim, heads, int(dim * mlp_ratio), activation="gelu", batch_first=True, norm_first=True
)
self.layers = nn.TransformerEncoder(layer, depth)
self.norm = nn.LayerNorm(dim)
def forward(self, values):
return self.norm(self.layers(values))
class DynamicEmbedding(nn.Module):
"""Create sensor-agnostic patches by conditioning per-band kernels on wavelength."""
def __init__(self, patch_size, embed_dim, wave_dim, wave_latents):
super().__init__()
self.patch_size = patch_size
self.wave_dim = wave_dim
self.wave_mlp = nn.Sequential(nn.Linear(wave_dim, wave_dim), nn.GELU(), nn.Linear(wave_dim, wave_dim))
self.latents = nn.Parameter(torch.randn(wave_latents, wave_dim) * 0.02)
self.cross_attention = nn.MultiheadAttention(wave_dim, 4, batch_first=True)
self.kernel = nn.Linear(wave_dim, patch_size * patch_size * embed_dim)
self.bias = nn.Parameter(torch.zeros(embed_dim))
self.embed_dim = embed_dim
def forward(self, pixels, wavelengths):
batch, channels, height, width = pixels.shape
if wavelengths.ndim == 1:
wavelengths = wavelengths[None].expand(batch, -1)
if wavelengths.shape != (batch, channels):
raise ValueError(f"wavelengths must have shape {(batch, channels)}, got {tuple(wavelengths.shape)}")
wave_features = fourier_encode(wavelengths / 1000.0, self.wave_dim)
wave_features = self.wave_mlp(wave_features)
queries = self.latents[None].expand(batch, -1, -1)
context = self.cross_attention(queries, wave_features, wave_features, need_weights=False)[0].mean(dim=1)
conditioned = wave_features + context[:, None]
kernels = self.kernel(conditioned).reshape(
batch, channels, self.embed_dim, self.patch_size, self.patch_size
)
patches = []
for index in range(batch):
patches.append(F.conv2d(pixels[index:index + 1], kernels[index].permute(1, 0, 2, 3),
stride=self.patch_size) + self.bias[None, :, None, None])
return torch.cat(patches).flatten(2).transpose(1, 2), conditioned
class DynamicDecoder(nn.Module):
def __init__(self, patch_size, decoder_dim, wave_dim):
super().__init__()
self.patch_size = patch_size
self.wave_mlp = nn.Sequential(nn.Linear(wave_dim, decoder_dim), nn.GELU(), nn.Linear(decoder_dim, decoder_dim))
self.output = nn.Linear(decoder_dim, patch_size * patch_size)
def forward(self, tokens, wave_features):
wave_context = self.wave_mlp(wave_features)
joint = tokens[:, :, None, :] + wave_context[:, None, :, :]
return self.output(joint).permute(0, 1, 2, 3).flatten(2)
class ClayFoundation(nn.Module):
def __init__(self, config):
super().__init__()
self.config = dict(config)
self.patch_size = int(config["patch_size"])
self.mask_ratio = float(config["mask_ratio"])
enc_dim, dec_dim = int(config["encoder_dim"]), int(config["decoder_dim"])
if enc_dim < 12 or (enc_dim - 8) % 4:
raise ValueError("encoder_dim - 8 must be positive and divisible by four")
if dec_dim < 12 or (dec_dim - 8) % 4:
raise ValueError("decoder_dim - 8 must be positive and divisible by four")
self.dynamic_embedding = DynamicEmbedding(
self.patch_size, enc_dim, int(config["wave_dim"]), int(config["wave_latents"])
)
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"]))
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"]))
self.wave_to_decoder = nn.Linear(int(config["wave_dim"]), int(config["wave_dim"]))
self.dynamic_decoder = DynamicDecoder(self.patch_size, dec_dim, int(config["wave_dim"]))
self.representation_head = nn.Linear(enc_dim, int(config["teacher_dim"]))
self.norm_pix_loss = bool(config.get("norm_pix_loss", False))
@staticmethod
def _metadata_encoding(time, latlon, dim):
values = torch.cat((time, latlon), dim=1)
widths = [dim // 4] * 4
for index in range(dim % 4):
widths[index] += 1
return torch.cat([fourier_encode(values[:, index], widths[index]) for index in range(4)], dim=1)
def _add_encoding(self, tokens, time, latlon, gsd):
batch, length, dim = tokens.shape
grid = int(math.sqrt(length))
if grid * grid != length:
raise ValueError("Clay reproduction requires a square patch grid")
spatial = position_encoding_2d(grid, grid, dim - 8, gsd, tokens.device, tokens.dtype)
metadata = self._metadata_encoding(time, latlon, 8)
encoding = torch.cat((spatial[None].expand(batch, -1, -1), metadata[:, None].expand(-1, length, -1)), dim=-1)
return tokens + encoding
def encode(self, pixels, time, latlon, gsd, wavelengths):
patches, _ = self.dynamic_embedding(pixels, wavelengths)
patches = self._add_encoding(patches, time, latlon, gsd)
cls = self.cls_token.expand(len(pixels), -1, -1)
encoded = self.encoder(torch.cat((cls, patches), dim=1))
return encoded[:, 0], encoded[:, 1:]
def forward(self, pixels, time, latlon, gsd, wavelengths, teacher_target=None, mask_ratio=None):
ratio = self.mask_ratio if mask_ratio is None else float(mask_ratio)
patches, wave_features = self.dynamic_embedding(pixels, wavelengths)
patches = self._add_encoding(patches, time, latlon, gsd)
batch, length, _ = patches.shape
keep = max(1, length - int(length * ratio))
noise = torch.rand(batch, length, device=pixels.device)
ordering = noise.argsort(dim=1)
unmasked_indices, masked_indices = ordering[:, :keep], ordering[:, keep:]
gather = unmasked_indices[:, :, None].expand(-1, -1, patches.shape[-1])
visible = patches.gather(1, gather)
encoded = self.encoder(torch.cat((self.cls_token.expand(batch, -1, -1), visible), dim=1))
embedding = encoded[:, 0]
decoded_visible = self.encoder_to_decoder(encoded[:, 1:])
decoder_tokens = self.mask_token.expand(batch, length, -1).clone()
decoder_tokens.scatter_(1, unmasked_indices[:, :, None].expand(-1, -1, decoded_visible.shape[-1]), decoded_visible)
grid = int(math.sqrt(length))
spatial = position_encoding_2d(grid, grid, decoder_tokens.shape[-1] - 8, gsd,
decoder_tokens.device, decoder_tokens.dtype)
metadata = self._metadata_encoding(time, latlon, 8)
decoder_tokens = decoder_tokens + torch.cat((spatial[None].expand(batch, -1, -1),
metadata[:, None].expand(-1, length, -1)), dim=-1)
decoded = self.decoder(decoder_tokens)
predicted_patches = self.dynamic_decoder(decoded, self.wave_to_decoder(wave_features))
target_patches = patchify(pixels, self.patch_size)
if self.norm_pix_loss:
mean = target_patches.mean(dim=-1, keepdim=True)
variance = target_patches.var(dim=-1, keepdim=True)
target_patches = (target_patches - mean) / (variance + 1e-6).sqrt()
mask = torch.zeros(batch, length, device=pixels.device)
mask.scatter_(1, masked_indices, 1.0)
patch_loss = (predicted_patches - target_patches).abs().mean(dim=-1)
reconstruction_loss = (patch_loss * mask).sum() / mask.sum().clamp_min(1)
projected = F.normalize(self.representation_head(embedding), dim=1)
if teacher_target is None:
representation_loss = embedding.new_zeros(())
else:
representation_loss = 1.0 - (projected * F.normalize(teacher_target, dim=1)).sum(dim=1).mean()
reconstruction = unpatchify(predicted_patches, pixels.shape[1], pixels.shape[2], pixels.shape[3], self.patch_size)
return {
"embedding": embedding,
"projected_embedding": projected,
"reconstruction": reconstruction,
"mask": mask,
"reconstruction_loss": reconstruction_loss,
"representation_loss": representation_loss,
}
def compute_loss(outputs, reconstruction_weight=0.95, representation_weight=0.05):
total = reconstruction_weight * outputs["reconstruction_loss"] + representation_weight * outputs["representation_loss"]
return total, {
"reconstruction": outputs["reconstruction_loss"],
"representation": outputs["representation_loss"],
"total": total,
}