"""From-scratch unimodal towers for the synthetic closed world. Vision: a compact ViT trained with InfoNCE over natural orbit positives -- two renders of the same scene, which differ exactly by the world's continuous nuisance. No flips (they would erase left/right semantics) and no color jitter (color is world content). Scene identity within the vision-only split is unimodal metadata. Text: a small word-level causal LM trained on the captions of the text-only split. Both towers therefore learn the same world from disjoint scenes and separate modalities, with dialable capacity. """ from __future__ import annotations import argparse import json import math import random from pathlib import Path import numpy as np import torch import torch.nn.functional as F from PIL import Image from torch import nn from .common import read_json, seed_everything def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--side", choices=["vision", "text"], required=True) parser.add_argument( "--objective", choices=["infonce", "simmim", "hybrid", "data2vec"], default="infonce" ) parser.add_argument("--mask-ratio", type=float, default=0.6) parser.add_argument("--recon-weight", type=float, default=25.0) parser.add_argument("--ema-decay", type=float, default=0.999) parser.add_argument("--data-dir", default="artifacts/synth_v0") parser.add_argument("--dim", type=int, default=192) parser.add_argument("--depth", type=int, default=6) parser.add_argument("--heads", type=int, default=3) parser.add_argument("--text-dim", type=int, default=256) parser.add_argument("--text-heads", type=int, default=4) parser.add_argument("--patch", type=int, default=16) parser.add_argument("--context", type=int, default=80) parser.add_argument("--epochs", type=int, default=80) parser.add_argument("--batch-size", type=int, default=256) parser.add_argument("--lr", type=float, default=1e-3) parser.add_argument("--temperature", type=float, default=0.2) parser.add_argument("--device", default="cuda:3") parser.add_argument("--seed", type=int, default=20260730) parser.add_argument("--output", required=True) return parser.parse_args() class Block(nn.Module): def __init__(self, dim: int, heads: int) -> None: super().__init__() self.norm1 = nn.LayerNorm(dim) self.attention = nn.MultiheadAttention(dim, heads, batch_first=True) self.norm2 = nn.LayerNorm(dim) self.mlp = nn.Sequential( nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim) ) def forward( self, x: torch.Tensor, mask: torch.Tensor | None = None ) -> torch.Tensor: normed = self.norm1(x) attended, _ = self.attention( normed, normed, normed, attn_mask=mask, need_weights=False ) x = x + attended return x + self.mlp(self.norm2(x)) class VisionTower(nn.Module): def __init__(self, image_size: int, patch: int, dim: int, depth: int, heads: int) -> None: super().__init__() self.patch = patch self.patch_embed = nn.Conv2d(3, dim, kernel_size=patch, stride=patch) tokens = (image_size // patch) ** 2 self.cls = nn.Parameter(torch.zeros(1, 1, dim)) self.mask_token = nn.Parameter(torch.zeros(1, 1, dim)) self.positions = nn.Parameter(torch.zeros(1, tokens + 1, dim)) nn.init.trunc_normal_(self.positions, std=0.02) nn.init.trunc_normal_(self.cls, std=0.02) nn.init.trunc_normal_(self.mask_token, std=0.02) self.blocks = nn.ModuleList(Block(dim, heads) for _ in range(depth)) self.norm = nn.LayerNorm(dim) self.head = nn.Sequential( nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, 128) ) self.reconstruction = nn.Linear(dim, patch * patch * 3) self.feature_head = nn.Linear(dim, dim) def encode( self, pixels: torch.Tensor, mask: torch.Tensor | None = None ) -> torch.Tensor: x = self.patch_embed(pixels).flatten(2).transpose(1, 2) if mask is not None: x = torch.where(mask[..., None], self.mask_token.expand_as(x), x) x = torch.cat([self.cls.expand(len(x), -1, -1), x], dim=1) x = x + self.positions for block in self.blocks: x = block(x) return self.norm(x) def forward(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: tokens = self.encode(pixels) return tokens[:, 0], self.head(tokens[:, 0]) class TextTower(nn.Module): def __init__(self, vocab: int, dim: int, depth: int, heads: int, context: int) -> None: super().__init__() self.embed = nn.Embedding(vocab, dim) self.positions = nn.Parameter(torch.zeros(1, context, dim)) nn.init.trunc_normal_(self.positions, std=0.02) self.blocks = nn.ModuleList(Block(dim, heads) for _ in range(depth)) self.norm = nn.LayerNorm(dim) self.context = context def forward(self, tokens: torch.Tensor) -> torch.Tensor: length = tokens.shape[1] x = self.embed(tokens) + self.positions[:, :length] mask = torch.triu( torch.full((length, length), float("-inf"), device=tokens.device), 1 ) for block in self.blocks: x = block(x, mask) return self.norm(x) def logits(self, hidden: torch.Tensor) -> torch.Tensor: return hidden @ self.embed.weight.T def load_image(path: Path) -> torch.Tensor: with Image.open(path) as image: array = np.asarray(image.convert("RGB"), dtype=np.float32) / 255.0 return torch.from_numpy(array).permute(2, 0, 1) def patchify(pixels: torch.Tensor, patch: int) -> torch.Tensor: batch, channels, height, width = pixels.shape grid = height // patch x = pixels.reshape(batch, channels, grid, patch, grid, patch) return x.permute(0, 2, 4, 3, 5, 1).reshape(batch, grid * grid, -1) def train_vision(args: argparse.Namespace) -> None: manifest = read_json(Path(args.data_dir, "manifest.json")) rows = manifest["vision_only_train"] views = manifest["visual_views"] image_dir = Path(manifest["image_dir"]) model = VisionTower( manifest["image_size"], args.patch, args.dim, args.depth, args.heads ).to(args.device) optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05) tokens = (manifest["image_size"] // args.patch) ** 2 if args.objective in ("infonce", "hybrid"): samples_per_epoch = len(rows) else: samples_per_epoch = len(rows) * views teacher = None if args.objective == "data2vec": import copy teacher = copy.deepcopy(model) for parameter in teacher.parameters(): parameter.requires_grad_(False) steps_per_epoch = max(1, samples_per_epoch // args.batch_size) schedule = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=args.epochs * steps_per_epoch ) rng = random.Random(args.seed) for epoch in range(args.epochs): total, count = 0.0, 0 if args.objective in ("infonce", "hybrid"): order = rows.copy() rng.shuffle(order) for start in range(0, len(order) - args.batch_size + 1, args.batch_size): batch_rows = order[start : start + args.batch_size] pairs = [] for row in batch_rows: first, second = rng.sample(range(views), 2) pairs.append( load_image(image_dir / f"scene{row:06d}_v{first}.png") ) pairs.append( load_image(image_dir / f"scene{row:06d}_v{second}.png") ) pixels = torch.stack(pairs).to(args.device) _, projected = model(pixels) projected = F.normalize(projected, dim=-1) logits = projected @ projected.T / args.temperature logits.fill_diagonal_(float("-inf")) targets = torch.arange(len(projected), device=args.device) ^ 1 loss = F.cross_entropy(logits, targets) if args.objective == "hybrid": mask = ( torch.rand(len(pixels), tokens, device=args.device) < args.mask_ratio ) encoded = model.encode(pixels, mask=mask) predicted = model.reconstruction(encoded[:, 1:][mask]) target_patches = patchify(pixels, args.patch)[mask] loss = loss + args.recon_weight * F.mse_loss( predicted, target_patches ) optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() schedule.step() total += float(loss) count += 1 else: jobs = [(row, view) for row in rows for view in range(views)] rng.shuffle(jobs) for start in range(0, len(jobs) - args.batch_size + 1, args.batch_size): batch = jobs[start : start + args.batch_size] pixels = torch.stack( [ load_image(image_dir / f"scene{row:06d}_v{view}.png") for row, view in batch ] ).to(args.device) mask = ( torch.rand(len(batch), tokens, device=args.device) < args.mask_ratio ) encoded = model.encode(pixels, mask=mask) if args.objective == "simmim": predicted = model.reconstruction(encoded[:, 1:][mask]) target = patchify(pixels, args.patch)[mask] loss = F.mse_loss(predicted, target) else: with torch.no_grad(): reference = teacher.encode(pixels)[:, 1:] reference = F.layer_norm( reference, reference.shape[-1:] ) predicted = model.feature_head(encoded[:, 1:][mask]) loss = F.smooth_l1_loss(predicted, reference[mask]) optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() schedule.step() if teacher is not None: with torch.no_grad(): for student_p, teacher_p in zip( model.parameters(), teacher.parameters() ): teacher_p.lerp_(student_p, 1.0 - args.ema_decay) total += float(loss) count += 1 if epoch % 5 == 0 or epoch == args.epochs - 1: print(json.dumps({"epoch": epoch, "loss": total / max(count, 1)})) torch.save( {"model": model.state_dict(), "args": vars(args), "side": "vision"}, args.output, ) print(f"Wrote {args.output}") def tokenize(caption: str, vocab: dict[str, int]) -> list[int]: tokens = caption.replace(",", " ").replace(".", " .").split() return [vocab[""]] + [vocab[token] for token in tokens] def build_vocab(manifest: dict) -> dict[str, int]: words = list(manifest["vocabulary"]) + ["."] vocab = {"": 0, "": 1} for word in sorted(set(words)): vocab.setdefault(word, len(vocab)) return vocab def train_text(args: argparse.Namespace) -> None: manifest = read_json(Path(args.data_dir, "manifest.json")) captions = read_json(Path(args.data_dir, "captions.json"))["captions"] vocab = build_vocab(manifest) sentences = [ tokenize(caption, vocab) for row in manifest["text_only_train"] for caption in captions[row] ] model = TextTower( len(vocab), args.text_dim, args.depth, args.text_heads, args.context ).to(args.device) optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01) rng = random.Random(args.seed) steps_per_epoch = max(1, len(sentences) // args.batch_size) schedule = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=args.epochs * steps_per_epoch ) for epoch in range(args.epochs): rng.shuffle(sentences) total, count = 0.0, 0 for start in range(0, len(sentences) - args.batch_size + 1, args.batch_size): batch = sentences[start : start + args.batch_size] longest = min(args.context, max(len(s) for s in batch)) tokens = torch.zeros(len(batch), longest, dtype=torch.long) for index, sentence in enumerate(batch): clipped = sentence[:longest] tokens[index, : len(clipped)] = torch.tensor(clipped) tokens = tokens.to(args.device) hidden = model(tokens) logits = model.logits(hidden[:, :-1]) targets = tokens[:, 1:] loss = F.cross_entropy( logits.reshape(-1, logits.shape[-1]), targets.reshape(-1), ignore_index=0, ) optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() schedule.step() total += float(loss) count += 1 if epoch % 5 == 0 or epoch == args.epochs - 1: print(json.dumps({"epoch": epoch, "loss": total / max(count, 1)})) torch.save( { "model": model.state_dict(), "vocab": vocab, "args": vars(args), "side": "text", }, args.output, ) print(f"Wrote {args.output}") def main() -> None: args = parse_args() seed_everything(args.seed) if args.side == "vision": train_vision(args) else: train_text(args) if __name__ == "__main__": main()