Download model/terramind.py from OneScience-Group/TerraMind: direct link, hf CLI and curl.
- Browser
- Download file 6.56 kB
-
https://huggingface.co/OneScience-Group/TerraMind/resolve/main/model/terramind.py
- Command line
-
hf download hf://OneScience-Group/TerraMind/model/terramind.py
-
curl -L -o terramind.py https://huggingface.co/OneScience-Group/TerraMind/resolve/main/model/terramind.py
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} | |
| 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) | |