summaryrefslogtreecommitdiff
path: root/worldalign/synth_set_battery.py
blob: 9ed66ec73dc0c8ae838c485116753ff8b64ed5d0 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
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()