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