diff options
Diffstat (limited to 'worldalign/synth_towers.py')
| -rw-r--r-- | worldalign/synth_towers.py | 348 |
1 files changed, 348 insertions, 0 deletions
diff --git a/worldalign/synth_towers.py b/worldalign/synth_towers.py new file mode 100644 index 0000000..15221ac --- /dev/null +++ b/worldalign/synth_towers.py @@ -0,0 +1,348 @@ +"""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["<bos>"]] + [vocab[token] for token in tokens] + + +def build_vocab(manifest: dict) -> dict[str, int]: + words = list(manifest["vocabulary"]) + ["."] + vocab = {"<pad>": 0, "<bos>": 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() |
