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