diff options
Diffstat (limited to 'worldalign/synth_set_battery.py')
| -rw-r--r-- | worldalign/synth_set_battery.py | 283 |
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() |
