summaryrefslogtreecommitdiff
path: root/worldalign/synth_slots.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/synth_slots.py')
-rw-r--r--worldalign/synth_slots.py264
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()