TerraMind / model /terramind.py
zhangrenchao's picture
Add engineering reproduction package
3571a70 verified
Raw
History Blame Contribute Delete
6.56 kB
"""Engineering reproduction of TerraMind dual-scale any-to-any pretraining."""
import math
import torch
from torch import nn
from torch.nn import functional as F
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 TerraMind(nn.Module):
def __init__(self, pixel_modalities, token_modalities, config):
super().__init__()
self.pixel_modalities = dict(pixel_modalities)
self.token_modalities = dict(token_modalities)
self.patch_size = int(config["patch_size"])
self.dim = int(config["dim"])
self.vocab = int(config["engineering_vocab_size"])
self.visible_fraction = float(config["visible_fraction"])
self.pixel_embeddings = nn.ModuleDict({
name: nn.Conv2d(channels, self.dim, self.patch_size, stride=self.patch_size)
for name, channels in self.pixel_modalities.items()
})
self.token_embeddings = nn.ModuleDict({
name: nn.Embedding(self.vocab, self.dim) for name in self.token_modalities
})
all_names = sorted(set(self.pixel_modalities) | set(self.token_modalities))
self.modality_ids = {name: index for index, name in enumerate(all_names)}
self.modality_embedding = nn.Embedding(len(all_names), self.dim)
self.position_embedding = nn.Parameter(torch.randn(1, 196, self.dim) * 0.02)
self.encoder = Transformer(self.dim, int(config["encoder_depth"]), int(config["heads"]),
float(config["mlp_ratio"]))
self.decoder = Transformer(self.dim, int(config["decoder_depth"]), int(config["heads"]),
float(config["mlp_ratio"]))
self.mask_tokens = nn.ParameterDict({name: nn.Parameter(torch.randn(1, 1, self.dim) * 0.02)
for name in self.token_modalities})
self.output_heads = nn.ModuleDict({name: nn.Linear(self.dim, self.vocab)
for name in self.token_modalities})
def _modality_bias(self, name, batch, length, device):
index = torch.full((batch, length), self.modality_ids[name], device=device, dtype=torch.long)
return self.modality_embedding(index)
def _visible(self, embedded, enabled):
if not enabled or embedded.shape[1] <= 2:
return embedded
keep = max(1, math.ceil(embedded.shape[1] * self.visible_fraction))
indices = torch.rand(len(embedded), embedded.shape[1], device=embedded.device).argsort(dim=1)[:, :keep]
return embedded.gather(1, indices[:, :, None].expand(-1, -1, embedded.shape[-1]))
def encode(self, pixels, tokens, apply_input_mask=False):
sequences, splits = [], {}
for name, values in pixels.items():
embedded = self.pixel_embeddings[name](values).flatten(2).transpose(1, 2)
embedded = embedded + self.position_embedding[:, :embedded.shape[1]]
embedded = embedded + self._modality_bias(name, len(values), embedded.shape[1], values.device)
embedded = self._visible(embedded, apply_input_mask)
splits[f"pixel_{name}"] = embedded.shape[1]
sequences.append(embedded)
for name, values in tokens.items():
embedded = self.token_embeddings[name](values % self.vocab)
if embedded.shape[1] == 196:
embedded = embedded + self.position_embedding
embedded = embedded + self._modality_bias(name, len(values), embedded.shape[1], values.device)
embedded = self._visible(embedded, apply_input_mask)
splits[f"token_{name}"] = embedded.shape[1]
sequences.append(embedded)
if not sequences:
raise ValueError("at least one conditioning modality is required")
return self.encoder(torch.cat(sequences, dim=1)), splits
def forward(self, pixels, tokens, target_modalities, input_token_modalities=None, apply_input_mask=True):
if target_modalities is None:
raise ValueError("target_modalities must be explicit")
if input_token_modalities is None:
input_token_modalities = [name for name in tokens if name not in target_modalities]
overlap = set(input_token_modalities) & set(target_modalities)
if overlap:
raise ValueError(f"input and target token modalities overlap: {sorted(overlap)}")
input_tokens = {name: tokens[name] for name in input_token_modalities}
encoded, splits = self.encode(pixels, input_tokens, apply_input_mask=apply_input_mask)
context = encoded.mean(dim=1, keepdim=True)
logits, losses = {}, {}
for name in target_modalities:
target = tokens[name] % self.vocab
length = target.shape[1]
query = self.mask_tokens[name].expand(len(target), length, -1)
query = query + context + self._modality_bias(name, len(target), length, target.device)
if length == 196:
query = query + self.position_embedding
prediction = self.output_heads[name](self.decoder(query))
logits[name] = prediction
losses[name] = F.cross_entropy(prediction.flatten(0, 1), target.flatten())
loss = torch.stack(list(losses.values())).mean()
return {"loss": loss, "losses": losses, "logits": logits, "embedding": encoded.mean(dim=1),
"encoder_tokens": encoded, "splits": splits}
@torch.no_grad()
def generate(self, pixels, tokens, target_modalities, input_token_modalities=None):
output = self.forward(pixels, tokens, target_modalities, input_token_modalities, apply_input_mask=False)
return {name: values.argmax(dim=-1) for name, values in output["logits"].items()}, output["embedding"]
def patch_tokens(values, patch_size, vocab):
pooled = F.avg_pool2d(values.float(), patch_size, stride=patch_size).mean(dim=1)
minimum = pooled.amin(dim=(1, 2), keepdim=True)
maximum = pooled.amax(dim=(1, 2), keepdim=True)
scaled = (pooled - minimum) / (maximum - minimum).clamp_min(1e-6)
return torch.round(scaled * (vocab - 1)).long().flatten(1)