diff options
Diffstat (limited to 'worldalign/synth_slots.py')
| -rw-r--r-- | worldalign/synth_slots.py | 264 |
1 files changed, 264 insertions, 0 deletions
diff --git a/worldalign/synth_slots.py b/worldalign/synth_slots.py new file mode 100644 index 0000000..72178d8 --- /dev/null +++ b/worldalign/synth_slots.py @@ -0,0 +1,264 @@ +"""Object-centric vision tower for the synthetic world: slot attention. + +A slot-attention autoencoder decomposes each render into K competing +slots that jointly reconstruct the image through per-slot alpha masks. +Reconstruction cannot shortcut (every object must be painted by some +slot), and the state is a SET of object vectors rather than a pooled +summary -- the literal form of the node-as-object-set doctrine. Training +is plain unimodal autoencoding on the vision-only split. + +Extraction emits a flickr-schema vision file (pooled states for the +existing gate stack) plus the slot sets and alpha masses for the +structure battery. +""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path + +import torch +import torch.nn.functional as F +from torch import nn +from tqdm import tqdm + +from .common import batch_indices, read_json, seed_everything +from .synth_towers import load_image + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--mode", choices=["train", "extract"], required=True) + parser.add_argument("--data-dir", default="artifacts/synth_v0") + parser.add_argument("--slots", type=int, default=6) + parser.add_argument("--slot-dim", type=int, default=64) + parser.add_argument("--iterations", type=int, default=3) + parser.add_argument("--epochs", type=int, default=40) + parser.add_argument("--batch-size", type=int, default=64) + parser.add_argument("--lr", type=float, default=4e-4) + parser.add_argument("--warmup-steps", type=int, default=1500) + parser.add_argument("--device", default="cuda:3") + parser.add_argument("--seed", type=int, default=20260730) + parser.add_argument("--checkpoint", default="artifacts/synth_v0/slot_tower.pt") + parser.add_argument( + "--vision-output", default="artifacts/synth_v0/vision_slots.pt" + ) + return parser.parse_args() + + +def coordinate_grid(size: int, device: torch.device) -> torch.Tensor: + axis = torch.linspace(0.0, 1.0, size, device=device) + y, x = torch.meshgrid(axis, axis, indexing="ij") + return torch.stack([x, y, 1 - x, 1 - y], dim=-1) + + +class SlotAttention(nn.Module): + def __init__(self, slots: int, dim: int, iterations: int) -> None: + super().__init__() + self.slots = slots + self.iterations = iterations + self.scale = dim**-0.5 + self.mu = nn.Parameter(torch.randn(1, 1, dim) * 0.02) + self.log_sigma = nn.Parameter(torch.zeros(1, 1, dim)) + self.norm_input = nn.LayerNorm(dim) + self.norm_slots = nn.LayerNorm(dim) + self.norm_mlp = nn.LayerNorm(dim) + self.project_q = nn.Linear(dim, dim, bias=False) + self.project_k = nn.Linear(dim, dim, bias=False) + self.project_v = nn.Linear(dim, dim, bias=False) + self.gru = nn.GRUCell(dim, dim) + self.mlp = nn.Sequential(nn.Linear(dim, dim * 2), nn.ReLU(), nn.Linear(dim * 2, dim)) + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + batch, _, dim = inputs.shape + inputs = self.norm_input(inputs) + k = self.project_k(inputs) + v = self.project_v(inputs) + slots = self.mu + self.log_sigma.exp() * torch.randn( + batch, self.slots, dim, device=inputs.device + ) + for _ in range(self.iterations): + previous = slots + q = self.project_q(self.norm_slots(slots)) + attention = F.softmax( + torch.einsum("bkd,bnd->bkn", q, k) * self.scale, dim=1 + ) + attention = attention / attention.sum(-1, keepdim=True).clamp_min(1e-8) + updates = torch.einsum("bkn,bnd->bkd", attention, v) + slots = self.gru( + updates.reshape(-1, dim), previous.reshape(-1, dim) + ).reshape(batch, self.slots, dim) + slots = slots + self.mlp(self.norm_mlp(slots)) + return slots + + +class SlotAutoencoder(nn.Module): + def __init__(self, image_size: int, slots: int, dim: int, iterations: int) -> None: + super().__init__() + self.image_size = image_size + self.encoder = nn.Sequential( + nn.Conv2d(3, dim, 5, 2, 2), nn.ReLU(), + nn.Conv2d(dim, dim, 5, 2, 2), nn.ReLU(), + nn.Conv2d(dim, dim, 5, 1, 2), nn.ReLU(), + nn.Conv2d(dim, dim, 5, 1, 2), nn.ReLU(), + ) + self.grid_size = image_size // 4 + self.position_encoder = nn.Linear(4, dim) + self.norm = nn.LayerNorm(dim) + self.pre_mlp = nn.Sequential(nn.Linear(dim, dim), nn.ReLU(), nn.Linear(dim, dim)) + self.slot_attention = SlotAttention(slots, dim, iterations) + self.broadcast = 8 + self.position_decoder = nn.Linear(4, dim) + self.decoder = nn.Sequential( + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.ConvTranspose2d(dim, dim, 4, 2, 1), nn.ReLU(), + nn.Conv2d(dim, 4, 3, 1, 1), + ) + + def encode(self, pixels: torch.Tensor) -> torch.Tensor: + features = self.encoder(pixels) + batch, dim = features.shape[:2] + features = features.flatten(2).transpose(1, 2) + grid = coordinate_grid(self.grid_size, pixels.device).reshape(1, -1, 4) + features = features + self.position_encoder(grid) + return self.slot_attention(self.pre_mlp(self.norm(features))) + + def decode(self, slots: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + batch, count, dim = slots.shape + x = slots.reshape(batch * count, dim, 1, 1).expand( + -1, -1, self.broadcast, self.broadcast + ) + grid = coordinate_grid(self.broadcast, slots.device).reshape(1, -1, 4) + position = self.position_decoder(grid).transpose(1, 2).reshape( + 1, dim, self.broadcast, self.broadcast + ) + decoded = self.decoder(x + position) + decoded = F.interpolate( + decoded, size=self.image_size, mode="bilinear", align_corners=False + ) + decoded = decoded.reshape(batch, count, 4, self.image_size, self.image_size) + rgb, alpha = decoded[:, :, :3], decoded[:, :, 3:] + alpha = F.softmax(alpha, dim=1) + return (rgb * alpha).sum(1), alpha + + def forward(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + slots = self.encode(pixels) + reconstruction, alpha = self.decode(slots) + return reconstruction, slots, alpha + + +def train(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 = SlotAutoencoder( + manifest["image_size"], args.slots, args.slot_dim, args.iterations + ).to(args.device) + optimizer = torch.optim.Adam(model.parameters(), lr=args.lr) + jobs = [(row, view) for row in rows for view in range(views)] + steps_total = args.epochs * (len(jobs) // args.batch_size) + schedule = torch.optim.lr_scheduler.LambdaLR( + optimizer, + lambda step: min(1.0, step / max(args.warmup_steps, 1)) + * 0.5 + * (1 + torch.cos(torch.tensor(step / max(steps_total, 1) * 3.14159)).item()), + ) + import random + + rng = random.Random(args.seed) + step = 0 + for epoch in range(args.epochs): + rng.shuffle(jobs) + total, count = 0.0, 0 + 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) + reconstruction, _, _ = model(pixels) + loss = F.mse_loss(reconstruction, pixels) + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + optimizer.step() + schedule.step() + step += 1 + 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)}, args.checkpoint) + print(f"Wrote {args.checkpoint}") + + +@torch.inference_mode() +def extract(args: argparse.Namespace) -> None: + manifest = read_json(Path(args.data_dir, "manifest.json")) + state = torch.load(args.checkpoint, map_location="cpu", weights_only=False) + saved = state["args"] + model = SlotAutoencoder( + manifest["image_size"], saved["slots"], saved["slot_dim"], saved["iterations"] + ).to(args.device) + model.load_state_dict(state["model"]) + model.eval() + views = manifest["visual_views"] + image_dir = Path(manifest["image_dir"]) + rendered_rows = sorted( + set(manifest["vision_only_train"]) | set(manifest["val"]) | set(manifest["test"]) + ) + jobs = [(row, view) for row in rendered_rows for view in range(views)] + slot_sets, masses = [], [] + for indices in tqdm(list(batch_indices(len(jobs), args.batch_size)), desc="slots"): + pixels = torch.stack( + [ + load_image(image_dir / f"scene{jobs[i][0]:06d}_v{jobs[i][1]}.png") + for i in indices + ] + ).to(args.device) + _, slots, alpha = model(pixels) + slot_sets.append(slots.float().cpu()) + masses.append(alpha.mean(dim=(2, 3, 4)).float().cpu()) + slot_sets = torch.cat(slot_sets).reshape( + len(rendered_rows), views, saved["slots"], saved["slot_dim"] + ) + masses = torch.cat(masses).reshape(len(rendered_rows), views, saved["slots"]) + # Pooled state: mass-weighted mean of the non-dominant slots. The + # largest-mass slot is the background in this world (the scene is + # mostly background) and would otherwise dominate the pool. + dominant = masses.argmax(-1, keepdim=True) + keep = torch.ones_like(masses).scatter(-1, dominant, 0.0) + weights = (masses * keep).clamp_min(1e-8) + weights = weights / weights.sum(-1, keepdim=True) + pooled = (slot_sets * weights[..., None]).sum(2).mean(1) + torch.save( + { + "model": "synth_slot_tower", + "rows": rendered_rows, + "features": F.normalize(pooled, dim=-1), + "slot_sets": slot_sets, + "slot_masses": masses, + "views_per_scene": views, + }, + args.vision_output, + ) + print(f"Wrote {args.vision_output}") + + +def main() -> None: + args = parse_args() + seed_everything(args.seed) + if args.mode == "train": + train(args) + else: + extract(args) + + +if __name__ == "__main__": + main() |
