"""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()