summaryrefslogtreecommitdiff
path: root/worldalign/synth_set_battery.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/synth_set_battery.py')
-rw-r--r--worldalign/synth_set_battery.py283
1 files changed, 283 insertions, 0 deletions
diff --git a/worldalign/synth_set_battery.py b/worldalign/synth_set_battery.py
new file mode 100644
index 0000000..9ed66ec
--- /dev/null
+++ b/worldalign/synth_set_battery.py
@@ -0,0 +1,283 @@
+"""Set-kernel relation fields for the synthetic world: the structure battery.
+
+The pooled readout of a set representation destroys it. Here scene states
+stay sets -- vision: slot vectors with alpha masses; text: per-group
+phrase states parsed from the caption's enumeration sentence and encoded
+individually -- and scene-to-scene relations are computed within each
+modality as set-matching similarities. The cross-modal gate then runs on
+these set-kernel relation fields exactly as on any relation channel.
+
+Phrase parsing reads only released captions; slot sets read only renders.
+Hidden pairs score orderings, as always.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from scipy.optimize import linear_sum_assignment
+from tqdm import tqdm
+
+from .common import batch_indices, read_json, seed_everything, write_json
+from .manifold_gate import standardize_relation
+from .ricci_control import run_gates
+from .synth_slots import SlotAutoencoder
+from .synth_towers import TextTower, load_image, tokenize
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument(
+ "--slots", default="artifacts/synth_v0/vision_slots.pt"
+ )
+ parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--mass-floor", type=float, default=0.02)
+ parser.add_argument(
+ "--vision-mode", choices=["slot_vectors", "sprites"], default="sprites"
+ )
+ parser.add_argument(
+ "--slot-tower", default="artifacts/synth_v0/slot_tower.pt"
+ )
+ parser.add_argument("--sprite-window", type=int, default=48)
+ parser.add_argument("--random-perms", type=int, default=300)
+ parser.add_argument("--descent-restarts", type=int, default=5)
+ parser.add_argument("--descent-max-steps", type=int, default=200000)
+ parser.add_argument(
+ "--descent-objective", default="mse", choices=["mse", "m30_total"]
+ )
+ parser.add_argument("--descent-verify-top", type=int, default=64)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v0/set_battery_gate.json"
+ )
+ return parser.parse_args()
+
+
+def parse_group_phrases(caption: str) -> list[str]:
+ """Group phrases from the enumeration sentence of a released caption."""
+ first = caption.split(".")[0]
+ for opener in ("there are ", "the picture shows ", "you can see "):
+ if first.startswith(opener):
+ first = first[len(opener):]
+ break
+ first = first.replace(" and ", ", ")
+ return [phrase.strip() for phrase in first.split(",") if phrase.strip()]
+
+
+@torch.inference_mode()
+def text_group_sets(
+ rows: list[int], captions: list[list[str]], args: argparse.Namespace
+) -> list[torch.Tensor]:
+ state = torch.load(args.text_tower, map_location="cpu", weights_only=False)
+ saved = state["args"]
+ vocab = state["vocab"]
+ model = TextTower(
+ len(vocab),
+ saved["text_dim"],
+ saved["depth"],
+ saved.get("text_heads", 4),
+ saved["context"],
+ ).to(args.device)
+ model.load_state_dict(state["model"])
+ model.eval()
+ phrases_per_row = [parse_group_phrases(captions[row][0]) for row in rows]
+ flat = [
+ (index, phrase)
+ for index, phrases in enumerate(phrases_per_row)
+ for phrase in phrases
+ ]
+ states: list[list[torch.Tensor]] = [[] for _ in rows]
+ for indices in tqdm(list(batch_indices(len(flat), 256)), desc="text sets"):
+ batch = [flat[i] for i in indices]
+ sequences = [tokenize(phrase, vocab) for _, phrase in batch]
+ longest = max(len(s) for s in sequences)
+ tokens = torch.zeros(len(batch), longest, dtype=torch.long)
+ for row, sequence in enumerate(sequences):
+ tokens[row, : len(sequence)] = torch.tensor(sequence)
+ tokens = tokens.to(args.device)
+ hidden = model(tokens)
+ mask = (tokens != 0).float()[..., None]
+ pooled = (hidden * mask).sum(1) / mask.sum(1).clamp_min(1.0)
+ for (index, _), vector in zip(batch, pooled.float().cpu()):
+ states[index].append(vector)
+ return [F.normalize(torch.stack(s), dim=-1) for s in states]
+
+
+def set_similarity_field(
+ sets: list[torch.Tensor], weights: list[torch.Tensor] | None = None
+) -> torch.Tensor:
+ """Symmetric matching-value similarity between all set pairs."""
+ n = len(sets)
+ field = torch.zeros(n, n)
+ for a in range(n):
+ for b in range(a, n):
+ similarity = sets[a] @ sets[b].T
+ if weights is not None:
+ similarity = similarity * torch.sqrt(
+ weights[a][:, None] * weights[b][None, :]
+ )
+ rows, cols = linear_sum_assignment(-similarity.numpy())
+ value = float(similarity[rows, cols].sum()) / max(
+ min(similarity.shape), 1
+ )
+ field[a, b] = field[b, a] = value
+ return field
+
+
+@torch.inference_mode()
+def decode_sprites(
+ rows: list[int],
+ lookup: dict[int, int],
+ slot_sets_all: torch.Tensor,
+ args: argparse.Namespace,
+) -> list[torch.Tensor]:
+ """Centered per-slot appearance sprites from the tower's own decoder.
+
+ The joint decode gives each slot an rgb map and a competitive alpha
+ mask; centering the masked appearance at the alpha centroid removes
+ layout, leaving color, shape, size, and multiplicity pattern.
+ """
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ state = torch.load(args.slot_tower, 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()
+ size = manifest["image_size"]
+ window = args.sprite_window
+ axis = torch.arange(size, dtype=torch.float32, device=args.device)
+ sprites: list[torch.Tensor] = []
+ for start in tqdm(range(0, len(rows), 64), desc="sprites"):
+ batch_rows = rows[start : start + 64]
+ slots = torch.stack(
+ [slot_sets_all[lookup[int(row)]][0] for row in batch_rows]
+ ).to(args.device)
+ rgb_alpha_rgb, alpha = model.decode(slots)
+ del rgb_alpha_rgb
+ # Re-decode retaining per-slot rgb: replicate decode internals.
+ batch, count, dim = slots.shape
+ x = slots.reshape(batch * count, dim, 1, 1).expand(
+ -1, -1, model.broadcast, model.broadcast
+ )
+ from .synth_slots import coordinate_grid
+
+ grid = coordinate_grid(model.broadcast, slots.device).reshape(1, -1, 4)
+ position = model.position_decoder(grid).transpose(1, 2).reshape(
+ 1, dim, model.broadcast, model.broadcast
+ )
+ decoded = model.decoder(x + position)
+ decoded = F.interpolate(
+ decoded, size=size, mode="bilinear", align_corners=False
+ ).reshape(batch, count, 4, size, size)
+ rgb = decoded[:, :, :3]
+ masked = rgb * alpha # [B, K, 3, H, W]
+ weight_y = alpha.squeeze(2).sum(-1) # [B, K, H]
+ weight_x = alpha.squeeze(2).sum(-2) # [B, K, W]
+ cy = (weight_y * axis).sum(-1) / weight_y.sum(-1).clamp_min(1e-6)
+ cx = (weight_x * axis).sum(-1) / weight_x.sum(-1).clamp_min(1e-6)
+ half = window // 2
+ batch_sprites = torch.zeros(batch, count, 3, window, window)
+ padded = F.pad(masked, (half, half, half, half))
+ for b in range(batch):
+ for k in range(count):
+ y0 = int(cy[b, k].round())
+ x0 = int(cx[b, k].round())
+ batch_sprites[b, k] = padded[
+ b, k, :, y0 : y0 + window, x0 : x0 + window
+ ].cpu()
+ sprites.extend(batch_sprites.flatten(2).unbind(0))
+ return sprites
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ captions = read_json(Path(args.data_dir, "captions.json"))["captions"]
+ rows = manifest[args.split][: args.samples]
+
+ slot_state = torch.load(args.slots, map_location="cpu", weights_only=False)
+ lookup = {int(row): i for i, row in enumerate(slot_state["rows"])}
+ slot_sets_all = slot_state["slot_sets"]
+ masses_all = slot_state["slot_masses"]
+ sprites_all = None
+ if args.vision_mode == "sprites":
+ sprites_all = decode_sprites(rows, lookup, slot_sets_all, args)
+ vision_sets, vision_weights = [], []
+ for position, row in enumerate(rows):
+ index = lookup[int(row)]
+ # Slot identity is not stable across forwards, so views cannot be
+ # averaged slot-wise; one view keeps object-slot binding intact.
+ slots = slot_sets_all[index][0] # [K, D]
+ mass = masses_all[index][0]
+ dominant = mass.argmax()
+ keep = torch.ones(len(mass), dtype=torch.bool)
+ keep[dominant] = False
+ keep &= mass > args.mass_floor
+ if not keep.any():
+ keep = torch.ones(len(mass), dtype=torch.bool)
+ if sprites_all is not None:
+ vision_sets.append(F.normalize(sprites_all[position][keep], dim=-1))
+ else:
+ vision_sets.append(F.normalize(slots[keep], dim=-1))
+ weight = mass[keep]
+ vision_weights.append(weight / weight.sum().clamp_min(1e-8))
+
+ text_sets = text_group_sets(rows, captions, args)
+
+ print(json.dumps({"building": "set similarity fields"}))
+ visual_field = set_similarity_field(vision_sets, vision_weights)
+ text_field = set_similarity_field(text_sets)
+
+ visual_channels = standardize_relation(visual_field.double())[0][None]
+ text_channels = standardize_relation(text_field.double())[0][None]
+ generator = torch.Generator().manual_seed(args.seed)
+ report = {
+ "protocol": (
+ "Scene states are sets (slot vectors; per-group phrase "
+ "states); within-modality relations are set-matching values; "
+ "hidden pairs score orderings only."
+ ),
+ "split": args.split,
+ "samples": len(rows),
+ "mean_vision_set_size": float(
+ torch.tensor([len(s) for s in vision_sets]).float().mean()
+ ),
+ "mean_text_set_size": float(
+ torch.tensor([len(s) for s in text_sets]).float().mean()
+ ),
+ **run_gates(text_channels, visual_channels, args, generator),
+ }
+ verdict = {
+ "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"],
+ "improving_fraction": report["gate_b"]["improving_fraction"],
+ "descent_keeps": report["descent_from_true"]["final_accuracy"],
+ "true_mse": report["gate_a"]["true"]["mse"],
+ "best_random_descent": min(
+ (r["final_objective"] for r in report["descent_from_random"]),
+ default=None,
+ ),
+ }
+ verdict["counterfeit_found"] = bool(
+ verdict["best_random_descent"] is not None
+ and verdict["best_random_descent"] < verdict["true_mse"]
+ )
+ report["verdict"] = verdict
+ print(json.dumps({"verdict": verdict}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()