summaryrefslogtreecommitdiff
path: root/worldalign
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign')
-rw-r--r--worldalign/__init__.py4
-rw-r--r--worldalign/anchor_battery.py274
-rw-r--r--worldalign/blind_recovery.py476
-rw-r--r--worldalign/common.py142
-rw-r--r--worldalign/content_gate.py164
-rw-r--r--worldalign/content_projection.py214
-rw-r--r--worldalign/cooccurrence_dictionary.py237
-rw-r--r--worldalign/diagnose.py76
-rw-r--r--worldalign/energy.py137
-rw-r--r--worldalign/energy_coupling.py205
-rw-r--r--worldalign/energy_functional_infer.py262
-rw-r--r--worldalign/energy_infer.py249
-rw-r--r--worldalign/evaluate.py149
-rw-r--r--worldalign/evaluate_energy_prefix.py151
-rw-r--r--worldalign/evaluate_prefix.py93
-rw-r--r--worldalign/extract_text.py95
-rw-r--r--worldalign/extract_text_orbits.py108
-rw-r--r--worldalign/extract_vision.py61
-rw-r--r--worldalign/gw.py79
-rw-r--r--worldalign/io.py30
-rw-r--r--worldalign/manifold_gate.py677
-rw-r--r--worldalign/models.py90
-rw-r--r--worldalign/natural_families.py264
-rw-r--r--worldalign/natural_fields.py224
-rw-r--r--worldalign/natural_objects.py262
-rw-r--r--worldalign/natural_pipeline.py204
-rw-r--r--worldalign/precompute_gw.py51
-rw-r--r--worldalign/prepare.py80
-rw-r--r--worldalign/ricci_control.py258
-rw-r--r--worldalign/spectral_match.py172
-rw-r--r--worldalign/synth_cc_battery.py377
-rw-r--r--worldalign/synth_decode.py150
-rw-r--r--worldalign/synth_deep_gate.py203
-rw-r--r--worldalign/synth_extract.py144
-rw-r--r--worldalign/synth_fast_gate.py345
-rw-r--r--worldalign/synth_precision_gate.py279
-rw-r--r--worldalign/synth_probes.py129
-rw-r--r--worldalign/synth_recovery.py184
-rw-r--r--worldalign/synth_set_battery.py283
-rw-r--r--worldalign/synth_slots.py264
-rw-r--r--worldalign/synth_towers.py348
-rw-r--r--worldalign/synth_transfer_energy.py276
-rw-r--r--worldalign/synth_triangle_gate.py296
-rw-r--r--worldalign/synth_triangle_recovery.py173
-rw-r--r--worldalign/synth_world.py506
-rw-r--r--worldalign/tier0_dictionary.py324
-rw-r--r--worldalign/tier0_pipeline.py310
-rw-r--r--worldalign/train_bridge.py255
-rw-r--r--worldalign/train_prefix.py141
-rw-r--r--worldalign/vg_attention_extract.py269
-rw-r--r--worldalign/vg_attention_probe.py170
-rw-r--r--worldalign/vg_diagnose.py200
-rw-r--r--worldalign/vg_extract_text.py111
-rw-r--r--worldalign/vg_extract_vision.py150
-rw-r--r--worldalign/vg_graph_diagnose.py222
-rw-r--r--worldalign/vg_prepare.py265
-rw-r--r--worldalign/vg_view_probe.py292
-rw-r--r--worldalign/view_anchor_affinity.py271
-rw-r--r--worldalign/view_anchor_battery.py365
-rw-r--r--worldalign/view_gate.py436
-rw-r--r--worldalign/world_match.py138
61 files changed, 13364 insertions, 0 deletions
diff --git a/worldalign/__init__.py b/worldalign/__init__.py
new file mode 100644
index 0000000..4013580
--- /dev/null
+++ b/worldalign/__init__.py
@@ -0,0 +1,4 @@
+"""Unpaired vision-language representation alignment experiments."""
+
+__version__ = "0.1.0"
+
diff --git a/worldalign/anchor_battery.py b/worldalign/anchor_battery.py
new file mode 100644
index 0000000..9594168
--- /dev/null
+++ b/worldalign/anchor_battery.py
@@ -0,0 +1,274 @@
+"""C5 battery: gauge-free unary anchors from raw observations.
+
+Two declared anchor classes, per `STRUCTURE_DESIGN.md`:
+
+1. pure invariants -- quantities whose cross-modal identity needs no
+ specification at all (numeral content of phrases);
+2. minimally specified physical invariants -- fixed world-knowledge maps
+ (color lexicon to hue bands, size words to box areas, luminance words
+ to pixel luminance), frozen before evaluation and never tuned against
+ pairs.
+
+Anchors are unary and computed independently per modality: text from the
+released phrase bundles, vision from cached image pixels and the region
+boxes of the private preprocessing file (vision-side observables). Hidden
+pairs score anchor quality; the anchor matrix itself is built without
+them and feeds the tempering search as side information.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from concurrent.futures import ThreadPoolExecutor
+from pathlib import Path
+
+import numpy as np
+import torch
+from PIL import Image
+
+from .common import write_json
+
+NUMBER_WORDS = {
+ "one": 1, "two": 2, "three": 3, "four": 4, "five": 5, "six": 6,
+ "seven": 7, "eight": 8, "nine": 9, "ten": 10, "eleven": 11,
+ "twelve": 12, "dozen": 12,
+}
+# Hue band centers in degrees on the HSV wheel; frozen world knowledge.
+COLOR_BANDS = {
+ "red": 0.0, "orange": 30.0, "yellow": 60.0, "green": 120.0,
+ "cyan": 180.0, "blue": 220.0, "purple": 275.0, "pink": 330.0,
+}
+ACHROMATIC = ("white", "black", "gray", "grey", "brown", "tan", "beige")
+SIZE_WORDS = {
+ "huge": 2.0, "large": 1.5, "big": 1.5, "tall": 1.0, "long": 1.0,
+ "small": -1.0, "little": -1.0, "tiny": -2.0, "short": -0.5,
+}
+LIGHT_WORDS = {
+ "white": 1.0, "bright": 1.0, "light": 0.5, "sunny": 1.0,
+ "black": -1.0, "dark": -1.0, "shadow": -0.5, "night": -1.0,
+}
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images")
+ parser.add_argument("--image-size", type=int, default=224)
+ parser.add_argument("--workers", type=int, default=16)
+ parser.add_argument(
+ "--output", default="artifacts/manifold_gate/anchor_battery.json"
+ )
+ parser.add_argument(
+ "--matrix-output", default="artifacts/manifold_gate/anchor_unary.pt"
+ )
+ return parser.parse_args()
+
+
+def read_jsonl(path: Path) -> list[dict]:
+ return [
+ json.loads(line)
+ for line in path.read_text(encoding="utf-8").splitlines()
+ if line.strip()
+ ]
+
+
+def text_anchors(phrases: list[str]) -> dict[str, np.ndarray | float]:
+ tokens = [
+ token
+ for phrase in phrases
+ for token in re.findall(r"[a-z]+|\d+", phrase.lower())
+ ]
+ numeral_total = 0.0
+ for token in tokens:
+ if token.isdigit() and int(token) <= 50:
+ numeral_total += int(token)
+ elif token in NUMBER_WORDS:
+ numeral_total += NUMBER_WORDS[token]
+ color_histogram = np.zeros(len(COLOR_BANDS))
+ for index, color in enumerate(COLOR_BANDS):
+ color_histogram[index] = sum(token == color for token in tokens)
+ achromatic = sum(token in ACHROMATIC for token in tokens)
+ chroma_score = color_histogram.sum() / max(color_histogram.sum() + achromatic, 1)
+ size_score = float(sum(SIZE_WORDS.get(token, 0.0) for token in tokens))
+ light_score = float(sum(LIGHT_WORDS.get(token, 0.0) for token in tokens))
+ return {
+ "numeral_total": float(numeral_total),
+ "color_histogram": color_histogram,
+ "chroma_score": float(chroma_score),
+ "size_score": size_score,
+ "light_score": light_score,
+ }
+
+
+def vision_anchors(record: dict, args: argparse.Namespace) -> dict:
+ path = Path(args.image_cache, f"{record['source_image_id']}.jpg")
+ with Image.open(path) as image:
+ hsv = image.convert("HSV").resize(
+ (args.image_size, args.image_size), Image.BILINEAR
+ )
+ values = np.asarray(hsv, dtype=np.float64)
+ hue = values[..., 0] * (360.0 / 255.0)
+ saturation = values[..., 1] / 255.0
+ value = values[..., 2] / 255.0
+ chromatic = (saturation > 0.25) & (value > 0.2)
+ hue_histogram = np.zeros(len(COLOR_BANDS))
+ for index, center in enumerate(COLOR_BANDS.values()):
+ distance = np.minimum(np.abs(hue - center), 360.0 - np.abs(hue - center))
+ hue_histogram[index] = float(((distance < 25.0) & chromatic).mean())
+ width, height = record["width"], record["height"]
+ areas = [
+ (region["width"] * region["height"]) / max(width * height, 1)
+ for region in record["regions"]
+ ]
+ return {
+ "hue_histogram": hue_histogram,
+ "chroma_fraction": float(chromatic.mean()),
+ "luminance_mean": float(value.mean()),
+ "mean_box_area": float(np.mean(areas)),
+ "log_mean_box_area": float(np.log(np.mean(areas) + 1e-6)),
+ }
+
+
+def rank_z(x: np.ndarray) -> np.ndarray:
+ order = x.argsort().argsort().astype(np.float64)
+ order = (order - order.mean()) / max(order.std(), 1e-9)
+ return order
+
+
+def spearman(x: np.ndarray, y: np.ndarray) -> float:
+ return float(np.corrcoef(rank_z(x), rank_z(y))[0, 1])
+
+
+def main() -> None:
+ args = parse_args()
+ text_nodes = read_jsonl(Path(args.vg_dir, "text_nodes.jsonl"))
+ vision_nodes = read_jsonl(Path(args.vg_dir, "vision_nodes.private.jsonl"))
+ truth = read_jsonl(Path(args.vg_dir, "ground_truth.private.jsonl"))
+
+ text_features = {
+ record["node_id"]: text_anchors(record["region_closed"])
+ for record in text_nodes
+ }
+ with ThreadPoolExecutor(max_workers=args.workers) as pool:
+ vision_results = list(
+ pool.map(lambda record: vision_anchors(record, args), vision_nodes)
+ )
+ vision_features = {
+ record["node_id"]: result
+ for record, result in zip(vision_nodes, vision_results)
+ }
+
+ # Aligned arrays in ground-truth order: index i pairs vision i, text i.
+ vision_ids = [pair["vision_node_id"] for pair in truth]
+ text_ids = [pair["text_node_id"] for pair in truth]
+ n = len(truth)
+
+ def text_array(key: str) -> np.ndarray:
+ return np.array([text_features[node][key] for node in text_ids])
+
+ def vision_array(key: str) -> np.ndarray:
+ return np.array([vision_features[node][key] for node in vision_ids])
+
+ declared_pairs = {
+ "color": ("color_histogram", "hue_histogram"),
+ "chroma": ("chroma_score", "chroma_fraction"),
+ "light": ("light_score", "luminance_mean"),
+ "size": ("size_score", "log_mean_box_area"),
+ "numerals_vs_area": ("numeral_total", "log_mean_box_area"),
+ }
+
+ report: dict = {
+ "protocol": (
+ "Unary anchors, computed independently per modality; hidden "
+ "pairs score anchor quality only. Lexicon-to-physics maps are "
+ "frozen world knowledge, declared before evaluation."
+ ),
+ "nodes": n,
+ "anchors": {},
+ }
+ channels: list[np.ndarray] = []
+ channel_names: list[str] = []
+
+ # Color: cosine between lexicon histogram and hue-band histogram under
+ # the frozen band map, evaluated for every cross pair.
+ text_color = np.stack([text_features[node]["color_histogram"] for node in text_ids])
+ vision_color = np.stack(
+ [vision_features[node]["hue_histogram"] for node in vision_ids]
+ )
+ text_color_norm = text_color / np.linalg.norm(text_color, axis=1, keepdims=True).clip(
+ min=1e-9
+ )
+ vision_color_norm = vision_color / np.linalg.norm(
+ vision_color, axis=1, keepdims=True
+ ).clip(min=1e-9)
+ color_similarity = vision_color_norm @ text_color_norm.T # [vision, text]
+ has_signal = (text_color.sum(1) > 0)[None, :] * np.ones((n, 1))
+ matched_color = np.diag(color_similarity)
+ informative = text_color.sum(1) > 0
+ shuffle = np.random.default_rng(0).permutation(n)
+ report["anchors"]["color_histogram"] = {
+ "informative_text_fraction": float(informative.mean()),
+ "matched_mean": float(matched_color[informative].mean()),
+ "shuffled_mean": float(color_similarity[shuffle, np.arange(n)][informative].mean()),
+ "matched_minus_shuffled_z": float(
+ (matched_color[informative] - color_similarity[shuffle, np.arange(n)][informative]).mean()
+ / (matched_color[informative] - color_similarity[shuffle, np.arange(n)][informative]).std()
+ * np.sqrt(informative.sum())
+ ),
+ }
+ channels.append(color_similarity * has_signal)
+ channel_names.append("color")
+
+ for name, (text_key, vision_key) in declared_pairs.items():
+ if name == "color":
+ continue
+ t = text_array(text_key)
+ v = vision_array(vision_key)
+ rho = spearman(t, v)
+ report["anchors"][name] = {"true_pairing_spearman": rho}
+ similarity = -np.abs(rank_z(v)[:, None] - rank_z(t)[None, :])
+ channels.append(similarity)
+ channel_names.append(name)
+
+ # Combined unary similarity: mean of per-channel rank-standardized maps.
+ stacked = []
+ for channel in channels:
+ flat = channel.flatten()
+ standardized = (channel - flat.mean()) / max(flat.std(), 1e-9)
+ stacked.append(standardized)
+ combined = np.mean(stacked, axis=0)
+ matched = np.diag(combined)
+ ranks = (combined >= matched[:, None]).sum(1)
+ report["combined_unary"] = {
+ "channels": channel_names,
+ "matched_mean": float(matched.mean()),
+ "grand_mean": float(combined.mean()),
+ "true_z": float(
+ (matched.mean() - combined.mean())
+ / combined.std()
+ * np.sqrt(n)
+ ),
+ "retrieval_r@10": float((ranks <= 10).mean()),
+ "retrieval_r@100": float((ranks <= 100).mean()),
+ "median_rank": float(np.median(ranks)),
+ "chance_r@10": 10.0 / n,
+ }
+ torch.save(
+ {
+ "vision_node_ids": vision_ids,
+ "text_node_ids": text_ids,
+ "order_note": "row i = vision node truth[i]; column j = text node truth[j]",
+ "channels": {name: torch.from_numpy(np.asarray(channel)) for name, channel in zip(channel_names, channels)},
+ "combined": torch.from_numpy(combined),
+ },
+ args.matrix_output,
+ )
+ write_json(args.output, report)
+ print(json.dumps(report, indent=2)[:2200])
+ print(f"Wrote {args.output} and {args.matrix_output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/blind_recovery.py b/worldalign/blind_recovery.py
new file mode 100644
index 0000000..625f407
--- /dev/null
+++ b/worldalign/blind_recovery.py
@@ -0,0 +1,476 @@
+"""Blind unpaired recovery on gate-passing states.
+
+The basin audit licenses this experiment: on content-projected VG states
+the true assignment sits below every quench, so the remaining question is
+search. Three arms attack the measured landscape shape (deep true basin,
+glassy surroundings, smooth coarse spectrum):
+
+A. entropic Sinkhorn annealing: soft coupling, temperature and entropy
+ schedules, barycentric language states;
+B. parallel tempering: replica-exchange Metropolis over permutations with
+ the exact closed-form swap deltas, imported from spin-glass practice;
+C. spectral-band homotopy: solve the coupling in a low-dimensional
+ spectral band first, then warm-start progressively finer bands --
+ a band-limited continuation in the sense of low-order Fourier terms on
+ the symmetric group, instantiated through the content spectrum.
+
+The text side enters through a hidden shuffle, so the optimizer's identity
+carries no information. The hidden truth scores results and is used
+nowhere else.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from scipy.optimize import linear_sum_assignment
+
+from .common import seed_everything, write_json
+from .content_gate import flickr_states, vg_states
+from .energy import log_sinkhorn
+from .manifold_gate import standardize_relation
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--dataset", choices=["flickr", "vg"], default="vg")
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt")
+ parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt")
+ parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt")
+ parser.add_argument(
+ "--vg-ground-truth", default="artifacts/vg_5k/ground_truth.private.jsonl"
+ )
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--subset-seed", type=int, default=0)
+ parser.add_argument("--dims", type=int, default=256)
+ parser.add_argument("--shrinkage", type=float, default=0.05)
+ parser.add_argument("--arms", default="tempering,sinkhorn,homotopy")
+ parser.add_argument("--replicas", type=int, default=8)
+ parser.add_argument("--tempering-rounds", type=int, default=40000)
+ parser.add_argument("--temp-high", type=float, default=3e-3)
+ parser.add_argument("--temp-low", type=float, default=1e-5)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--sinkhorn-steps", type=int, default=2500)
+ parser.add_argument("--sinkhorn-restarts", type=int, default=4)
+ parser.add_argument("--homotopy-bands", default="4,8,16,32,64,full")
+ parser.add_argument(
+ "--unary-matrix",
+ default="",
+ help="Optional anchor similarity matrix (.pt) in ground-truth node "
+ "order; injected as a unary field on the tempering energy.",
+ )
+ parser.add_argument("--unary-weight", type=float, default=2.0)
+ parser.add_argument("--init", choices=["random", "unary"], default="random")
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260730)
+ parser.add_argument("--output", required=True)
+ return parser.parse_args()
+
+
+def relation_mse_soft(
+ language_states: torch.Tensor, visual_standardized: torch.Tensor
+) -> torch.Tensor:
+ """Differentiable standardized-relation MSE for soft states."""
+ normalized = F.normalize(language_states, dim=-1)
+ relation = normalized @ normalized.T
+ mask = ~torch.eye(len(relation), dtype=torch.bool, device=relation.device)
+ values = relation[mask]
+ standardized = (values - values.mean()) / values.std().clamp_min(1e-6)
+ return (standardized - visual_standardized[mask]).square().mean()
+
+
+def permutation_energy_batch(
+ text_standardized: torch.Tensor,
+ visual_standardized: torch.Tensor,
+ permutations: torch.Tensor,
+) -> torch.Tensor:
+ fields = text_standardized[permutations[:, :, None], permutations[:, None, :]]
+ mask = ~torch.eye(
+ text_standardized.shape[-1], dtype=torch.bool, device=fields.device
+ )
+ return ((fields - visual_standardized) ** 2)[:, mask].mean(-1)
+
+
+def proposal_swap_deltas(
+ permuted_fields: torch.Tensor,
+ visual_standardized: torch.Tensor,
+ pairs_p: torch.Tensor,
+ pairs_q: torch.Tensor,
+) -> torch.Tensor:
+ """Exact deltas for proposed swaps only, O(N) per proposal.
+
+ permuted_fields: [R, N, N] text fields under each replica's current
+ permutation; pairs_p/pairs_q: [R, P] proposal endpoints.
+ """
+ size = permuted_fields.shape[-1]
+ count = size * (size - 1)
+ replica_index = torch.arange(len(permuted_fields), device=permuted_fields.device)
+ rows_p = permuted_fields[replica_index[:, None], pairs_p] # [R, P, N]
+ rows_q = permuted_fields[replica_index[:, None], pairs_q]
+ visual_p = visual_standardized[pairs_p]
+ visual_q = visual_standardized[pairs_q]
+ self_p = (rows_p * visual_p).sum(-1)
+ self_q = (rows_q * visual_q).sum(-1)
+ cross_pq = (rows_p * visual_q).sum(-1)
+ cross_qp = (rows_q * visual_p).sum(-1)
+ direct = permuted_fields[replica_index[:, None], pairs_p, pairs_q]
+ direct_visual = visual_standardized[pairs_p, pairs_q]
+ total = self_p + self_q - cross_pq - cross_qp - 2.0 * direct * direct_visual
+ delta = (4.0 / count) * total
+ return delta.masked_fill(pairs_p == pairs_q, float("inf"))
+
+
+def score(
+ permutation: torch.Tensor, truth: torch.Tensor, energy: float, true_energy: float
+) -> dict:
+ return {
+ "accuracy": float((permutation.cpu() == truth.cpu()).double().mean()),
+ "energy": energy,
+ "energy_over_true": energy / true_energy - 1.0,
+ }
+
+
+def coupling_metrics(coupling: torch.Tensor, truth: torch.Tensor) -> dict:
+ n = len(coupling)
+ truth = truth.to(coupling.device)
+ mass_at_truth = float(coupling[torch.arange(n, device=coupling.device), truth].mean())
+ rows, cols = linear_sum_assignment(-coupling.detach().cpu().numpy())
+ permutation = torch.from_numpy(cols)
+ forward = coupling.argmax(-1).cpu()
+ backward = coupling.argmax(0).cpu()
+ mutual = backward[forward] == torch.arange(n)
+ sorted_mass = coupling.sort(-1, descending=True).values
+ margin = (sorted_mass[:, 0] - sorted_mass[:, 1]).cpu()
+ top = margin.argsort(descending=True)[: max(1, n // 10)]
+ return {
+ "mass_at_truth": mass_at_truth,
+ "hungarian_accuracy": float((permutation == truth.cpu()).double().mean()),
+ "mutual_nn_count": int(mutual.sum()),
+ "mutual_nn_precision": float(
+ (forward[mutual] == truth.cpu()[mutual]).double().mean()
+ )
+ if mutual.any()
+ else None,
+ "top_margin_decile_precision": float(
+ (forward[top] == truth.cpu()[top]).double().mean()
+ ),
+ }
+
+
+def arm_tempering(
+ text_standardized: torch.Tensor,
+ visual_standardized: torch.Tensor,
+ truth: torch.Tensor,
+ true_energy: float,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+ unary: torch.Tensor | None = None,
+) -> dict:
+ device = text_standardized.device
+ size = text_standardized.shape[-1]
+ replicas = args.replicas
+ weight = args.unary_weight if unary is not None else 0.0
+
+ def total_energy(permutations: torch.Tensor) -> torch.Tensor:
+ energy = permutation_energy_batch(
+ text_standardized, visual_standardized, permutations
+ )
+ if unary is not None:
+ index = torch.arange(size, device=device)
+ energy = energy - weight * unary[index, permutations].mean(-1)
+ return energy
+
+ temperatures = torch.logspace(
+ torch.log10(torch.tensor(args.temp_low)),
+ torch.log10(torch.tensor(args.temp_high)),
+ replicas,
+ ).to(device)
+ if args.init == "unary" and unary is not None:
+ rows, cols = linear_sum_assignment(-unary.cpu().numpy())
+ start = torch.from_numpy(cols).to(device)
+ permutations = torch.stack([start.clone() for _ in range(replicas)])
+ else:
+ permutations = torch.stack(
+ [
+ torch.randperm(size, generator=generator).to(device)
+ for _ in range(replicas)
+ ]
+ )
+ energies = total_energy(permutations)
+ best = {"energy": float("inf"), "permutation": permutations[0].clone()}
+ accepted = 0
+ for round_index in range(args.tempering_rounds):
+ fields = text_standardized[
+ permutations[:, :, None], permutations[:, None, :]
+ ]
+ pairs_p = torch.randint(0, size, (replicas, 24), generator=generator).to(
+ device
+ )
+ pairs_q = torch.randint(0, size, (replicas, 24), generator=generator).to(
+ device
+ )
+ proposal_deltas = proposal_swap_deltas(
+ fields, visual_standardized, pairs_p, pairs_q
+ )
+ if unary is not None:
+ replica_index = torch.arange(replicas, device=device)[:, None]
+ assigned_p = permutations[replica_index, pairs_p]
+ assigned_q = permutations[replica_index, pairs_q]
+ unary_delta = (
+ unary[pairs_p, assigned_q]
+ + unary[pairs_q, assigned_p]
+ - unary[pairs_p, assigned_p]
+ - unary[pairs_q, assigned_q]
+ )
+ proposal_deltas = proposal_deltas - (weight / size) * unary_delta
+ noise = torch.rand(replicas, 24, generator=generator).to(device)
+ acceptable = (
+ proposal_deltas < -temperatures[:, None] * noise.clamp_min(1e-12).log()
+ ) & torch.isfinite(proposal_deltas)
+ for replica in range(replicas):
+ hits = torch.nonzero(acceptable[replica])
+ if not len(hits):
+ continue
+ first = int(hits[0, 0])
+ p = int(pairs_p[replica, first])
+ q = int(pairs_q[replica, first])
+ permutations[replica][[p, q]] = permutations[replica][[q, p]]
+ energies[replica] = energies[replica] + proposal_deltas[replica, first]
+ accepted += 1
+ if round_index % args.exchange_every == 0:
+ for replica in range(replicas - 1):
+ gap = (energies[replica] - energies[replica + 1]) * (
+ 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1]
+ )
+ if gap > 0 or torch.rand(1, generator=generator).item() < float(
+ gap.exp()
+ ):
+ permutations[[replica, replica + 1]] = permutations[
+ [replica + 1, replica]
+ ]
+ energies[[replica, replica + 1]] = energies[
+ [replica + 1, replica]
+ ]
+ cold = int(energies.argmin())
+ if float(energies[cold]) < best["energy"]:
+ best = {
+ "energy": float(energies[cold]),
+ "permutation": permutations[cold].clone(),
+ }
+ if round_index % 5000 == 0:
+ energies = total_energy(permutations) # refresh against drift
+ print(
+ json.dumps(
+ {
+ "tempering_round": round_index,
+ "cold_energy": float(energies.min()),
+ "cold_accuracy": float(
+ (permutations[int(energies.argmin())].cpu() == truth)
+ .double()
+ .mean()
+ ),
+ "accepted": accepted,
+ }
+ )
+ )
+ final = [
+ score(
+ permutations[r],
+ truth,
+ float(
+ permutation_energy_batch(
+ text_standardized, visual_standardized, permutations[r : r + 1]
+ )[0]
+ ),
+ true_energy,
+ )
+ for r in range(replicas)
+ ]
+ best_score = score(
+ best["permutation"],
+ truth,
+ float(
+ permutation_energy_batch(
+ text_standardized, visual_standardized, best["permutation"][None]
+ )[0]
+ ),
+ true_energy,
+ )
+ return {"replicas": final, "best": best_score, "accepted_moves": accepted}
+
+
+def arm_sinkhorn(
+ text_states: torch.Tensor,
+ visual_standardized: torch.Tensor,
+ truth: torch.Tensor,
+ true_energy: float,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+ bands: list[int | None] | None = None,
+) -> dict:
+ device = text_states.device
+ size = len(text_states)
+ restarts = []
+ for restart in range(args.sinkhorn_restarts if bands is None else 1):
+ logits = torch.nn.Parameter(
+ 0.01
+ * torch.randn(size, size, generator=generator).to(device)
+ )
+ optimizer = torch.optim.Adam([logits], lr=0.08)
+ stage_states = text_states
+ stages = bands if bands is not None else [None]
+ for stage_index, band in enumerate(stages):
+ if band is not None:
+ keep = min(band, text_states.shape[-1])
+ stage_states = text_states[:, :keep]
+ steps = args.sinkhorn_steps // len(stages)
+ for step in range(steps):
+ progress = step / max(steps - 1, 1)
+ temperature = 0.5 * (0.05 / 0.5) ** progress
+ coupling = log_sinkhorn(logits, temperature, iterations=12)
+ barycenter = coupling @ stage_states
+ loss = relation_mse_soft(barycenter, visual_standardized)
+ entropy = -(coupling * coupling.clamp_min(1e-12).log()).sum(-1).mean()
+ loss = loss + (0.002 + 0.02 * progress) * entropy
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ optimizer.step()
+ with torch.no_grad():
+ coupling = log_sinkhorn(logits, 0.05, iterations=30)
+ metrics = coupling_metrics(coupling, truth)
+ rows, cols = linear_sum_assignment(-coupling.detach().cpu().numpy())
+ permutation = torch.from_numpy(cols).to(device)
+ rounded_energy = float(
+ permutation_energy_batch(
+ standardize_relation(
+ F.normalize(text_states, dim=-1)
+ @ F.normalize(text_states, dim=-1).T
+ )[0][None].squeeze(0),
+ visual_standardized,
+ permutation[None],
+ )[0]
+ )
+ metrics.update(
+ score(permutation, truth, rounded_energy, true_energy)
+ )
+ restarts.append(metrics)
+ print(json.dumps({"sinkhorn_restart": restart, **metrics}))
+ return {"restarts": restarts}
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ if args.dataset == "flickr":
+ text_states, visual_states = flickr_states(args)
+ else:
+ text_states, visual_states = vg_states(args)
+ device = torch.device(args.device)
+ generator = torch.Generator().manual_seed(args.seed)
+
+ size = len(text_states)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden) # input row j holds true counterpart hidden[j]
+ text_input = text_states[hidden].float().to(device)
+ visual_states = visual_states.float().to(device)
+
+ unary = None
+ if args.unary_matrix:
+ anchor_state = torch.load(
+ args.unary_matrix, map_location="cpu", weights_only=False
+ )
+ combined = anchor_state["combined"].double()
+ subset_generator = torch.Generator().manual_seed(args.subset_seed)
+ subset = torch.randperm(len(combined), generator=subset_generator)[
+ : args.samples
+ ]
+ unary_subset = combined[subset][:, subset]
+ unary_subset = (unary_subset - unary_subset.mean()) / unary_subset.std()
+ unary = unary_subset[:, hidden].float().to(device)
+
+ text_standardized = standardize_relation(
+ F.normalize(text_input, dim=-1) @ F.normalize(text_input, dim=-1).T
+ )[0]
+ visual_standardized = standardize_relation(
+ F.normalize(visual_states, dim=-1) @ F.normalize(visual_states, dim=-1).T
+ )[0]
+ true_energy = float(
+ permutation_energy_batch(
+ text_standardized, visual_standardized, truth[None].to(device)
+ )[0]
+ )
+ report: dict = {
+ "protocol": (
+ "The text side enters through a hidden shuffle; the optimizer "
+ "never sees pair, order, or truth information. Hidden truth "
+ "scores outcomes only."
+ ),
+ "dataset": args.dataset,
+ "dims": args.dims,
+ "samples": size,
+ "true_energy": true_energy,
+ "chance_accuracy": 1.0 / size,
+ "arms": {},
+ }
+ arms = [arm.strip() for arm in args.arms.split(",")]
+ if "tempering" in arms:
+ report["arms"]["tempering"] = arm_tempering(
+ text_standardized,
+ visual_standardized,
+ truth,
+ true_energy,
+ args,
+ generator,
+ unary=unary,
+ )
+ if unary is not None:
+ report["unary"] = {
+ "weight": args.unary_weight,
+ "matched_mean": float(
+ unary[torch.arange(size, device=unary.device), truth.to(unary.device)].mean()
+ ),
+ "grand_mean": float(unary.mean()),
+ "init": args.init,
+ }
+ if "sinkhorn" in arms:
+ report["arms"]["sinkhorn"] = arm_sinkhorn(
+ text_input, visual_standardized, truth, true_energy, args, generator
+ )
+ if "homotopy" in arms:
+ bands: list[int | None] = [
+ None if band == "full" else int(band)
+ for band in args.homotopy_bands.split(",")
+ ]
+ report["arms"]["homotopy"] = arm_sinkhorn(
+ text_input,
+ visual_standardized,
+ truth,
+ true_energy,
+ args,
+ generator,
+ bands=bands,
+ )
+ summary = {}
+ for arm, result in report["arms"].items():
+ candidates = result.get("restarts") or result.get("replicas") or []
+ if "best" in result:
+ candidates = candidates + [result["best"]]
+ best = max(candidates, key=lambda c: c["accuracy"], default=None)
+ summary[arm] = best
+ report["summary"] = summary
+ print(json.dumps({"summary": summary}))
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/common.py b/worldalign/common.py
new file mode 100644
index 0000000..aa7b174
--- /dev/null
+++ b/worldalign/common.py
@@ -0,0 +1,142 @@
+from __future__ import annotations
+
+import json
+import math
+import os
+import random
+from pathlib import Path
+from typing import Any
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+
+DATASET_NAME = "nlphuji/flickr30k"
+DATASET_SPLIT = "test"
+
+
+def seed_everything(seed: int) -> None:
+ random.seed(seed)
+ np.random.seed(seed)
+ torch.manual_seed(seed)
+ if torch.cuda.is_available():
+ torch.cuda.manual_seed_all(seed)
+
+
+def read_json(path: str | os.PathLike[str]) -> dict[str, Any]:
+ with open(path, encoding="utf-8") as f:
+ return json.load(f)
+
+
+def write_json(path: str | os.PathLike[str], value: Any) -> None:
+ path = Path(path)
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with open(path, "w", encoding="utf-8") as f:
+ json.dump(value, f, indent=2, ensure_ascii=False)
+
+
+def normalized(x: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
+ return F.normalize(x.float(), dim=-1, eps=eps)
+
+
+def pairwise_cosine_distance(x: np.ndarray) -> np.ndarray:
+ x = x.astype(np.float64, copy=False)
+ x /= np.linalg.norm(x, axis=1, keepdims=True).clip(min=1e-12)
+ d = 1.0 - x @ x.T
+ np.fill_diagonal(d, 0.0)
+ scale = np.median(d[d > 0])
+ return d / max(float(scale), 1e-12)
+
+
+def linear_cka(x: torch.Tensor, y: torch.Tensor) -> float:
+ x = x.float() - x.float().mean(0, keepdim=True)
+ y = y.float() - y.float().mean(0, keepdim=True)
+ xty = x.T @ y
+ numerator = (xty * xty).sum()
+ xx = x.T @ x
+ yy = y.T @ y
+ denominator = torch.sqrt((xx * xx).sum() * (yy * yy).sum())
+ return float((numerator / denominator.clamp_min(1e-12)).item())
+
+
+def retrieval_metrics(
+ image_features: torch.Tensor,
+ text_features: torch.Tensor,
+ ks: tuple[int, ...] = (1, 5, 10),
+) -> dict[str, float]:
+ image_features = normalized(image_features)
+ text_features = normalized(text_features)
+ similarities = image_features @ text_features.T
+ n = similarities.shape[0]
+ truth = torch.arange(n, device=similarities.device)
+
+ i2t_order = similarities.argsort(dim=1, descending=True)
+ t2i_order = similarities.T.argsort(dim=1, descending=True)
+ i2t_rank = (i2t_order == truth[:, None]).nonzero()[:, 1]
+ t2i_rank = (t2i_order == truth[:, None]).nonzero()[:, 1]
+
+ result: dict[str, float] = {}
+ for k in ks:
+ result[f"i2t_r@{k}"] = float((i2t_rank < k).float().mean().item())
+ result[f"t2i_r@{k}"] = float((t2i_rank < k).float().mean().item())
+ result["i2t_median_rank"] = float(i2t_rank.float().median().item() + 1)
+ result["t2i_median_rank"] = float(t2i_rank.float().median().item() + 1)
+ result["chance_r@1"] = 1.0 / max(n, 1)
+ return result
+
+
+def batch_indices(n: int, batch_size: int, shuffle: bool = False, seed: int = 0):
+ order = np.arange(n)
+ if shuffle:
+ rng = np.random.default_rng(seed)
+ rng.shuffle(order)
+ for start in range(0, n, batch_size):
+ yield order[start : start + batch_size]
+
+
+def sliced_wasserstein(
+ x: torch.Tensor,
+ y: torch.Tensor,
+ num_projections: int = 64,
+) -> torch.Tensor:
+ """Differentiable empirical sliced W2 for equal-sized minibatches."""
+ n = min(x.shape[0], y.shape[0])
+ x = x[:n]
+ y = y[:n]
+ directions = torch.randn(
+ x.shape[-1], num_projections, device=x.device, dtype=x.dtype
+ )
+ directions = F.normalize(directions, dim=0)
+ x_proj = (x @ directions).sort(dim=0).values
+ y_proj = (y @ directions).sort(dim=0).values
+ return (x_proj - y_proj).square().mean()
+
+
+def cosine_isometry_loss(source: torch.Tensor, mapped: torch.Tensor) -> torch.Tensor:
+ source = normalized(source)
+ mapped = normalized(mapped)
+ source_gram = source @ source.T
+ mapped_gram = mapped @ mapped.T
+ mask = ~torch.eye(source.shape[0], dtype=torch.bool, device=source.device)
+ return (source_gram[mask] - mapped_gram[mask]).square().mean()
+
+
+def cosine_loss(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
+ return 1.0 - F.cosine_similarity(x.float(), y.float(), dim=-1).mean()
+
+
+def dtype_for_device(device: str) -> torch.dtype:
+ return torch.bfloat16 if device.startswith("cuda") else torch.float32
+
+
+def parameter_count(module: torch.nn.Module) -> int:
+ return sum(p.numel() for p in module.parameters())
+
+
+def cosine_schedule(step: int, steps: int, warmup: int) -> float:
+ if step < warmup:
+ return (step + 1) / max(1, warmup)
+ progress = (step - warmup) / max(1, steps - warmup)
+ return 0.5 * (1.0 + math.cos(math.pi * progress))
+
diff --git a/worldalign/content_gate.py b/worldalign/content_gate.py
new file mode 100644
index 0000000..4e1383a
--- /dev/null
+++ b/worldalign/content_gate.py
@@ -0,0 +1,164 @@
+"""Full assignment gate on content-projected states.
+
+The R6 battery collapsed the improving-swap fraction by one to two orders
+of magnitude. This runs the complete gate -- global ranking, exact
+transposition enumeration, descent from the truth, and counterfeit search
+from random starts -- on the projected states, which the fraction alone
+cannot decide.
+
+Projection hygiene: Flickr directions are fitted on the text-only
+training orbits and applied to held-out test states. VG directions are
+fitted only on nodes outside the evaluated subset. No pairs anywhere in
+the fit; hidden pairs score orderings only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+
+from .common import read_json, seed_everything, write_json
+from .content_projection import content_directions, project
+from .io import load_feature_pair, select_rows
+from .manifold_gate import standardize_relation
+from .ricci_control import run_gates
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--dataset", choices=["flickr", "vg"], default="flickr")
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt")
+ parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt")
+ parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt")
+ parser.add_argument(
+ "--vg-ground-truth", default="artifacts/vg_5k/ground_truth.private.jsonl"
+ )
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--subset-seed", type=int, default=0)
+ parser.add_argument("--dims", type=int, default=32)
+ parser.add_argument("--shrinkage", type=float, default=0.05)
+ parser.add_argument("--random-perms", type=int, default=1000)
+ parser.add_argument("--descent-restarts", type=int, default=3)
+ 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("--seed", type=int, default=20260729)
+ parser.add_argument("--output", required=True)
+ return parser.parse_args()
+
+
+def flickr_states(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]:
+ manifest = read_json(args.manifest)
+ vision, _, vision_lookup, _ = load_feature_pair(args.vision, args.text)
+ orbits = torch.load(args.text_orbits, map_location="cpu", weights_only=False)
+ lookup = {int(row): i for i, row in enumerate(orbits["rows"])}
+ features = F.normalize(orbits["features"].double(), dim=-1)
+ train_views = features[[lookup[int(r)] for r in manifest["text_only_train"]]]
+ mean, directions = content_directions(train_views, args.shrinkage)
+ rows = manifest[args.split][: args.samples]
+ text_states = project(
+ features[[lookup[int(r)] for r in rows]].mean(1), mean, directions, args.dims
+ )
+ visual_states = F.normalize(
+ select_rows(vision["features"], vision_lookup, rows).double(), dim=-1
+ )
+ return text_states, visual_states
+
+
+def vg_states(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]:
+ vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False)
+ text = torch.load(args.vg_text, map_location="cpu", weights_only=False)
+ vision_key = "context_states" if "context_states" in vision else "region_features"
+ text_key = "context_states" if "context_states" in text else "region_features"
+ pairs = [
+ json.loads(line)
+ for line in open(args.vg_ground_truth, encoding="utf-8")
+ if line.strip()
+ ]
+ vision_index = {node: i for i, node in enumerate(vision["node_ids"])}
+ text_index = {node: i for i, node in enumerate(text["node_ids"])}
+ vision_order = [vision_index[p["vision_node_id"]] for p in pairs]
+ text_order = [text_index[p["text_node_id"]] for p in pairs]
+ visual_views = F.normalize(vision[vision_key].double(), dim=-1)[vision_order]
+ text_views = F.normalize(text[text_key].double(), dim=-1)[text_order]
+ generator = torch.Generator().manual_seed(args.subset_seed)
+ order = torch.randperm(len(visual_views), generator=generator)
+ subset = order[: args.samples]
+ holdout = order[args.samples :]
+ visual_mean, visual_directions = content_directions(
+ visual_views[holdout], args.shrinkage
+ )
+ text_mean, text_directions = content_directions(
+ text_views[holdout], args.shrinkage
+ )
+ text_states = project(
+ text_views[subset].mean(1), text_mean, text_directions, args.dims
+ )
+ visual_states = project(
+ visual_views[subset].mean(1), visual_mean, visual_directions, args.dims
+ )
+ return text_states, visual_states
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ if args.dataset == "flickr":
+ text_states, visual_states = flickr_states(args)
+ else:
+ text_states, visual_states = vg_states(args)
+ text_channels = standardize_relation(text_states @ text_states.T)[0][None]
+ visual_channels = standardize_relation(visual_states @ visual_states.T)[0][None]
+ generator = torch.Generator().manual_seed(args.seed)
+ report = {
+ "protocol": (
+ "Content-projected states, directions fitted without pairs on "
+ "held-out scenes; the complete assignment gate is scored with "
+ "hidden pairs."
+ ),
+ "dataset": args.dataset,
+ "dims": args.dims,
+ "samples": args.samples,
+ **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"],
+ "identity_strict_local_min": report["gate_b"]["identity_is_local_min_mse"],
+ "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,
+ ),
+ "best_random_accuracy": max(
+ (r["final_accuracy"] 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"]
+ and verdict["best_random_accuracy"] < 0.5
+ )
+ report["verdict"] = verdict
+ print(json.dumps({"verdict": verdict}))
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/content_projection.py b/worldalign/content_projection.py
new file mode 100644
index 0000000..b3e7422
--- /dev/null
+++ b/worldalign/content_projection.py
@@ -0,0 +1,214 @@
+"""R6 battery: cross-view predictable (content) subspace projection.
+
+Views of the same scene share content and differ in style. Directions that
+maximize between-scene over within-scene variance are estimated from
+within-modality orbit structure alone (no pairs anywhere), then states are
+projected onto the top content directions. Hidden pairs are used only to
+score the cross-modal effect.
+
+Note the contrast with population whitening, which was destructive: the
+generalized eigenproblem whitens the within-scene (view-noise) covariance,
+not the total covariance, so directions where redescriptions of the same
+scene agree are amplified rather than equalized away.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from scipy.linalg import eigh
+
+from .common import read_json, write_json
+from .io import load_feature_pair, select_rows
+from .manifold_gate import all_transposition_delta_mse, standardize_relation
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt")
+ parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt")
+ parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt")
+ parser.add_argument(
+ "--vg-ground-truth", default="artifacts/vg_5k/ground_truth.private.jsonl"
+ )
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--dims", default="8,16,32,64,128,256")
+ parser.add_argument("--shrinkage", type=float, default=0.05)
+ parser.add_argument(
+ "--output", default="artifacts/manifold_gate/content_projection.json"
+ )
+ return parser.parse_args()
+
+
+def content_directions(
+ views: torch.Tensor, shrinkage: float
+) -> tuple[np.ndarray, np.ndarray]:
+ """Generalized eigenvectors of between-scene vs within-scene covariance.
+
+ views: [scenes, views_per_scene, dim], any consistent preprocessing.
+ Returns (mean, directions) with directions sorted by decreasing ratio.
+ """
+ flat = views.reshape(-1, views.shape[-1]).double().numpy()
+ mean = flat.mean(0)
+ scene_means = views.double().mean(1).numpy()
+ within = views.double().numpy() - scene_means[:, None, :]
+ within = within.reshape(-1, views.shape[-1])
+ sigma_within = within.T @ within / max(len(within) - 1, 1)
+ centered_means = scene_means - mean
+ sigma_between = (
+ centered_means.T @ centered_means / max(len(centered_means) - 1, 1)
+ )
+ trace_scale = np.trace(sigma_within) / len(sigma_within)
+ regularized = sigma_within + shrinkage * trace_scale * np.eye(len(sigma_within))
+ values, vectors = eigh(sigma_between, regularized)
+ order = np.argsort(values)[::-1]
+ return mean, vectors[:, order]
+
+
+def project(
+ states: torch.Tensor, mean: np.ndarray, directions: np.ndarray, dims: int
+) -> torch.Tensor:
+ basis = torch.from_numpy(directions[:, :dims]).double()
+ centered = states.double() - torch.from_numpy(mean).double()
+ return F.normalize(centered @ basis, dim=-1)
+
+
+def relation_spearman(text_states: torch.Tensor, visual_states: torch.Tensor) -> float:
+ text_relation = text_states @ text_states.T
+ visual_relation = visual_states @ visual_states.T
+ mask = ~torch.eye(len(text_relation), dtype=torch.bool)
+ t, v = text_relation[mask], visual_relation[mask]
+ ranks = torch.stack(
+ [t.argsort().argsort().double(), v.argsort().argsort().double()]
+ )
+ return float(torch.corrcoef(ranks)[0, 1])
+
+
+def improving_fraction(
+ text_states: torch.Tensor, visual_states: torch.Tensor
+) -> float:
+ text_field, _, _ = standardize_relation(text_states @ text_states.T)
+ visual_field, _, _ = standardize_relation(visual_states @ visual_states.T)
+ delta = all_transposition_delta_mse(text_field, visual_field)
+ upper = torch.triu(torch.ones_like(delta, dtype=torch.bool), diagonal=1)
+ return float((delta[upper] < 0).double().mean())
+
+
+def flickr_battery(args: argparse.Namespace, dims: list[int]) -> dict:
+ manifest = read_json(args.manifest)
+ vision, _, vision_lookup, _ = load_feature_pair(args.vision, args.text)
+ orbits = torch.load(args.text_orbits, map_location="cpu", weights_only=False)
+ lookup = {int(row): i for i, row in enumerate(orbits["rows"])}
+ features = F.normalize(orbits["features"].double(), dim=-1)
+
+ train_rows = [int(r) for r in manifest["text_only_train"]]
+ test_rows = manifest["test"][: args.samples]
+ train_views = features[[lookup[r] for r in train_rows]]
+ mean, directions = content_directions(train_views, args.shrinkage)
+
+ visual = select_rows(vision["features"], vision_lookup, test_rows).double()
+ visual = F.normalize(visual, dim=-1)
+ test_views = features[[lookup[int(r)] for r in test_rows]]
+ orbit_mean = F.normalize(test_views.mean(1), dim=-1)
+ single = test_views[:, 0]
+
+ report = {
+ "baseline_orbit_mean": {
+ "spearman": relation_spearman(orbit_mean, visual),
+ "improving_fraction": improving_fraction(orbit_mean, visual),
+ },
+ "baseline_single": {
+ "spearman": relation_spearman(F.normalize(single, dim=-1), visual),
+ "improving_fraction": improving_fraction(
+ F.normalize(single, dim=-1), visual
+ ),
+ },
+ "projected": {},
+ }
+ for k in dims:
+ projected_mean = project(test_views.mean(1), mean, directions, k)
+ projected_single = project(single, mean, directions, k)
+ report["projected"][k] = {
+ "orbit_mean_spearman": relation_spearman(projected_mean, visual),
+ "orbit_mean_improving_fraction": improving_fraction(
+ projected_mean, visual
+ ),
+ "single_spearman": relation_spearman(projected_single, visual),
+ }
+ return report
+
+
+def vg_battery(args: argparse.Namespace, dims: list[int]) -> dict:
+ vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False)
+ text = torch.load(args.vg_text, map_location="cpu", weights_only=False)
+ pairs = [
+ json.loads(line)
+ for line in open(args.vg_ground_truth, encoding="utf-8")
+ if line.strip()
+ ]
+ vision_index = {node: i for i, node in enumerate(vision["node_ids"])}
+ text_index = {node: i for i, node in enumerate(text["node_ids"])}
+ vision_order = [vision_index[p["vision_node_id"]] for p in pairs]
+ text_order = [text_index[p["text_node_id"]] for p in pairs]
+ visual_views = F.normalize(vision["region_features"].double(), dim=-1)[
+ vision_order
+ ]
+ text_views = F.normalize(text["region_features"].double(), dim=-1)[text_order]
+
+ visual_mean, visual_directions = content_directions(visual_views, args.shrinkage)
+ text_mean, text_directions = content_directions(text_views, args.shrinkage)
+
+ generator = torch.Generator().manual_seed(0)
+ subset = torch.randperm(len(visual_views), generator=generator)[: args.samples]
+ visual_node = F.normalize(visual_views[subset].mean(1), dim=-1)
+ text_node = F.normalize(text_views[subset].mean(1), dim=-1)
+
+ report = {
+ "baseline": {
+ "spearman": relation_spearman(text_node, visual_node),
+ "improving_fraction": improving_fraction(text_node, visual_node),
+ },
+ "projected": {},
+ }
+ for k in dims:
+ projected_text = project(
+ text_views[subset].mean(1), text_mean, text_directions, k
+ )
+ projected_visual = project(
+ visual_views[subset].mean(1), visual_mean, visual_directions, k
+ )
+ report["projected"][k] = {
+ "both_sides_spearman": relation_spearman(projected_text, projected_visual),
+ "both_sides_improving_fraction": improving_fraction(
+ projected_text, projected_visual
+ ),
+ "text_only_spearman": relation_spearman(projected_text, visual_node),
+ }
+ return report
+
+
+def main() -> None:
+ args = parse_args()
+ dims = [int(d) for d in args.dims.split(",")]
+ report = {
+ "protocol": (
+ "Content directions maximize between-scene over within-scene "
+ "variance of view states, fitted per modality on unpaired "
+ "orbit structure only. Hidden pairs score the effect."
+ ),
+ "flickr": flickr_battery(args, dims),
+ "vg_region_closed": vg_battery(args, dims),
+ }
+ write_json(args.output, report)
+ print(json.dumps(report, indent=2))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/cooccurrence_dictionary.py b/worldalign/cooccurrence_dictionary.py
new file mode 100644
index 0000000..4c338e4
--- /dev/null
+++ b/worldalign/cooccurrence_dictionary.py
@@ -0,0 +1,237 @@
+"""Derive the cross-modal value correspondence from joint structure.
+
+Marginal frequency pairs values only when the two corpora rank them the
+same way, which reporting bias erodes: text mentions the salient, pixels
+count the common. Joint structure is sturdier. Which colours co-occur
+with which shapes is a property of the world that both modalities
+observe, and reporting bias distorts the marginals long before it
+scrambles that pattern.
+
+The correspondence is therefore recovered by matching two small
+value-level graphs -- colour-by-shape co-occurrence, estimated per
+modality on its own corpus -- with the same spectral-plus-refinement
+solver used at scene level. Both corpora stay disjoint; no pair is read.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from collections import Counter
+from pathlib import Path
+
+import numpy as np
+import torch
+from scipy.cluster.hierarchy import fcluster, linkage
+from scipy.optimize import linear_sum_assignment
+from sklearn.cluster import KMeans
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .synth_set_battery import parse_group_phrases
+from .synth_towers import load_image
+from .tier0_pipeline import (
+ appearance,
+ extract_objects,
+ group_objects,
+)
+from .tier0_dictionary import partition_text_vocabulary, text_token_statistics
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v3")
+ parser.add_argument("--fit-scenes", type=int, default=3000)
+ parser.add_argument("--peak-distance", type=int, default=5)
+ parser.add_argument("--group-threshold", type=float, default=0.35)
+ parser.add_argument("--shape-classes", type=int, default=6)
+ parser.add_argument("--restarts", type=int, default=40)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v3/cooccurrence_dictionary.json"
+ )
+ return parser.parse_args()
+
+
+def text_joint(
+ captions: list[list[str]],
+ rows: list[int],
+ colour_words: list[str],
+ shape_words: list[str],
+) -> np.ndarray:
+ joint = np.zeros((len(colour_words), len(shape_words)))
+ colour_index = {word: i for i, word in enumerate(colour_words)}
+ shape_index = {word: i for i, word in enumerate(shape_words)}
+ for row in rows:
+ for phrase in parse_group_phrases(captions[row][0]):
+ tokens = phrase.split()
+ colour = next((colour_index[t] for t in tokens if t in colour_index), None)
+ shape = next(
+ (
+ shape_index[t.rstrip("es") if t.rstrip("es") in shape_index else t]
+ for t in tokens
+ if t in shape_index or t.rstrip("es") in shape_index
+ ),
+ None,
+ )
+ if colour is not None and shape is not None:
+ joint[colour, shape] += 1.0
+ return joint
+
+
+def vision_joint(
+ rows: list[int],
+ image_dir: Path,
+ args: argparse.Namespace,
+ colour_classes: int,
+) -> tuple[np.ndarray, KMeans, KMeans, dict[int, int]]:
+ groups, shapes = [], []
+ for row in tqdm(rows, desc="vision joint"):
+ objects = extract_objects(
+ load_image(image_dir / f"scene{row:06d}_v0.png"), args.peak_distance
+ )
+ members = group_objects(objects, args.group_threshold)
+ if not members:
+ continue
+ features = np.stack([appearance(item) for item in objects])
+ if len(objects) == 1:
+ labels = np.array([0])
+ else:
+ labels = fcluster(
+ linkage(features, "complete"), args.group_threshold, "distance"
+ )
+ buckets: dict[int, list[dict]] = {}
+ for item, label in zip(objects, labels):
+ buckets.setdefault(int(label), []).append(item)
+ for group, bucket in zip(members, buckets.values()):
+ groups.append(group)
+ shapes.append(np.mean([item["shape"] for item in bucket], axis=0))
+ colour_model = KMeans(colour_classes, n_init=10, random_state=args.seed).fit(
+ np.stack([group["rgb"] for group in groups]).astype(np.float64)
+ )
+ shape_model = KMeans(args.shape_classes, n_init=10, random_state=args.seed).fit(
+ np.stack(shapes).astype(np.float64)
+ )
+ colour_labels = colour_model.labels_
+ shape_labels = shape_model.labels_
+ frequency = Counter(colour_labels.tolist())
+ rank = {label: index for index, (label, _) in enumerate(frequency.most_common())}
+ joint = np.zeros((colour_classes, args.shape_classes))
+ for colour, shape in zip(colour_labels, shape_labels):
+ joint[rank[int(colour)], int(shape)] += 1.0
+ return joint, colour_model, shape_model, rank
+
+
+def normalise_joint(joint: np.ndarray) -> np.ndarray:
+ """Row-normalised joint: the shape profile of each colour."""
+ return joint / joint.sum(axis=1, keepdims=True).clip(1e-9)
+
+
+def match_values(
+ text: np.ndarray, vision: np.ndarray, restarts: int, seed: int
+) -> tuple[np.ndarray, np.ndarray, float]:
+ """Align colour rows and shape columns of two joint tables.
+
+ Alternating assignment: given a column correspondence, rows are
+ matched by profile similarity; given rows, columns are rematched.
+ Iterating from many random column starts and keeping the best
+ agreement avoids the trivial fixed point. Row indices are returned as
+ vision-colour to text-colour.
+ """
+ generator = np.random.default_rng(seed)
+ text_profile = normalise_joint(text)
+ vision_profile = normalise_joint(vision)
+ shapes = text.shape[1]
+ best_row, best_column, best_value = None, None, -np.inf
+ for restart in range(restarts):
+ column = (
+ np.arange(shapes) if restart == 0 else generator.permutation(shapes)
+ )
+ row = None
+ for _ in range(20):
+ # Rows: vision colour i against text colour j under current columns.
+ score = vision_profile @ text_profile[:, column].T
+ _, row = linear_sum_assignment(-score)
+ # Columns: vision shape a against text shape b under current rows.
+ column_score = vision_profile.T @ text_profile[row]
+ _, new_column = linear_sum_assignment(-column_score)
+ if np.array_equal(new_column, column):
+ break
+ column = new_column
+ value = float((vision_profile * text_profile[row][:, column]).sum())
+ if value > best_value:
+ best_row, best_column, best_value = row, column, value
+ return best_row, best_column, best_value
+
+
+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"]
+ scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"]
+ image_dir = Path(manifest["image_dir"])
+
+ families = partition_text_vocabulary(
+ text_token_statistics(captions, manifest["text_only_train"])
+ )
+ colour_words = families["colour_words"]
+ from .synth_world import SHAPES
+
+ shape_words = list(SHAPES)
+ text_table = text_joint(
+ captions, manifest["text_only_train"], colour_words, shape_words
+ )
+ vision_table, colour_model, _, rank = vision_joint(
+ manifest["vision_only_train"][: args.fit_scenes],
+ image_dir,
+ args,
+ len(colour_words),
+ )
+
+ row, column, agreement = match_values(
+ text_table, vision_table, args.restarts, args.seed
+ )
+
+ # Evaluation only: names attached to vision clusters by nearest true RGB.
+ from .synth_world import COLORS
+
+ names = list(COLORS)
+ reference = np.asarray([COLORS[name] for name in names], dtype=np.float64) / 255.0
+ inverse_rank = {index: label for label, index in rank.items()}
+ cluster_name = {
+ index: names[
+ int(np.argmin(((colour_model.cluster_centers_[inverse_rank[index]] - reference) ** 2).sum(1)))
+ ]
+ for index in range(len(colour_words))
+ }
+ marginal_correct = sum(
+ 1
+ for index, word in enumerate(colour_words)
+ if index < len(cluster_name) and word == cluster_name[index]
+ )
+ joint_correct = sum(
+ 1
+ for vision_index, text_index in enumerate(row)
+ if colour_words[text_index] == cluster_name[vision_index]
+ )
+ report = {
+ "protocol": (
+ "Colour-by-shape joint tables are estimated per modality on "
+ "disjoint corpora and aligned by alternating assignment; "
+ "truth is read only to count correct entries."
+ ),
+ "colour_words": colour_words,
+ "agreement": agreement,
+ "marginal_rank_correct": marginal_correct,
+ "joint_structure_correct": joint_correct,
+ "classes": len(colour_words),
+ }
+ print(json.dumps({k: report[k] for k in
+ ("marginal_rank_correct", "joint_structure_correct", "classes")}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/diagnose.py b/worldalign/diagnose.py
new file mode 100644
index 0000000..2d05fd9
--- /dev/null
+++ b/worldalign/diagnose.py
@@ -0,0 +1,76 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import numpy as np
+import torch
+from scipy.stats import spearmanr
+
+from .common import linear_cka, normalized, read_json, write_json
+from .io import load_feature_pair, select_rows
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--split", choices=["val", "test"], default="val")
+ p.add_argument("--max-samples", type=int, default=1_000)
+ p.add_argument("--permutations", type=int, default=20)
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument("--output", default="artifacts/diagnostics.json")
+ return p.parse_args()
+
+
+def upper_triangle(x: torch.Tensor) -> np.ndarray:
+ n = x.shape[0]
+ i, j = torch.triu_indices(n, n, offset=1)
+ return x[i, j].cpu().numpy()
+
+
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
+ rows = manifest[args.split][: args.max_samples]
+ x = select_rows(vision["features"], vlookup, rows)
+ y = select_rows(text["features"], tlookup, rows)
+
+ gx = normalized(x) @ normalized(x).T
+ gy = normalized(y) @ normalized(y).T
+ gx_upper = upper_triangle(gx)
+ gy_upper = upper_triangle(gy)
+ rho = spearmanr(gx_upper, gy_upper).statistic
+ generator = torch.Generator().manual_seed(args.seed)
+ shuffled_rhos = []
+ for _ in range(args.permutations):
+ permutation = torch.randperm(len(y), generator=generator)
+ shuffled_gy = gy[permutation][:, permutation]
+ shuffled_rhos.append(
+ float(spearmanr(gx_upper, upper_triangle(shuffled_gy)).statistic)
+ )
+ shuffled_mean = float(np.mean(shuffled_rhos))
+ shuffled_std = float(np.std(shuffled_rhos))
+ result = {
+ "split": args.split,
+ "samples": len(rows),
+ "linear_cka": linear_cka(x, y),
+ "pairwise_cosine_spearman": float(rho),
+ "shuffled_spearman_mean": shuffled_mean,
+ "shuffled_spearman_std": shuffled_std,
+ "spearman_shuffle_z": float(
+ (rho - shuffled_mean) / max(shuffled_std, 1e-12)
+ ),
+ "permutations": args.permutations,
+ "vision_dim": x.shape[-1],
+ "text_dim": y.shape[-1],
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(result)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/energy.py b/worldalign/energy.py
new file mode 100644
index 0000000..796c41a
--- /dev/null
+++ b/worldalign/energy.py
@@ -0,0 +1,137 @@
+from __future__ import annotations
+
+import torch
+import torch.nn.functional as F
+
+
+def off_diagonal_mask(size: int, device: torch.device | str) -> torch.Tensor:
+ return ~torch.eye(size, dtype=torch.bool, device=device)
+
+
+def standardized_relation(
+ features: torch.Tensor, mask: torch.Tensor | None = None
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Cosine relation field and its standardized off-diagonal values."""
+ features = F.normalize(features.float(), dim=-1)
+ relation = features @ features.T
+ if mask is None:
+ mask = off_diagonal_mask(len(features), features.device)
+ values = relation[mask]
+ standardized = (values - values.mean()) / values.std().clamp_min(1e-6)
+ return relation, standardized
+
+
+def relation_field_energy(
+ visual_relation: torch.Tensor,
+ visual_standardized: torch.Tensor,
+ language_particles: torch.Tensor,
+ temperatures: tuple[float, ...] = (0.03, 0.07, 0.15),
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Second-order and multiscale conditional relation energies."""
+ language_relation, language_standardized = standardized_relation(
+ language_particles
+ )
+ mse = F.mse_loss(language_standardized, visual_standardized)
+ diagonal = torch.eye(
+ len(language_particles),
+ dtype=torch.bool,
+ device=language_particles.device,
+ )
+ conditional_kl = language_particles.new_zeros(())
+ for temperature in temperatures:
+ visual_logits = (visual_relation / temperature).masked_fill(
+ diagonal, -1e4
+ )
+ language_logits = (language_relation / temperature).masked_fill(
+ diagonal, -1e4
+ )
+ visual_probability = F.softmax(visual_logits, dim=-1)
+ conditional_kl = conditional_kl + (
+ visual_probability
+ * (
+ F.log_softmax(visual_logits, dim=-1)
+ - F.log_softmax(language_logits, dim=-1)
+ )
+ ).sum(-1).mean()
+ return mse, conditional_kl
+
+
+def projection_quantile_target(
+ text_features: torch.Tensor,
+ particles: int,
+ projections: int,
+ generator: torch.Generator,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Fixed sliced-distribution target from an unpaired text population."""
+ directions = F.normalize(
+ torch.randn(
+ text_features.shape[-1],
+ projections,
+ generator=generator,
+ device=text_features.device,
+ ),
+ dim=0,
+ )
+ projected = (text_features @ directions).sort(dim=0).values
+ quantile_indices = (
+ torch.linspace(
+ 0,
+ len(projected) - 1,
+ particles,
+ device=text_features.device,
+ )
+ .round()
+ .long()
+ )
+ return directions, projected[quantile_indices]
+
+
+def sliced_distribution_energy(
+ particles: torch.Tensor,
+ directions: torch.Tensor,
+ target_quantiles: torch.Tensor,
+) -> torch.Tensor:
+ projected = (F.normalize(particles, dim=-1) @ directions).sort(
+ dim=0
+ ).values
+ return F.mse_loss(projected, target_quantiles)
+
+
+def prototype_manifold_energy(
+ particles: torch.Tensor, prototypes: torch.Tensor
+) -> torch.Tensor:
+ particles = F.normalize(particles, dim=-1)
+ prototypes = F.normalize(prototypes, dim=-1)
+ return (1 - (particles @ prototypes.T).max(dim=-1).values).mean()
+
+
+def log_sinkhorn(
+ logits: torch.Tensor, temperature: float, iterations: int = 12
+) -> torch.Tensor:
+ """Doubly stochastic coupling with differentiable log-domain updates."""
+ log_coupling = logits / temperature
+ for _ in range(iterations):
+ log_coupling = log_coupling - torch.logsumexp(
+ log_coupling, dim=1, keepdim=True
+ )
+ log_coupling = log_coupling - torch.logsumexp(
+ log_coupling, dim=0, keepdim=True
+ )
+ return log_coupling.exp()
+
+
+def retrieval_metrics(
+ particles: torch.Tensor, paired_text: torch.Tensor
+) -> dict[str, float]:
+ particles = F.normalize(particles.float(), dim=-1)
+ paired_text = F.normalize(paired_text.float(), dim=-1)
+ similarity = particles @ paired_text.T
+ target = similarity.diagonal()
+ ranks = (similarity > target[:, None]).sum(-1) + 1
+ return {
+ "r@1": float((ranks <= 1).float().mean()),
+ "r@5": float((ranks <= 5).float().mean()),
+ "r@10": float((ranks <= 10).float().mean()),
+ "median_rank": float(ranks.float().median()),
+ "paired_cosine": float(target.mean()),
+ }
diff --git a/worldalign/energy_coupling.py b/worldalign/energy_coupling.py
new file mode 100644
index 0000000..cf5b859
--- /dev/null
+++ b/worldalign/energy_coupling.py
@@ -0,0 +1,205 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+
+from .common import read_json, seed_everything, write_json
+from .energy import (
+ log_sinkhorn,
+ relation_field_energy,
+ retrieval_metrics,
+ standardized_relation,
+)
+from .io import load_feature_pair, select_rows
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument(
+ "--reservoir",
+ choices=["same_set_shuffled", "unpaired"],
+ default="same_set_shuffled",
+ )
+ parser.add_argument("--steps", type=int, default=1_200)
+ parser.add_argument("--lr", type=float, default=0.12)
+ parser.add_argument("--temperature-start", type=float, default=0.8)
+ parser.add_argument("--temperature-end", type=float, default=0.12)
+ parser.add_argument("--conditional-weight", type=float, default=0.08)
+ parser.add_argument("--entropy-weight-start", type=float, default=0.005)
+ parser.add_argument("--entropy-weight-end", type=float, default=0.055)
+ parser.add_argument("--device", default="cuda:1")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument(
+ "--output", default="artifacts/energy_coupling_test.pt"
+ )
+ parser.add_argument(
+ "--metrics-output", default="artifacts/energy_coupling_test.json"
+ )
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ vision, text, vision_lookup, text_lookup = load_feature_pair(
+ args.vision, args.text
+ )
+ rows = manifest[args.split][: args.samples]
+ visual = select_rows(
+ vision["features"], vision_lookup, rows
+ ).to(args.device)
+ paired_text = select_rows(
+ text["features"], text_lookup, rows
+ ).to(args.device)
+ text_population = select_rows(
+ text["features"], text_lookup, manifest["text_only_train"]
+ ).to(args.device)
+ if args.text_orbits:
+ state = torch.load(
+ args.text_orbits, map_location="cpu", weights_only=False
+ )
+ lookup = {
+ int(row): index for index, row in enumerate(state["rows"])
+ }
+ orbit_mean = F.normalize(
+ state["features"].float().mean(1), dim=-1
+ )
+ paired_text = select_rows(
+ orbit_mean, lookup, rows
+ ).to(args.device)
+ text_population = select_rows(
+ orbit_mean, lookup, manifest["text_only_train"]
+ ).to(args.device)
+
+ generator = torch.Generator(device=args.device).manual_seed(args.seed)
+ target: torch.Tensor | None = None
+ if args.reservoir == "same_set_shuffled":
+ permutation = torch.randperm(
+ len(paired_text), generator=generator, device=args.device
+ )
+ anchors = paired_text[permutation]
+ target = torch.argsort(permutation)
+ else:
+ anchor_indices = torch.randperm(
+ len(text_population),
+ generator=generator,
+ device=args.device,
+ )[: len(visual)]
+ anchors = text_population[anchor_indices]
+
+ initial_permutation = torch.randperm(
+ len(visual), generator=generator, device=args.device
+ )
+ logits = torch.nn.Parameter(
+ 0.02
+ * torch.randn(
+ len(visual),
+ len(visual),
+ generator=generator,
+ device=args.device,
+ )
+ )
+ with torch.no_grad():
+ logits[torch.arange(len(visual), device=args.device), initial_permutation] += 3
+ visual_relation, visual_standardized = standardized_relation(visual)
+ optimizer = torch.optim.Adam([logits], lr=args.lr)
+ history: list[dict] = []
+
+ for step in range(args.steps + 1):
+ progress = min(step / max(args.steps, 1), 1.0)
+ temperature = max(
+ args.temperature_end,
+ args.temperature_start
+ + progress
+ * (args.temperature_end - args.temperature_start),
+ )
+ coupling = log_sinkhorn(logits, temperature, iterations=15)
+ particles = F.normalize(coupling @ anchors, dim=-1)
+ relation, conditional = relation_field_energy(
+ visual_relation, visual_standardized, particles
+ )
+ entropy = -(
+ coupling * coupling.clamp_min(1e-12).log()
+ ).sum(-1).mean()
+ entropy_weight = (
+ args.entropy_weight_start
+ + progress
+ * (args.entropy_weight_end - args.entropy_weight_start)
+ )
+ loss = (
+ relation
+ + args.conditional_weight * conditional
+ + entropy_weight * entropy
+ )
+ if step % 100 == 0 or step == args.steps:
+ record = {
+ "step": step,
+ "total": float(loss.detach()),
+ "relation": float(relation.detach()),
+ "conditional": float(conditional.detach()),
+ "entropy": float(entropy.detach()),
+ "temperature": temperature,
+ "paired_evaluation_only": retrieval_metrics(
+ particles.detach(), paired_text
+ ),
+ }
+ if target is not None:
+ prediction = coupling.argmax(-1)
+ record["exact_coupling_evaluation_only"] = {
+ "accuracy": float((prediction == target).float().mean()),
+ "true_mass": float(
+ coupling[
+ torch.arange(
+ len(coupling), device=args.device
+ ),
+ target,
+ ].mean()
+ ),
+ }
+ history.append(record)
+ print(json.dumps(record))
+ if step == args.steps:
+ break
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ torch.nn.utils.clip_grad_norm_([logits], 5.0)
+ optimizer.step()
+
+ result = {
+ "protocol": (
+ "The optimizer sees only frozen within-modality states and "
+ "energy. Pair identity and paired text are used only for "
+ "evaluation diagnostics."
+ ),
+ "mode": "mass_conserving_language_state_coupling",
+ "reservoir": args.reservoir,
+ "split": args.split,
+ "rows": rows,
+ "args": vars(args),
+ "history": history,
+ }
+ state = {
+ **result,
+ "anchors": anchors.detach().cpu(),
+ "final_coupling": coupling.detach().cpu(),
+ "final_particles": particles.detach().cpu(),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ write_json(args.metrics_output, result)
+ print(f"Wrote {args.output} and {args.metrics_output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/energy_functional_infer.py b/worldalign/energy_functional_infer.py
new file mode 100644
index 0000000..6864011
--- /dev/null
+++ b/worldalign/energy_functional_infer.py
@@ -0,0 +1,262 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import dtype_for_device, read_json, seed_everything, write_json
+from .energy import (
+ projection_quantile_target,
+ relation_field_energy,
+ retrieval_metrics,
+ sliced_distribution_energy,
+ standardized_relation,
+)
+from .extract_text import mean_pool
+from .io import load_feature_pair, select_rows
+from .models import load_prefix
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument(
+ "--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt"
+ )
+ parser.add_argument("--prefix", default="artifacts/prefix.pt")
+ parser.add_argument("--split", choices=["val", "test"], default="val")
+ parser.add_argument("--samples", type=int, default=128)
+ parser.add_argument("--steps", type=int, default=60)
+ parser.add_argument("--refresh-steps", type=int, default=20)
+ parser.add_argument("--max-new-tokens", type=int, default=24)
+ parser.add_argument("--lr", type=float, default=0.03)
+ parser.add_argument("--projections", type=int, default=128)
+ parser.add_argument("--relation-weight", type=float, default=1.0)
+ parser.add_argument("--conditional-weight", type=float, default=0.08)
+ parser.add_argument("--distribution-weight", type=float, default=80.0)
+ parser.add_argument("--functional-weight", type=float, default=10.0)
+ parser.add_argument("--device", default="cuda:1")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument(
+ "--output", default="artifacts/energy_functional_val.pt"
+ )
+ parser.add_argument(
+ "--metrics-output", default="artifacts/energy_functional_val.json"
+ )
+ return parser.parse_args()
+
+
+@torch.no_grad()
+def refresh_functional_anchor(
+ particles: torch.Tensor,
+ prefix: torch.nn.Module,
+ lm: AutoModelForCausalLM,
+ tokenizer: AutoTokenizer,
+ dtype: torch.dtype,
+ max_new_tokens: int,
+) -> tuple[torch.Tensor, list[str]]:
+ captions: list[str] = []
+ for chunk in particles.split(16):
+ prefix_embedding = prefix(chunk).to(dtype)
+ attention = torch.ones(
+ prefix_embedding.shape[:2],
+ dtype=torch.long,
+ device=particles.device,
+ )
+ output = lm.generate(
+ inputs_embeds=prefix_embedding,
+ attention_mask=attention,
+ max_new_tokens=max_new_tokens,
+ do_sample=False,
+ eos_token_id=tokenizer.eos_token_id,
+ pad_token_id=tokenizer.pad_token_id,
+ )
+ captions.extend(
+ tokenizer.batch_decode(output, skip_special_tokens=True)
+ )
+ anchor_parts: list[torch.Tensor] = []
+ for start in range(0, len(captions), 32):
+ tokens = tokenizer(
+ captions[start : start + 32],
+ padding=True,
+ truncation=True,
+ max_length=64,
+ return_tensors="pt",
+ )
+ tokens = {
+ key: value.to(particles.device)
+ for key, value in tokens.items()
+ }
+ result = lm(
+ **tokens, output_hidden_states=True, return_dict=True
+ )
+ anchor_parts.append(
+ mean_pool(
+ result.hidden_states[-1], tokens["attention_mask"]
+ ).float()
+ )
+ return F.normalize(torch.cat(anchor_parts), dim=-1), captions
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ vision, text, vision_lookup, _ = load_feature_pair(
+ args.vision, args.text
+ )
+ orbit_state = torch.load(
+ args.text_orbits, map_location="cpu", weights_only=False
+ )
+ orbit_lookup = {
+ int(row): index for index, row in enumerate(orbit_state["rows"])
+ }
+ orbit_mean = F.normalize(
+ orbit_state["features"].float().mean(1), dim=-1
+ )
+ rows = manifest[args.split][: args.samples]
+ visual = select_rows(
+ vision["features"], vision_lookup, rows
+ ).to(args.device)
+ paired_text = select_rows(
+ orbit_mean, orbit_lookup, rows
+ ).to(args.device)
+ text_population = select_rows(
+ orbit_mean, orbit_lookup, manifest["text_only_train"]
+ ).to(args.device)
+
+ prefix, prefix_state = load_prefix(args.prefix, args.device)
+ if prefix.semantic_dim != text_population.shape[-1]:
+ raise ValueError("Prefix and text-orbit dimensions do not match")
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+ for parameter in lm.parameters():
+ parameter.requires_grad_(False)
+ for parameter in prefix.parameters():
+ parameter.requires_grad_(False)
+
+ generator = torch.Generator(device=args.device).manual_seed(args.seed)
+ initial_indices = torch.randperm(
+ len(text_population),
+ generator=generator,
+ device=args.device,
+ )[: len(visual)]
+ particles = torch.nn.Parameter(text_population[initial_indices].clone())
+ initial_particles = particles.detach().cpu()
+ directions, target_quantiles = projection_quantile_target(
+ text_population,
+ len(particles),
+ args.projections,
+ generator,
+ )
+ visual_relation, visual_standardized = standardized_relation(visual)
+ optimizer = torch.optim.Adam([particles], lr=args.lr)
+ functional_anchor, functional_captions = refresh_functional_anchor(
+ particles.detach(),
+ prefix,
+ lm,
+ tokenizer,
+ dtype,
+ args.max_new_tokens,
+ )
+
+ history: list[dict] = []
+ refresh_captions: dict[int, list[str]] = {
+ 0: functional_captions[:25]
+ }
+ for step in range(args.steps + 1):
+ if step > 0 and step % args.refresh_steps == 0:
+ functional_anchor, functional_captions = (
+ refresh_functional_anchor(
+ particles.detach(),
+ prefix,
+ lm,
+ tokenizer,
+ dtype,
+ args.max_new_tokens,
+ )
+ )
+ refresh_captions[step] = functional_captions[:25]
+ relation, conditional = relation_field_energy(
+ visual_relation, visual_standardized, particles
+ )
+ distribution = sliced_distribution_energy(
+ particles, directions, target_quantiles
+ )
+ functional = (
+ 1
+ - (
+ F.normalize(particles, dim=-1)
+ * functional_anchor.detach()
+ ).sum(-1)
+ ).mean()
+ loss = (
+ args.relation_weight * relation
+ + args.conditional_weight * conditional
+ + args.distribution_weight * distribution
+ + args.functional_weight * functional
+ )
+ if step % 10 == 0 or step == args.steps:
+ record = {
+ "step": step,
+ "total": float(loss.detach()),
+ "relation": float(relation.detach()),
+ "conditional": float(conditional.detach()),
+ "distribution": float(distribution.detach()),
+ "functional": float(functional.detach()),
+ "paired_evaluation_only": retrieval_metrics(
+ particles.detach(), paired_text
+ ),
+ }
+ history.append(record)
+ print(json.dumps(record))
+ if step == args.steps:
+ break
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ torch.nn.utils.clip_grad_norm_([particles], 2.0)
+ optimizer.step()
+ with torch.no_grad():
+ particles.copy_(F.normalize(particles, dim=-1))
+
+ result = {
+ "protocol": (
+ "No image-text pair and no cross-modal parameter is used. The "
+ "functional language energy is a frozen decode-reencode cycle "
+ "through a text-only prefix interface and frozen Qwen."
+ ),
+ "mode": "functional_cycle_language_latent_particles",
+ "split": args.split,
+ "rows": rows,
+ "args": vars(args),
+ "history": history,
+ "prefix_training": prefix_state["training"],
+ "refresh_caption_examples": refresh_captions,
+ }
+ state = {
+ **result,
+ "initial_particles": initial_particles,
+ "final_particles": F.normalize(
+ particles.detach(), dim=-1
+ ).cpu(),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ write_json(args.metrics_output, result)
+ print(f"Wrote {args.output} and {args.metrics_output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/energy_infer.py b/worldalign/energy_infer.py
new file mode 100644
index 0000000..d5a67b6
--- /dev/null
+++ b/worldalign/energy_infer.py
@@ -0,0 +1,249 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+
+from .common import read_json, seed_everything, write_json
+from .energy import (
+ projection_quantile_target,
+ prototype_manifold_energy,
+ relation_field_energy,
+ retrieval_metrics,
+ sliced_distribution_energy,
+ standardized_relation,
+)
+from .io import load_feature_pair, select_rows
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument(
+ "--text-orbits",
+ help=(
+ "Optional multi-description feature cache. When supplied, each "
+ "language particle is the normalized mean of its observation "
+ "orbit."
+ ),
+ )
+ parser.add_argument("--prototypes", default="artifacts/gw.pt")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--steps", type=int, default=100)
+ parser.add_argument("--lr", type=float, default=0.03)
+ parser.add_argument("--projections", type=int, default=128)
+ parser.add_argument("--relation-weight-start", type=float, default=0.4)
+ parser.add_argument("--relation-weight-end", type=float, default=2.0)
+ parser.add_argument("--conditional-weight-start", type=float, default=0.05)
+ parser.add_argument("--conditional-weight-end", type=float, default=0.2)
+ parser.add_argument("--distribution-weight", type=float, default=80.0)
+ parser.add_argument("--manifold-weight", type=float, default=1.0)
+ parser.add_argument("--noise", type=float, default=0.01)
+ parser.add_argument(
+ "--shuffle-visual-energy",
+ action="store_true",
+ help=(
+ "Evaluation control: permute image particles before constructing "
+ "the energy while leaving evaluation rows unchanged."
+ ),
+ )
+ parser.add_argument("--device", default="cuda:1")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument(
+ "--output", default="artifacts/energy_free_test.pt"
+ )
+ parser.add_argument(
+ "--metrics-output", default="artifacts/energy_free_test.json"
+ )
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ vision, text, vision_lookup, text_lookup = load_feature_pair(
+ args.vision, args.text
+ )
+ rows = manifest[args.split][: args.samples]
+ visual = select_rows(
+ vision["features"], vision_lookup, rows
+ ).to(args.device)
+ paired_text = select_rows(
+ text["features"], text_lookup, rows
+ ).to(args.device)
+ text_population = select_rows(
+ text["features"], text_lookup, manifest["text_only_train"]
+ ).to(args.device)
+ if args.text_orbits:
+ orbit_state = torch.load(
+ args.text_orbits, map_location="cpu", weights_only=False
+ )
+ orbit_lookup = {
+ int(row): index
+ for index, row in enumerate(orbit_state["rows"])
+ }
+ orbit_mean = F.normalize(
+ orbit_state["features"].float().mean(1), dim=-1
+ )
+ paired_text = select_rows(
+ orbit_mean, orbit_lookup, rows
+ ).to(args.device)
+ text_population = select_rows(
+ orbit_mean, orbit_lookup, manifest["text_only_train"]
+ ).to(args.device)
+ if args.shuffle_visual_energy:
+ control_generator = torch.Generator(
+ device=args.device
+ ).manual_seed(args.seed + 10_000)
+ visual = visual[
+ torch.randperm(
+ len(visual),
+ generator=control_generator,
+ device=args.device,
+ )
+ ]
+ prototype_state = torch.load(
+ args.prototypes, map_location="cpu", weights_only=False
+ )
+ prototypes = prototype_state["text_centers"].to(args.device)
+ if prototypes.shape[-1] != text_population.shape[-1]:
+ raise ValueError(
+ "Prototype dimension does not match text features; use the GW "
+ "cache built from the selected text backbone."
+ )
+
+ generator = torch.Generator(device=args.device).manual_seed(args.seed)
+ initial_indices = torch.randperm(
+ len(text_population),
+ generator=generator,
+ device=args.device,
+ )[: len(visual)]
+ initial = text_population[initial_indices].clone()
+ initial = initial + args.noise * torch.randn(
+ initial.shape, generator=generator, device=args.device
+ )
+ particles = torch.nn.Parameter(F.normalize(initial, dim=-1))
+ directions, target_quantiles = projection_quantile_target(
+ text_population,
+ particles=len(particles),
+ projections=args.projections,
+ generator=generator,
+ )
+ visual_relation, visual_standardized = standardized_relation(visual)
+ optimizer = torch.optim.Adam([particles], lr=args.lr)
+
+ def energy_values() -> tuple[torch.Tensor, ...]:
+ relation, conditional = relation_field_energy(
+ visual_relation, visual_standardized, particles
+ )
+ distribution = sliced_distribution_energy(
+ particles, directions, target_quantiles
+ )
+ manifold = prototype_manifold_energy(particles, prototypes)
+ return relation, conditional, distribution, manifold
+
+ history: list[dict] = []
+ initial_cpu = F.normalize(particles.detach(), dim=-1).cpu()
+ for step in range(args.steps + 1):
+ relation, conditional, distribution, manifold = energy_values()
+ progress = min(step / max(args.steps, 1), 1.0)
+ relation_weight = (
+ args.relation_weight_start
+ + progress
+ * (args.relation_weight_end - args.relation_weight_start)
+ )
+ conditional_weight = (
+ args.conditional_weight_start
+ + progress
+ * (
+ args.conditional_weight_end
+ - args.conditional_weight_start
+ )
+ )
+ loss = (
+ relation_weight * relation
+ + conditional_weight * conditional
+ + args.distribution_weight * distribution
+ + args.manifold_weight * manifold
+ )
+ if step % 20 == 0 or step == args.steps:
+ history.append(
+ {
+ "step": step,
+ "total": float(loss.detach()),
+ "relation": float(relation.detach()),
+ "conditional": float(conditional.detach()),
+ "distribution": float(distribution.detach()),
+ "manifold": float(manifold.detach()),
+ "paired_evaluation_only": retrieval_metrics(
+ particles.detach(), paired_text
+ ),
+ }
+ )
+ print(json.dumps(history[-1]))
+ if step == args.steps:
+ break
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ torch.nn.utils.clip_grad_norm_([particles], 2.0)
+ optimizer.step()
+ with torch.no_grad():
+ particles.copy_(F.normalize(particles, dim=-1))
+
+ with torch.no_grad():
+ oracle_relation, oracle_conditional = relation_field_energy(
+ visual_relation, visual_standardized, paired_text
+ )
+ oracle_distribution = sliced_distribution_energy(
+ paired_text, directions, target_quantiles
+ )
+ oracle_manifold = prototype_manifold_energy(
+ paired_text, prototypes
+ )
+ result = {
+ "protocol": (
+ "No cross-modal map and no image-text pair is used by the energy "
+ "or optimizer. Paired text is loaded only for trajectory and "
+ "oracle diagnostics; the final step is fixed by CLI arguments."
+ ),
+ "mode": "free_language_latent_particles",
+ "split": args.split,
+ "rows": rows,
+ "vision_model": vision["model"],
+ "text_model": text["model"],
+ "text_observation": (
+ "multi-description orbit mean"
+ if args.text_orbits
+ else "single description"
+ ),
+ "args": vars(args),
+ "history": history,
+ "oracle_energy_components": {
+ "relation": float(oracle_relation),
+ "conditional": float(oracle_conditional),
+ "distribution": float(oracle_distribution),
+ "manifold": float(oracle_manifold),
+ },
+ }
+ state = {
+ **result,
+ "initial_particles": initial_cpu,
+ "final_particles": F.normalize(
+ particles.detach(), dim=-1
+ ).cpu(),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ write_json(args.metrics_output, result)
+ print(f"Wrote {args.output} and {args.metrics_output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/evaluate.py b/worldalign/evaluate.py
new file mode 100644
index 0000000..6f97873
--- /dev/null
+++ b/worldalign/evaluate.py
@@ -0,0 +1,149 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+import re
+from collections import Counter
+
+import torch
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import (
+ dtype_for_device,
+ read_json,
+ retrieval_metrics,
+ write_json,
+)
+from .io import load_feature_pair, select_rows
+from .models import load_bridge, load_prefix
+
+
+TOKEN_RE = re.compile(r"[a-z0-9]+")
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--bridge", required=True)
+ p.add_argument("--prefix")
+ p.add_argument("--split", choices=["val", "test"], default="test")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--generation-samples", type=int, default=100)
+ p.add_argument("--max-new-tokens", type=int, default=32)
+ p.add_argument(
+ "--shuffle-mapped",
+ action="store_true",
+ help="Permute image-conditioned latents before retrieval/generation as a null control.",
+ )
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument("--output", default="artifacts/evaluation.json")
+ return p.parse_args()
+
+
+def unigram_f1(candidate: str, references: list[str]) -> float:
+ candidate_tokens = TOKEN_RE.findall(candidate.lower())
+ if not candidate_tokens:
+ return 0.0
+ candidate_count = Counter(candidate_tokens)
+ best = 0.0
+ for reference in references:
+ reference_count = Counter(TOKEN_RE.findall(reference.lower()))
+ overlap = sum((candidate_count & reference_count).values())
+ precision = overlap / max(sum(candidate_count.values()), 1)
+ recall = overlap / max(sum(reference_count.values()), 1)
+ f1 = 2 * precision * recall / max(precision + recall, 1e-12)
+ best = max(best, f1)
+ return best
+
+
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
+ rows = manifest[args.split]
+ x = select_rows(vision["features"], vlookup, rows)
+ y = select_rows(text["features"], tlookup, rows)
+
+ bridge, bridge_state = load_bridge(args.bridge, args.device)
+ mapped = []
+ with torch.inference_mode():
+ for chunk in x.split(512):
+ mapped.append(bridge(chunk.to(args.device)).cpu())
+ mapped = torch.cat(mapped)
+ if args.shuffle_mapped:
+ generator = torch.Generator().manual_seed(args.seed)
+ mapped = mapped[torch.randperm(len(mapped), generator=generator)]
+ result: dict = {
+ "split": args.split,
+ "samples": len(rows),
+ "bridge_mode": bridge_state["mode"],
+ "shuffle_mapped": args.shuffle_mapped,
+ "retrieval": retrieval_metrics(mapped, y),
+ }
+
+ if args.prefix:
+ prefix, prefix_state = load_prefix(args.prefix, args.device)
+ if prefix_state["text_model"] != text["model"]:
+ raise ValueError("Prefix adapter and text feature model differ")
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+ generated: list[str] = []
+ n = min(args.generation_samples, len(rows))
+ with torch.inference_mode():
+ for chunk in mapped[:n].split(16):
+ prefix_embeds = prefix(chunk.to(args.device)).to(dtype)
+ attention = torch.ones(
+ prefix_embeds.shape[:2],
+ dtype=torch.long,
+ device=args.device,
+ )
+ output = lm.generate(
+ inputs_embeds=prefix_embeds,
+ attention_mask=attention,
+ max_new_tokens=args.max_new_tokens,
+ do_sample=False,
+ eos_token_id=tokenizer.eos_token_id,
+ pad_token_id=tokenizer.pad_token_id,
+ )
+ generated.extend(tokenizer.batch_decode(output, skip_special_tokens=True))
+
+ caption_lookup = {
+ int(row): caps
+ for row, caps in zip(text["rows"], text["all_captions"])
+ }
+ records = []
+ scores = []
+ for row, caption in zip(rows[:n], generated):
+ refs = caption_lookup[int(row)]
+ score = unigram_f1(caption, refs)
+ scores.append(score)
+ records.append(
+ {
+ "row": int(row),
+ "generated": caption,
+ "references": refs,
+ "unigram_f1": score,
+ }
+ )
+ result["generation"] = {
+ "samples": n,
+ "mean_best_reference_unigram_f1": sum(scores) / max(len(scores), 1),
+ "examples": records[:25],
+ }
+
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2, ensure_ascii=False))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/evaluate_energy_prefix.py b/worldalign/evaluate_energy_prefix.py
new file mode 100644
index 0000000..2f3af7a
--- /dev/null
+++ b/worldalign/evaluate_energy_prefix.py
@@ -0,0 +1,151 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import dtype_for_device, write_json
+from .evaluate import unigram_f1
+from .io import load_feature_pair, select_rows
+from .models import load_prefix
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--energy", default="artifacts/energy_free_test.pt")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits")
+ parser.add_argument("--prefix", default="artifacts/prefix.pt")
+ parser.add_argument("--samples", type=int, default=100)
+ parser.add_argument("--max-new-tokens", type=int, default=32)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument(
+ "--output", default="artifacts/energy_prefix_evaluation.json"
+ )
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ energy = torch.load(args.energy, map_location="cpu", weights_only=False)
+ _, text, _, text_lookup = load_feature_pair(args.vision, args.text)
+ rows = energy["rows"][: args.samples]
+ initial = energy["initial_particles"][: args.samples]
+ final = energy["final_particles"][: args.samples]
+ oracle = select_rows(text["features"], text_lookup, rows)
+ if args.text_orbits:
+ orbit_state = torch.load(
+ args.text_orbits, map_location="cpu", weights_only=False
+ )
+ orbit_lookup = {
+ int(row): index
+ for index, row in enumerate(orbit_state["rows"])
+ }
+ orbit_mean = torch.nn.functional.normalize(
+ orbit_state["features"].float().mean(1), dim=-1
+ )
+ oracle = select_rows(orbit_mean, orbit_lookup, rows)
+ reference_lookup = {
+ int(row): captions
+ for row, captions in zip(text["rows"], text["all_captions"])
+ }
+
+ prefix, prefix_state = load_prefix(args.prefix, args.device)
+ if prefix.semantic_dim != initial.shape[-1]:
+ raise ValueError("Energy latent and text-only prefix dimensions differ")
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+
+ conditions = {"initial": initial, "final": final, "oracle": oracle}
+ decoded: dict[str, list[str]] = {key: [] for key in conditions}
+ with torch.inference_mode():
+ for name, semantic in conditions.items():
+ for chunk in semantic.split(16):
+ embeds = prefix(chunk.to(args.device)).to(dtype)
+ attention = torch.ones(
+ embeds.shape[:2], dtype=torch.long, device=args.device
+ )
+ output = lm.generate(
+ inputs_embeds=embeds,
+ attention_mask=attention,
+ max_new_tokens=args.max_new_tokens,
+ do_sample=False,
+ eos_token_id=tokenizer.eos_token_id,
+ pad_token_id=tokenizer.pad_token_id,
+ )
+ decoded[name].extend(
+ tokenizer.batch_decode(
+ output, skip_special_tokens=True
+ )
+ )
+
+ scores = {}
+ per_condition: dict[str, list[float]] = {}
+ for name, generations in decoded.items():
+ values = [
+ unigram_f1(generation, reference_lookup[int(row)])
+ for row, generation in zip(rows, generations)
+ ]
+ per_condition[name] = values
+ scores[name] = sum(values) / max(len(values), 1)
+ difference = np.asarray(per_condition["final"]) - np.asarray(
+ per_condition["initial"]
+ )
+ bootstrap_generator = np.random.default_rng(20260729)
+ bootstrap = difference[
+ bootstrap_generator.integers(
+ 0, len(difference), size=(10_000, len(difference))
+ )
+ ].mean(1)
+ result = {
+ "samples": len(rows),
+ "mean_best_reference_unigram_f1": scores,
+ "paired_final_minus_initial": {
+ "mean": float(difference.mean()),
+ "bootstrap_95_percentile_interval": [
+ float(np.quantile(bootstrap, 0.025)),
+ float(np.quantile(bootstrap, 0.975)),
+ ],
+ "improved": int((difference > 0).sum()),
+ "tied": int((difference == 0).sum()),
+ "worsened": int((difference < 0).sum()),
+ },
+ "energy_protocol": energy["protocol"],
+ "prefix_training": prefix_state["training"],
+ "per_sample_unigram_f1": [
+ {
+ "row": int(row),
+ **{
+ name: per_condition[name][index]
+ for name in per_condition
+ },
+ }
+ for index, row in enumerate(rows)
+ ],
+ "examples": [
+ {
+ "row": int(row),
+ "references": reference_lookup[int(row)],
+ **{name: decoded[name][i] for name in decoded},
+ }
+ for i, row in enumerate(rows[:25])
+ ],
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2, ensure_ascii=False))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/evaluate_prefix.py b/worldalign/evaluate_prefix.py
new file mode 100644
index 0000000..2d53a51
--- /dev/null
+++ b/worldalign/evaluate_prefix.py
@@ -0,0 +1,93 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import dtype_for_device, read_json, write_json
+from .evaluate import unigram_f1
+from .io import load_feature_pair, select_rows
+from .models import load_prefix
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--prefix", default="artifacts/prefix.pt")
+ p.add_argument("--split", choices=["val", "test"], default="val")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--samples", type=int, default=100)
+ p.add_argument("--max-new-tokens", type=int, default=32)
+ p.add_argument("--output", default="artifacts/prefix_evaluation.json")
+ return p.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ _, text, _, tlookup = load_feature_pair(args.vision, args.text)
+ rows = manifest[args.split][: args.samples]
+ semantic = select_rows(text["features"], tlookup, rows)
+ caption_lookup = {
+ int(row): caps for row, caps in zip(text["rows"], text["all_captions"])
+ }
+ prefix, prefix_state = load_prefix(args.prefix, args.device)
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+
+ generated = []
+ with torch.inference_mode():
+ for chunk in semantic.split(16):
+ embeds = prefix(chunk.to(args.device)).to(dtype)
+ attention = torch.ones(
+ embeds.shape[:2], dtype=torch.long, device=args.device
+ )
+ output = lm.generate(
+ inputs_embeds=embeds,
+ attention_mask=attention,
+ max_new_tokens=args.max_new_tokens,
+ do_sample=False,
+ eos_token_id=tokenizer.eos_token_id,
+ pad_token_id=tokenizer.pad_token_id,
+ )
+ generated.extend(tokenizer.batch_decode(output, skip_special_tokens=True))
+
+ records = []
+ for row, caption in zip(rows, generated):
+ references = caption_lookup[int(row)]
+ records.append(
+ {
+ "row": int(row),
+ "generated": caption,
+ "references": references,
+ "unigram_f1": unigram_f1(caption, references),
+ }
+ )
+ result = {
+ "split": args.split,
+ "samples": len(records),
+ "mean_best_reference_unigram_f1": sum(
+ r["unigram_f1"] for r in records
+ )
+ / max(len(records), 1),
+ "examples": records[:25],
+ "training": prefix_state["training"],
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2, ensure_ascii=False))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/extract_text.py b/worldalign/extract_text.py
new file mode 100644
index 0000000..8586e21
--- /dev/null
+++ b/worldalign/extract_text.py
@@ -0,0 +1,95 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import torch
+from datasets import load_dataset
+from tqdm import tqdm
+from transformers import AutoModel, AutoTokenizer
+
+from .common import batch_indices, dtype_for_device, read_json
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--output", default="artifacts/text.pt")
+ p.add_argument("--model", default="Qwen/Qwen2.5-0.5B")
+ p.add_argument("--device", default="cuda:3")
+ p.add_argument("--batch-size", type=int, default=96)
+ p.add_argument("--max-length", type=int, default=64)
+ p.add_argument(
+ "--layer",
+ type=int,
+ default=-1,
+ help="Hidden-state index; -1 is the final transformer output.",
+ )
+ p.add_argument("--limit", type=int)
+ return p.parse_args()
+
+
+def mean_pool(hidden: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
+ weights = mask.to(hidden.dtype).unsqueeze(-1)
+ return (hidden * weights).sum(1) / weights.sum(1).clamp_min(1)
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ rows = manifest["all_rows"]
+ if args.limit:
+ rows = rows[: args.limit]
+ dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"])
+ # Avoid decoding the 4.3GB image column when only captions are requested.
+ dataset = dataset.remove_columns("image")
+ tokenizer = AutoTokenizer.from_pretrained(args.model)
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ tokenizer.padding_side = "right"
+ dtype = dtype_for_device(args.device)
+ # The backbone is sufficient for hidden states. Using a causal-LM wrapper
+ # would also materialize vocabulary logits that are discarded here.
+ model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(args.device)
+ model.eval()
+
+ outputs: list[torch.Tensor] = []
+ captions: list[str] = []
+ all_captions: list[list[str]] = []
+ for ids in tqdm(
+ batch_indices(len(rows), args.batch_size), desc="Qwen text features"
+ ):
+ batch_captions = [dataset[int(rows[i])]["caption"] for i in ids]
+ texts = [caps[0] for caps in batch_captions]
+ tokens = tokenizer(
+ texts,
+ padding=True,
+ truncation=True,
+ max_length=args.max_length,
+ return_tensors="pt",
+ )
+ tokens = {k: v.to(args.device) for k, v in tokens.items()}
+ result = model(**tokens, output_hidden_states=True, return_dict=True)
+ feature = mean_pool(
+ result.hidden_states[args.layer], tokens["attention_mask"]
+ )
+ outputs.append(feature.float().cpu())
+ captions.extend(texts)
+ all_captions.extend(batch_captions)
+
+ value = {
+ "model": args.model,
+ "layer": args.layer,
+ "rows": rows,
+ "features": torch.cat(outputs),
+ "captions": captions,
+ "all_captions": all_captions,
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(value, args.output)
+ print(f"Wrote {args.output}: {tuple(value['features'].shape)}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/extract_text_orbits.py b/worldalign/extract_text_orbits.py
new file mode 100644
index 0000000..43b9020
--- /dev/null
+++ b/worldalign/extract_text_orbits.py
@@ -0,0 +1,108 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+from datasets import load_dataset
+import torch
+from tqdm import tqdm
+from transformers import AutoModel, AutoTokenizer
+
+from .common import batch_indices, dtype_for_device, read_json
+from .extract_text import mean_pool
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument(
+ "--output", default="artifacts/text_orbits_qwen0p5b.pt"
+ )
+ parser.add_argument("--model", default="Qwen/Qwen2.5-0.5B")
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--batch-size", type=int, default=128)
+ parser.add_argument("--max-length", type=int, default=64)
+ parser.add_argument("--layer", type=int, default=-1)
+ parser.add_argument(
+ "--row-groups",
+ default="text_only_train,val,test",
+ help="Comma-separated manifest row lists to encode.",
+ )
+ parser.add_argument("--limit", type=int)
+ return parser.parse_args()
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ groups = [group.strip() for group in args.row_groups.split(",")]
+ rows = list(
+ dict.fromkeys(
+ int(row)
+ for group in groups
+ for row in manifest[group]
+ )
+ )
+ if args.limit:
+ rows = rows[: args.limit]
+ dataset = load_dataset(
+ manifest["dataset"], split=manifest["dataset_split"]
+ ).remove_columns("image")
+ captions = [dataset[row]["caption"] for row in rows]
+ views = len(captions[0])
+ if any(len(items) != views for items in captions):
+ raise ValueError("Every text orbit must have the same view count")
+ flat_text = [text for items in captions for text in items]
+
+ tokenizer = AutoTokenizer.from_pretrained(args.model)
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ tokenizer.padding_side = "right"
+ dtype = dtype_for_device(args.device)
+ model = AutoModel.from_pretrained(
+ args.model, torch_dtype=dtype
+ ).to(args.device)
+ model.eval()
+
+ output: list[torch.Tensor] = []
+ for indices in tqdm(
+ batch_indices(len(flat_text), args.batch_size),
+ desc="Qwen text orbits",
+ ):
+ tokens = tokenizer(
+ [flat_text[index] for index in indices],
+ padding=True,
+ truncation=True,
+ max_length=args.max_length,
+ return_tensors="pt",
+ )
+ tokens = {key: value.to(args.device) for key, value in tokens.items()}
+ result = model(
+ **tokens, output_hidden_states=True, return_dict=True
+ )
+ output.append(
+ mean_pool(
+ result.hidden_states[args.layer],
+ tokens["attention_mask"],
+ )
+ .float()
+ .cpu()
+ )
+ features = torch.cat(output).reshape(len(rows), views, -1)
+ state = {
+ "model": args.model,
+ "layer": args.layer,
+ "rows": rows,
+ "features": features,
+ "captions": captions,
+ "views_per_orbit": views,
+ "row_groups": groups,
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ print(f"Wrote {args.output}: {tuple(features.shape)}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/extract_vision.py b/worldalign/extract_vision.py
new file mode 100644
index 0000000..348c6a4
--- /dev/null
+++ b/worldalign/extract_vision.py
@@ -0,0 +1,61 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import torch
+from datasets import load_dataset
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel
+
+from .common import batch_indices, dtype_for_device, read_json
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--output", default="artifacts/vision.pt")
+ p.add_argument("--model", default="facebook/dinov2-small")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--batch-size", type=int, default=96)
+ p.add_argument("--limit", type=int)
+ return p.parse_args()
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ rows = manifest["all_rows"]
+ if args.limit:
+ rows = rows[: args.limit]
+ dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"])
+ processor = AutoImageProcessor.from_pretrained(args.model)
+ dtype = dtype_for_device(args.device)
+ model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(args.device)
+ model.eval()
+
+ outputs: list[torch.Tensor] = []
+ for ids in tqdm(
+ batch_indices(len(rows), args.batch_size), desc="DINO image features"
+ ):
+ images = [dataset[int(rows[i])]["image"].convert("RGB") for i in ids]
+ batch = processor(images=images, return_tensors="pt")
+ batch = {k: v.to(args.device) for k, v in batch.items()}
+ result = model(**batch)
+ feature = result.last_hidden_state[:, 0]
+ outputs.append(feature.float().cpu())
+
+ value = {
+ "model": args.model,
+ "rows": rows,
+ "features": torch.cat(outputs),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(value, args.output)
+ print(f"Wrote {args.output}: {tuple(value['features'].shape)}")
+
+
+if __name__ == "__main__":
+ main()
+
diff --git a/worldalign/gw.py b/worldalign/gw.py
new file mode 100644
index 0000000..62c3f5c
--- /dev/null
+++ b/worldalign/gw.py
@@ -0,0 +1,79 @@
+from __future__ import annotations
+
+import numpy as np
+import ot
+import torch
+from sklearn.cluster import MiniBatchKMeans
+
+from .common import pairwise_cosine_distance
+
+
+def gw_pseudo_targets(
+ vision: torch.Tensor,
+ text: torch.Tensor,
+ clusters: int,
+ seed: int,
+ max_iter: int = 100,
+) -> dict[str, torch.Tensor | float | int]:
+ """Build structure-only pseudo-correspondences between modality prototypes."""
+ x = vision.float().numpy()
+ y = text.float().numpy()
+ k = min(clusters, len(x), len(y))
+ vx = MiniBatchKMeans(
+ n_clusters=k,
+ random_state=seed,
+ batch_size=min(4096, len(x)),
+ n_init=3,
+ max_iter=200,
+ ).fit(x)
+ ty = MiniBatchKMeans(
+ n_clusters=k,
+ random_state=seed + 1,
+ batch_size=min(4096, len(y)),
+ n_init=3,
+ max_iter=200,
+ ).fit(y)
+
+ cx = pairwise_cosine_distance(vx.cluster_centers_)
+ cy = pairwise_cosine_distance(ty.cluster_centers_)
+ p = np.bincount(vx.labels_, minlength=k).astype(np.float64)
+ q = np.bincount(ty.labels_, minlength=k).astype(np.float64)
+ p /= p.sum()
+ q /= q.sum()
+
+ coupling, log = ot.gromov.gromov_wasserstein(
+ cx,
+ cy,
+ p,
+ q,
+ loss_fun="square_loss",
+ armijo=False,
+ log=True,
+ max_iter=max_iter,
+ tol_rel=1e-8,
+ tol_abs=1e-8,
+ )
+ target_centers = coupling @ ty.cluster_centers_
+ target_centers /= p[:, None].clip(min=1e-12)
+ target_centers /= np.linalg.norm(target_centers, axis=1, keepdims=True).clip(
+ min=1e-12
+ )
+ row_entropy = -np.sum(
+ (coupling / p[:, None].clip(min=1e-12))
+ * np.log(
+ (coupling / p[:, None].clip(min=1e-12)).clip(min=1e-12)
+ ),
+ axis=1,
+ ).mean()
+
+ return {
+ "vision_centers": torch.from_numpy(vx.cluster_centers_).float(),
+ "text_centers": torch.from_numpy(ty.cluster_centers_).float(),
+ "target_centers": torch.from_numpy(target_centers).float(),
+ "vision_assignments": torch.from_numpy(vx.labels_).long(),
+ "coupling": torch.from_numpy(coupling).float(),
+ "gw_distance": float(log["gw_dist"]),
+ "coupling_row_entropy": float(row_entropy),
+ "clusters": k,
+ }
+
diff --git a/worldalign/io.py b/worldalign/io.py
new file mode 100644
index 0000000..6cac007
--- /dev/null
+++ b/worldalign/io.py
@@ -0,0 +1,30 @@
+from __future__ import annotations
+
+import torch
+
+from .common import normalized
+
+
+def load_feature_pair(
+ vision_path: str, text_path: str
+) -> tuple[dict, dict, dict[int, int], dict[int, int]]:
+ vision = torch.load(vision_path, map_location="cpu", weights_only=False)
+ text = torch.load(text_path, map_location="cpu", weights_only=False)
+ vision["features"] = normalized(vision["features"])
+ text["features"] = normalized(text["features"])
+ vision_lookup = {int(row): i for i, row in enumerate(vision["rows"])}
+ text_lookup = {int(row): i for i, row in enumerate(text["rows"])}
+ return vision, text, vision_lookup, text_lookup
+
+
+def select_rows(
+ feature: torch.Tensor, lookup: dict[int, int], rows: list[int]
+) -> torch.Tensor:
+ missing = [int(row) for row in rows if int(row) not in lookup]
+ if missing:
+ raise KeyError(
+ f"{len(missing)} requested rows are absent from feature cache; "
+ f"first missing rows: {missing[:5]}"
+ )
+ return feature[[lookup[int(row)] for row in rows]]
+
diff --git a/worldalign/manifold_gate.py b/worldalign/manifold_gate.py
new file mode 100644
index 0000000..30acceb
--- /dev/null
+++ b/worldalign/manifold_gate.py
@@ -0,0 +1,677 @@
+"""On-manifold identifiability gate for assignment energies.
+
+The configuration space is restricted to permutations of real frozen text
+states. On this space the language-only energy terms (sliced distribution,
+prototype manifold) depend only on the set of states and are therefore
+constant; the only varying terms are the cross-modal relation MSE and the
+multiscale conditional KL from ``energy.relation_field_energy``. Hidden
+pairs are used only to score orderings, never inside the energy.
+
+Gates, in increasing strictness:
+
+A. global ranking: energy of the true assignment against random and
+ structured permutations;
+B. local identifiability: exact delta energy of every transposition of the
+ true assignment, via a closed form that one matrix product evaluates for
+ all N(N-1)/2 swaps;
+C. basin audit: exact steepest 2-swap descent from the true assignment and
+ from random assignments, with a local-minimum certificate. Descent from
+ random assignments doubles as a blind transductive recovery baseline and
+ as a search for on-manifold counterfeits.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+
+from .common import read_json, seed_everything, write_json
+from .io import load_feature_pair, select_rows
+
+TEMPERATURES = (0.03, 0.07, 0.15)
+M30_RELATION_WEIGHT = 2.0
+M30_CONDITIONAL_WEIGHT = 0.2
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--dataset", choices=["flickr", "vg"], default="flickr")
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt")
+ parser.add_argument(
+ "--text-mode",
+ choices=["single", "orbit_mean"],
+ default="orbit_mean",
+ help="Flickr language node state definition.",
+ )
+ parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt")
+ parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt")
+ parser.add_argument(
+ "--vg-ground-truth",
+ default="artifacts/vg_5k/ground_truth.private.jsonl",
+ help="Private pairing, loaded only to construct the evaluation order.",
+ )
+ parser.add_argument(
+ "--vg-bundle-channels",
+ action="store_true",
+ help="Add std/q10/q90 view-pair relation channels to the scalar mean.",
+ )
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--subset-seed", type=int, default=0)
+ parser.add_argument("--random-perms", type=int, default=1000)
+ parser.add_argument("--derangement-samples", type=int, default=200)
+ parser.add_argument("--descent-restarts", type=int, default=3)
+ parser.add_argument("--descent-max-steps", type=int, default=200000)
+ parser.add_argument(
+ "--descent-objective",
+ choices=["mse", "m30_total"],
+ default="m30_total",
+ help="m30_total preselects swaps by closed-form MSE and verifies the "
+ "exact weighted MSE+KL objective on the best candidates.",
+ )
+ parser.add_argument("--descent-verify-top", type=int, default=64)
+ parser.add_argument("--device", default="cpu")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument("--output", default="artifacts/manifold_gate/gate.json")
+ parser.add_argument("--trajectory-output")
+ return parser.parse_args()
+
+
+def cosine_relation(features: torch.Tensor) -> torch.Tensor:
+ features = F.normalize(features.double(), dim=-1)
+ return features @ features.T
+
+
+def offdiag_mask(size: int, device: torch.device | str) -> torch.Tensor:
+ return ~torch.eye(size, dtype=torch.bool, device=device)
+
+
+def standardize_relation(relation: torch.Tensor) -> tuple[torch.Tensor, float, float]:
+ """Standardized copy with zeroed diagonal.
+
+ The mean and std are taken over off-diagonal values, which are a
+ permutation-invariant set, so the same constants apply to every
+ assignment of the same states.
+ """
+ mask = offdiag_mask(len(relation), relation.device)
+ values = relation[mask]
+ mean = values.mean()
+ std = values.std().clamp_min(1e-6)
+ standardized = (relation - mean) / std
+ standardized = standardized.masked_fill(~mask, 0.0)
+ return standardized, float(mean), float(std)
+
+
+def relation_mse(
+ text_standardized: torch.Tensor, visual_standardized: torch.Tensor
+) -> torch.Tensor:
+ mask = offdiag_mask(len(visual_standardized), visual_standardized.device)
+ return (text_standardized[mask] - visual_standardized[mask]).square().mean()
+
+
+def conditional_kl(
+ text_relation: torch.Tensor, visual_relation: torch.Tensor
+) -> torch.Tensor:
+ """Multiscale conditional KL, identical to energy.relation_field_energy."""
+ diagonal = torch.eye(
+ len(visual_relation), dtype=torch.bool, device=visual_relation.device
+ )
+ total = visual_relation.new_zeros(())
+ for temperature in TEMPERATURES:
+ visual_logits = (visual_relation / temperature).masked_fill(diagonal, -1e4)
+ text_logits = (text_relation / temperature).masked_fill(diagonal, -1e4)
+ visual_probability = F.softmax(visual_logits, dim=-1)
+ total = total + (
+ visual_probability
+ * (
+ F.log_softmax(visual_logits, dim=-1)
+ - F.log_softmax(text_logits, dim=-1)
+ )
+ ).sum(-1).mean()
+ return total
+
+
+def permuted(relation: torch.Tensor, permutation: torch.Tensor) -> torch.Tensor:
+ return relation[permutation][:, permutation]
+
+
+def assignment_energy(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ text_relation: torch.Tensor,
+ visual_relation: torch.Tensor,
+ permutation: torch.Tensor,
+) -> dict[str, float]:
+ """Exact energy of one assignment. Channel 0 is the scalar relation."""
+ mse_channels = [
+ float(relation_mse(permuted(text_channels[c], permutation), visual_channels[c]))
+ for c in range(len(text_channels))
+ ]
+ kl = float(conditional_kl(permuted(text_relation, permutation), visual_relation))
+ mse = mse_channels[0]
+ return {
+ "mse": mse,
+ "mse_channels": mse_channels,
+ "mse_channel_mean": sum(mse_channels) / len(mse_channels),
+ "conditional_kl": kl,
+ "m30_total": M30_RELATION_WEIGHT * mse + M30_CONDITIONAL_WEIGHT * kl,
+ }
+
+
+def all_transposition_delta_mse(
+ text_standardized: torch.Tensor, visual_standardized: torch.Tensor
+) -> torch.Tensor:
+ """Exact MSE change for every transposition of the current assignment.
+
+ Swapping nodes p and q changes rows/columns p and q of the permuted text
+ relation. In the squared error the quadratic text terms cancel, leaving
+ delta(p, q) = (4 / M) * sum_{k not in {p, q}}
+ (T_pk - T_qk)(V_pk - V_qk),
+ with M the off-diagonal count and both matrices standardized with zeroed
+ diagonals. One matrix product evaluates the sum for all pairs.
+ """
+ size = len(text_standardized)
+ count = size * (size - 1)
+ cross = text_standardized @ visual_standardized # (T V)_pq
+ self_terms = (text_standardized * visual_standardized).sum(-1) # s_i
+ corrections = 2.0 * text_standardized * visual_standardized # k in {p, q}
+ total = self_terms[:, None] + self_terms[None, :] - cross - cross.T - corrections
+ delta = (4.0 / count) * total
+ delta.fill_diagonal_(0.0)
+ return delta
+
+
+def sum_channel_delta(
+ text_channels: torch.Tensor, visual_channels: torch.Tensor
+) -> torch.Tensor:
+ delta = all_transposition_delta_mse(text_channels[0], visual_channels[0])
+ for c in range(1, len(text_channels)):
+ delta = delta + all_transposition_delta_mse(
+ text_channels[c], visual_channels[c]
+ )
+ return delta / len(text_channels)
+
+
+def random_permutations(
+ count: int, size: int, generator: torch.Generator
+) -> torch.Tensor:
+ return torch.argsort(torch.rand(count, size, generator=generator), dim=-1)
+
+
+def k_derangement(
+ size: int, k: int, generator: torch.Generator
+) -> torch.Tensor:
+ """Identity with a random cyclic derangement on k random positions."""
+ permutation = torch.arange(size)
+ chosen = torch.randperm(size, generator=generator)[:k]
+ permutation[chosen] = chosen.roll(1)
+ return permutation
+
+
+def gate_a_global_ranking(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ text_relation: torch.Tensor,
+ visual_relation: torch.Tensor,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+) -> dict:
+ size = len(visual_relation)
+ identity = torch.arange(size)
+ true_energy = assignment_energy(
+ text_channels, visual_channels, text_relation, visual_relation, identity
+ )
+ keys = ("mse", "mse_channel_mean", "conditional_kl", "m30_total")
+ samples: dict[str, list[float]] = {key: [] for key in keys}
+ for index in range(args.random_perms):
+ permutation = random_permutations(1, size, generator)[0]
+ energy = assignment_energy(
+ text_channels, visual_channels, text_relation, visual_relation, permutation
+ )
+ for key in keys:
+ samples[key].append(energy[key])
+ # Structured negative: cyclic shift along the text-similarity order, a
+ # systematic misassignment that preserves neighborhood smoothness.
+ order = text_relation.sum(-1).argsort()
+ shift = torch.empty_like(order)
+ shift[order] = order.roll(1)
+ shifted_energy = assignment_energy(
+ text_channels, visual_channels, text_relation, visual_relation, shift
+ )
+ report: dict = {
+ "true": true_energy,
+ "similarity_shift": shifted_energy,
+ "random": {},
+ }
+ for key in keys:
+ values = torch.tensor(samples[key])
+ z = (values.mean() - true_energy[key]) / values.std().clamp_min(1e-12)
+ rank = int((values <= true_energy[key]).sum())
+ report["random"][key] = {
+ "mean": float(values.mean()),
+ "std": float(values.std()),
+ "min": float(values.min()),
+ "true_z": float(z),
+ "true_rank_among_random": rank,
+ "count": args.random_perms,
+ }
+ return report
+
+
+def gate_b_transpositions(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ text_relation: torch.Tensor,
+ visual_relation: torch.Tensor,
+ captions: list[str] | None,
+) -> dict:
+ size = len(visual_relation)
+ delta = sum_channel_delta(text_channels, visual_channels)
+ upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1)
+ values = delta[upper]
+ improving = values < 0
+ report: dict = {
+ "pairs": int(values.numel()),
+ "improving_pairs": int(improving.sum()),
+ "improving_fraction": float(improving.double().mean()),
+ "delta_mean": float(values.mean()),
+ "delta_min": float(values.min()),
+ "identity_is_local_min_mse": bool(improving.sum() == 0),
+ }
+ if improving.any():
+ flat = delta.masked_fill(~upper, float("inf")).flatten()
+ worst = flat.argsort()[:20]
+ offenders = []
+ for index in worst.tolist():
+ p, q = divmod(index, size)
+ if flat[index] == float("inf"):
+ break
+ exact = {
+ "pair": [p, q],
+ "delta_mse": float(delta[p, q]),
+ "text_cosine": float(text_relation[p, q]),
+ "visual_cosine": float(visual_relation[p, q]),
+ }
+ if captions is not None:
+ exact["captions"] = [captions[p][:90], captions[q][:90]]
+ offenders.append(exact)
+ report["worst_improving_swaps"] = offenders
+ return report
+
+
+def derangement_curve(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+) -> list[dict]:
+ size = len(text_channels[0])
+ identity_mse = float(
+ relation_mse(text_channels[0], visual_channels[0])
+ )
+ curve = []
+ k = 2
+ while k <= size:
+ deltas = []
+ for _ in range(args.derangement_samples):
+ permutation = k_derangement(size, k, generator)
+ mse = float(
+ relation_mse(
+ permuted(text_channels[0], permutation), visual_channels[0]
+ )
+ )
+ deltas.append(mse - identity_mse)
+ values = torch.tensor(deltas)
+ curve.append(
+ {
+ "k": k,
+ "delta_mean": float(values.mean()),
+ "delta_std": float(values.std()),
+ "improving_fraction": float((values < 0).double().mean()),
+ }
+ )
+ k *= 2
+ return curve
+
+
+def steepest_descent(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ text_relation: torch.Tensor,
+ visual_relation: torch.Tensor,
+ start: torch.Tensor,
+ args: argparse.Namespace,
+) -> dict:
+ """Exact steepest 2-swap descent with a local-minimum certificate.
+
+ Every step evaluates the closed-form MSE delta of all transpositions of
+ the current assignment. With the m30_total objective the best candidates
+ by MSE delta are re-scored with the exact weighted MSE+KL objective, so
+ an accepted move always lowers the reported objective.
+ """
+ permutation = start.clone()
+ identity = torch.arange(len(start))
+ trajectory = []
+
+ def objective(perm: torch.Tensor) -> float:
+ energy = assignment_energy(
+ text_channels, visual_channels, text_relation, visual_relation, perm
+ )
+ return energy["m30_total" if args.descent_objective == "m30_total" else "mse"]
+
+ current = objective(permutation)
+ accepted_moves = 0
+ for step in range(args.descent_max_steps):
+ perm_text = torch.stack(
+ [permuted(channel, permutation) for channel in text_channels]
+ )
+ delta = sum_channel_delta(perm_text, visual_channels)
+ upper = torch.triu(torch.ones_like(delta, dtype=torch.bool), diagonal=1)
+ masked = delta.masked_fill(~upper, float("inf"))
+ if args.descent_objective == "mse":
+ best = masked.flatten().argmin()
+ p, q = divmod(int(best), len(permutation))
+ if masked[p, q] >= 0:
+ break
+ permutation[[p, q]] = permutation[[q, p]]
+ current = objective(permutation)
+ accepted_moves += 1
+ else:
+ candidates = masked.flatten().argsort()[: args.descent_verify_top]
+ accepted = False
+ for index in candidates.tolist():
+ p, q = divmod(index, len(permutation))
+ if masked[p, q] == float("inf"):
+ break
+ trial = permutation.clone()
+ trial[[p, q]] = trial[[q, p]]
+ value = objective(trial)
+ if value < current - 1e-12:
+ permutation = trial
+ current = value
+ accepted = True
+ accepted_moves += 1
+ break
+ if not accepted:
+ break
+ if step % 50 == 0:
+ trajectory.append(
+ {
+ "step": step,
+ "objective": current,
+ "accuracy": float((permutation == identity).double().mean()),
+ }
+ )
+ final_delta = sum_channel_delta(
+ torch.stack([permuted(channel, permutation) for channel in text_channels]),
+ visual_channels,
+ )
+ upper = torch.triu(torch.ones_like(final_delta, dtype=torch.bool), diagonal=1)
+ certificate = bool((final_delta[upper] >= 0).all())
+ return {
+ "start_accuracy": float((start == identity).double().mean()),
+ "final_accuracy": float((permutation == identity).double().mean()),
+ "final_objective": current,
+ "final_energy": assignment_energy(
+ text_channels, visual_channels, text_relation, visual_relation, permutation
+ ),
+ "accepted_moves": accepted_moves,
+ "moved_fraction": float((permutation != start).double().mean()),
+ "mse_local_min_certificate": certificate,
+ "trajectory": trajectory,
+ "final_permutation": permutation.tolist(),
+ }
+
+
+def load_flickr(args: argparse.Namespace) -> dict:
+ manifest = read_json(args.manifest)
+ vision, text, vision_lookup, text_lookup = load_feature_pair(
+ args.vision, args.text
+ )
+ rows = manifest[args.split][: args.samples]
+ visual_states = select_rows(vision["features"], vision_lookup, rows)
+ captions = None
+ if args.text_mode == "single":
+ text_states = select_rows(text["features"], text_lookup, rows)
+ captions = [
+ text["captions"][text_lookup[int(row)]] for row in rows
+ ]
+ else:
+ state = torch.load(args.text_orbits, map_location="cpu", weights_only=False)
+ lookup = {int(row): index for index, row in enumerate(state["rows"])}
+ orbit_mean = F.normalize(state["features"].float().mean(1), dim=-1)
+ text_states = select_rows(orbit_mean, lookup, rows)
+ captions = [state["captions"][lookup[int(row)]][0] for row in rows]
+ return {
+ "visual_views": visual_states[:, None, :],
+ "text_views": text_states[:, None, :],
+ "captions": captions,
+ "meta": {
+ "dataset": "flickr30k",
+ "split": args.split,
+ "samples": len(rows),
+ "text_mode": args.text_mode,
+ "rows": rows,
+ },
+ }
+
+
+def load_vg(args: argparse.Namespace) -> dict:
+ vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False)
+ text = torch.load(args.vg_text, map_location="cpu", weights_only=False)
+ pairs = [
+ json.loads(line)
+ for line in Path(args.vg_ground_truth).read_text().splitlines()
+ if line.strip()
+ ]
+ vision_index = {node: i for i, node in enumerate(vision["node_ids"])}
+ text_index = {node: i for i, node in enumerate(text["node_ids"])}
+ vision_order = [vision_index[pair["vision_node_id"]] for pair in pairs]
+ text_order = [text_index[pair["text_node_id"]] for pair in pairs]
+ visual_views = F.normalize(vision["region_features"].float(), dim=-1)[vision_order]
+ text_views = F.normalize(text["region_features"].float(), dim=-1)[text_order]
+ if args.samples and args.samples < len(visual_views):
+ generator = torch.Generator().manual_seed(args.subset_seed)
+ subset = torch.randperm(len(visual_views), generator=generator)[: args.samples]
+ visual_views = visual_views[subset]
+ text_views = text_views[subset]
+ return {
+ "visual_views": visual_views,
+ "text_views": text_views,
+ "captions": None,
+ "meta": {
+ "dataset": "visual_genome_5k",
+ "tier": text.get("tier"),
+ "samples": len(visual_views),
+ "subset_seed": args.subset_seed,
+ "bundle_channels": bool(args.vg_bundle_channels),
+ },
+ }
+
+
+def view_bundle_channels(views: torch.Tensor) -> torch.Tensor:
+ """Distribution-valued relation field from per-node view sets.
+
+ Channel order: mean, std, q10, q90 of the view-pair cosine distribution
+ between two nodes. The scalar mean channel equals the relation of the
+ (unnormalized) view-mean embeddings; the remaining channels carry
+ information a single pooled vector cannot.
+ """
+ nodes, view_count, _ = views.shape
+ views = views.double()
+ pair_cosines = torch.einsum("aud,bvd->abuv", views, views).reshape(
+ nodes, nodes, view_count * view_count
+ )
+ mean = pair_cosines.mean(-1)
+ std = pair_cosines.std(-1)
+ q10 = pair_cosines.quantile(0.10, dim=-1)
+ q90 = pair_cosines.quantile(0.90, dim=-1)
+ return torch.stack([mean, std, q10, q90])
+
+
+def build_channels(
+ views: torch.Tensor, bundle: bool
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Standardized relation channels and the raw scalar relation."""
+ node_states = F.normalize(views.double().mean(1), dim=-1)
+ scalar = node_states @ node_states.T
+ if bundle and views.shape[1] > 1:
+ raw = view_bundle_channels(views)
+ else:
+ raw = scalar[None]
+ channels = []
+ for c in range(len(raw)):
+ standardized, _, _ = standardize_relation(raw[c])
+ channels.append(standardized)
+ return torch.stack(channels), scalar
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ data = load_flickr(args) if args.dataset == "flickr" else load_vg(args)
+ device = torch.device(args.device)
+ bundle = args.dataset == "vg" and args.vg_bundle_channels
+ text_channels, text_relation = build_channels(
+ data["text_views"].to(device), bundle
+ )
+ visual_channels, visual_relation = build_channels(
+ data["visual_views"].to(device), bundle
+ )
+ generator = torch.Generator().manual_seed(args.seed)
+
+ report: dict = {
+ "protocol": (
+ "Assignments permute real frozen text states; hidden pairs are "
+ "used only to place the true assignment in the ranking. The "
+ "energy terms are the cross-modal relation MSE and conditional "
+ "KL; the language-only terms of the falsified free-particle "
+ "energy are permutation-invariant on this space."
+ ),
+ "meta": data["meta"],
+ "args": {
+ key: value
+ for key, value in vars(args).items()
+ if key not in ("manifest", "vision", "text")
+ },
+ "channel_names": (
+ ["mean", "std", "q10", "q90"] if bundle else ["mean"]
+ ),
+ }
+
+ report["gate_a_global_ranking"] = gate_a_global_ranking(
+ text_channels, visual_channels, text_relation, visual_relation, args, generator
+ )
+ print(json.dumps({"gate_a": report["gate_a_global_ranking"]["random"]}))
+
+ report["gate_b_transpositions"] = gate_b_transpositions(
+ text_channels, visual_channels, text_relation, visual_relation, data["captions"]
+ )
+ print(
+ json.dumps(
+ {
+ "gate_b": {
+ key: value
+ for key, value in report["gate_b_transpositions"].items()
+ if key != "worst_improving_swaps"
+ }
+ }
+ )
+ )
+
+ report["derangement_curve"] = derangement_curve(
+ text_channels, visual_channels, args, generator
+ )
+
+ identity = torch.arange(len(visual_relation))
+ report["gate_c_descent_from_true"] = steepest_descent(
+ text_channels,
+ visual_channels,
+ text_relation,
+ visual_relation,
+ identity,
+ args,
+ )
+ print(
+ json.dumps(
+ {
+ "gate_c_from_true": {
+ key: value
+ for key, value in report["gate_c_descent_from_true"].items()
+ if key not in ("trajectory", "final_permutation")
+ }
+ }
+ )
+ )
+
+ restarts = []
+ for restart in range(args.descent_restarts):
+ start = random_permutations(1, len(visual_relation), generator)[0]
+ result = steepest_descent(
+ text_channels,
+ visual_channels,
+ text_relation,
+ visual_relation,
+ start,
+ args,
+ )
+ result.pop("final_permutation")
+ restarts.append(result)
+ print(
+ json.dumps(
+ {
+ "gate_c_from_random": {
+ "restart": restart,
+ "final_objective": result["final_objective"],
+ "final_accuracy": result["final_accuracy"],
+ }
+ }
+ )
+ )
+ report["gate_c_descent_from_random"] = restarts
+
+ true_total = report["gate_a_global_ranking"]["true"]["m30_total"]
+ counterfeit = [
+ restart
+ for restart in restarts
+ if restart["final_energy"]["m30_total"] < true_total
+ and restart["final_accuracy"] < 0.5
+ ]
+ report["verdict"] = {
+ "true_m30_total": true_total,
+ "identity_is_local_min_mse": report["gate_b_transpositions"][
+ "identity_is_local_min_mse"
+ ],
+ "descent_from_true_stays": report["gate_c_descent_from_true"][
+ "final_accuracy"
+ ],
+ "on_manifold_counterfeit_found": bool(counterfeit),
+ "best_random_descent_m30_total": min(
+ (restart["final_energy"]["m30_total"] for restart in restarts),
+ default=None,
+ ),
+ }
+ print(json.dumps({"verdict": report["verdict"]}))
+
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, report)
+ if args.trajectory_output:
+ torch.save(
+ {
+ "from_true": report["gate_c_descent_from_true"],
+ "meta": data["meta"],
+ },
+ args.trajectory_output,
+ )
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/models.py b/worldalign/models.py
new file mode 100644
index 0000000..5f98605
--- /dev/null
+++ b/worldalign/models.py
@@ -0,0 +1,90 @@
+from __future__ import annotations
+
+import torch
+from torch import nn
+import torch.nn.functional as F
+
+
+class Bridge(nn.Module):
+ def __init__(
+ self,
+ vision_dim: int,
+ text_dim: int,
+ hidden_dim: int = 1536,
+ linear: bool = False,
+ ):
+ super().__init__()
+ self.vision_dim = vision_dim
+ self.text_dim = text_dim
+ self.hidden_dim = hidden_dim
+ self.linear = linear
+ if linear:
+ self.network = nn.Linear(vision_dim, text_dim, bias=False)
+ else:
+ self.network = nn.Sequential(
+ nn.LayerNorm(vision_dim),
+ nn.Linear(vision_dim, hidden_dim),
+ nn.GELU(),
+ nn.Linear(hidden_dim, text_dim),
+ )
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return F.normalize(self.network(x).float(), dim=-1)
+
+ def config(self) -> dict:
+ return {
+ "vision_dim": self.vision_dim,
+ "text_dim": self.text_dim,
+ "hidden_dim": self.hidden_dim,
+ "linear": self.linear,
+ }
+
+
+class PrefixAdapter(nn.Module):
+ def __init__(
+ self,
+ semantic_dim: int,
+ lm_dim: int,
+ prefix_length: int = 8,
+ hidden_dim: int = 2048,
+ ):
+ super().__init__()
+ self.semantic_dim = semantic_dim
+ self.lm_dim = lm_dim
+ self.prefix_length = prefix_length
+ self.hidden_dim = hidden_dim
+ self.network = nn.Sequential(
+ nn.LayerNorm(semantic_dim),
+ nn.Linear(semantic_dim, hidden_dim),
+ nn.GELU(),
+ nn.Linear(hidden_dim, prefix_length * lm_dim),
+ )
+
+ def forward(self, semantic: torch.Tensor) -> torch.Tensor:
+ prefix = self.network(semantic.float())
+ return prefix.view(-1, self.prefix_length, self.lm_dim)
+
+ def config(self) -> dict:
+ return {
+ "semantic_dim": self.semantic_dim,
+ "lm_dim": self.lm_dim,
+ "prefix_length": self.prefix_length,
+ "hidden_dim": self.hidden_dim,
+ }
+
+
+def load_bridge(path: str, device: str = "cpu") -> tuple[Bridge, dict]:
+ state = torch.load(path, map_location="cpu", weights_only=False)
+ model = Bridge(**state["config"])
+ model.load_state_dict(state["state_dict"])
+ model.to(device).eval()
+ return model, state
+
+
+def load_prefix(path: str, device: str = "cpu") -> tuple[PrefixAdapter, dict]:
+ state = torch.load(path, map_location="cpu", weights_only=False)
+ model = PrefixAdapter(**state["config"])
+ model.load_state_dict(state["state_dict"])
+ model.to(device).eval()
+ return model, state
+
diff --git a/worldalign/natural_families.py b/worldalign/natural_families.py
new file mode 100644
index 0000000..a13a1d9
--- /dev/null
+++ b/worldalign/natural_families.py
@@ -0,0 +1,264 @@
+"""Do factor families survive real language?
+
+The synthetic derivation found colour, count, and size families because
+templated phrases place exactly one member of each family in a fixed
+slot, so family members never co-occur. Real descriptions break every
+part of that: free word order, stacked adjectives, synonyms, and phrases
+that mention no attribute at all. This measures how much of the
+mutual-exclusivity signal survives, on Visual Genome region descriptions,
+using no lexicon and no labels.
+
+Families are recovered as low-co-occurrence, high-context-similarity
+groups: two words of one factor rarely modify the same head, and when
+they do appear they appear in the same distributional company. That is
+the paradigmatic relation of distributional semantics, and the greedy
+exclusivity pass of the synthetic pipeline is its degenerate case.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from collections import Counter, defaultdict
+from pathlib import Path
+
+import numpy as np
+
+from .common import read_json, seed_everything, write_json
+
+# Reference sets, used only to score the recovered families, never to find
+# them. Membership is checked after the fact.
+REFERENCE = {
+ "colour": {
+ "black", "white", "red", "blue", "green", "yellow", "brown", "gray",
+ "grey", "orange", "purple", "pink", "tan", "beige", "silver", "gold",
+ "golden", "dark", "light",
+ },
+ "number": {
+ "one", "two", "three", "four", "five", "six", "seven", "eight",
+ "nine", "ten", "a", "an", "the", "some", "many",
+ },
+ "size": {"small", "large", "big", "little", "tiny", "huge", "tall", "short", "long"},
+ "material": {"wooden", "metal", "plastic", "glass", "brick", "stone", "leather"},
+}
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--nodes", default="artifacts/vg_5k/text_nodes.jsonl")
+ parser.add_argument("--field", default="region_closed")
+ parser.add_argument("--min-count", type=int, default=200)
+ parser.add_argument("--max-vocabulary", type=int, default=300)
+ parser.add_argument("--context-window", type=int, default=2)
+ parser.add_argument("--exclusivity", type=float, default=0.25,
+ help="Maximum observed-over-expected co-occurrence for "
+ "two words to count as mutually exclusive.")
+ parser.add_argument(
+ "--method",
+ choices=["exclusivity", "distributional"],
+ default="distributional",
+ help="Mutual exclusivity is a template artifact; real paradigms are "
+ "found by distributional similarity, of which it is a special case.",
+ )
+ parser.add_argument("--components", type=int, default=64)
+ parser.add_argument("--clusters", type=int, default=30)
+ parser.add_argument("--min-family", type=int, default=3)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument(
+ "--output", default="artifacts/vg_5k/natural_families.json"
+ )
+ return parser.parse_args()
+
+
+def load_phrases(path: Path, field: str) -> list[list[str]]:
+ phrases = []
+ for line in path.read_text(encoding="utf-8").splitlines():
+ if not line.strip():
+ continue
+ record = json.loads(line)
+ for phrase in record[field]:
+ tokens = re.findall(r"[a-z]+", phrase.lower())
+ if tokens:
+ phrases.append(tokens)
+ return phrases
+
+
+def statistics(
+ phrases: list[list[str]], vocabulary: set[str], window: int
+) -> tuple[Counter, Counter, dict[str, Counter]]:
+ unigram: Counter = Counter()
+ pair: Counter = Counter()
+ context: dict[str, Counter] = defaultdict(Counter)
+ for tokens in phrases:
+ present = [token for token in tokens if token in vocabulary]
+ for token in set(present):
+ unigram[token] += 1
+ for first in set(present):
+ for second in set(present):
+ if first < second:
+ pair[(first, second)] += 1
+ for index, token in enumerate(tokens):
+ if token not in vocabulary:
+ continue
+ for offset in range(1, window + 1):
+ for neighbour_index in (index - offset, index + offset):
+ if 0 <= neighbour_index < len(tokens):
+ context[token][tokens[neighbour_index]] += 1
+ return unigram, pair, context
+
+
+def exclusivity_ratio(
+ first: str, second: str, unigram: Counter, pair: Counter, total: int
+) -> float:
+ """Observed co-occurrence over the independent expectation."""
+ expected = unigram[first] * unigram[second] / max(total, 1)
+ key = (first, second) if first < second else (second, first)
+ return pair[key] / max(expected, 1e-9)
+
+
+def context_similarity(first: Counter, second: Counter) -> float:
+ keys = set(first) | set(second)
+ a = np.array([first[k] for k in keys], dtype=float)
+ b = np.array([second[k] for k in keys], dtype=float)
+ a /= max(np.linalg.norm(a), 1e-9)
+ b /= max(np.linalg.norm(b), 1e-9)
+ return float(a @ b)
+
+
+def build_families(
+ words: list[str],
+ unigram: Counter,
+ pair: Counter,
+ context: dict[str, Counter],
+ total: int,
+ args: argparse.Namespace,
+) -> list[list[str]]:
+ """Greedy paradigmatic grouping: exclusive and distributionally alike."""
+ families: list[list[str]] = []
+ for word in words:
+ best_family, best_score = None, 0.0
+ for family in families:
+ ratios = [
+ exclusivity_ratio(word, member, unigram, pair, total)
+ for member in family
+ ]
+ if max(ratios) > args.exclusivity:
+ continue
+ similarity = float(
+ np.mean(
+ [context_similarity(context[word], context[member])
+ for member in family]
+ )
+ )
+ if similarity > best_score:
+ best_family, best_score = family, similarity
+ if best_family is not None and best_score > 0.15:
+ best_family.append(word)
+ else:
+ families.append([word])
+ return [family for family in families if len(family) >= args.min_family]
+
+
+def distributional_families(
+ words: list[str], context: dict[str, Counter], args: argparse.Namespace
+) -> list[list[str]]:
+ """Word classes from context vectors: positive PMI, truncated SVD, k-means.
+
+ The standard recipe of distributional word-class induction. Words of
+ one factor modify the same heads and so share company, whether or not
+ they exclude each other -- real colour terms co-occur freely.
+ """
+ from sklearn.cluster import KMeans
+
+ features = sorted({key for word in words for key in context[word]})
+ index = {key: position for position, key in enumerate(features)}
+ matrix = np.zeros((len(words), len(features)))
+ for row, word in enumerate(words):
+ for key, value in context[word].items():
+ matrix[row, index[key]] = value
+ total = matrix.sum()
+ row_sum = matrix.sum(1, keepdims=True)
+ column_sum = matrix.sum(0, keepdims=True)
+ expected = row_sum * column_sum / max(total, 1e-9)
+ pmi = np.log(np.maximum(matrix, 1e-9) / np.maximum(expected, 1e-9))
+ pmi[matrix == 0] = 0.0
+ pmi = np.maximum(pmi, 0.0)
+ left, values, _ = np.linalg.svd(pmi, full_matrices=False)
+ embedding = left[:, : args.components] * values[: args.components]
+ embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9)
+ labels = KMeans(args.clusters, n_init=10, random_state=args.seed).fit_predict(
+ embedding
+ )
+ families: dict[int, list[str]] = defaultdict(list)
+ for word, label in zip(words, labels):
+ families[int(label)].append(word)
+ return [family for family in families.values() if len(family) >= args.min_family]
+
+
+def label_family(family: list[str]) -> tuple[str, float]:
+ best_name, best_purity = "unlabelled", 0.0
+ for name, reference in REFERENCE.items():
+ purity = sum(1 for word in family if word in reference) / len(family)
+ if purity > best_purity:
+ best_name, best_purity = name, purity
+ return best_name, best_purity
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ phrases = load_phrases(Path(args.nodes), args.field)
+ counts = Counter(token for tokens in phrases for token in set(tokens))
+ vocabulary = [
+ word
+ for word, count in counts.most_common(args.max_vocabulary)
+ if count >= args.min_count
+ ]
+ unigram, pair, context = statistics(phrases, set(vocabulary), args.context_window)
+ if args.method == "distributional":
+ families = distributional_families(vocabulary, context, args)
+ else:
+ families = build_families(
+ vocabulary, unigram, pair, context, len(phrases), args
+ )
+
+ described = []
+ for family in sorted(families, key=len, reverse=True):
+ name, purity = label_family(family)
+ described.append(
+ {
+ "size": len(family),
+ "label": name,
+ "purity": purity,
+ "words": sorted(family, key=lambda w: -unigram[w])[:14],
+ }
+ )
+ report = {
+ "protocol": (
+ "Families are recovered from region descriptions alone by "
+ "mutual exclusivity plus distributional similarity. Reference "
+ "word sets are read only to label the result."
+ ),
+ "phrases": len(phrases),
+ "vocabulary": len(vocabulary),
+ "families": described,
+ "recovered_labels": Counter(item["label"] for item in described),
+ }
+ for item in described[:12]:
+ print(
+ json.dumps(
+ {
+ "label": item["label"],
+ "purity": round(item["purity"], 2),
+ "size": item["size"],
+ "words": item["words"][:10],
+ }
+ )
+ )
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/natural_fields.py b/worldalign/natural_fields.py
new file mode 100644
index 0000000..7221a3a
--- /dev/null
+++ b/worldalign/natural_fields.py
@@ -0,0 +1,224 @@
+"""Natural-data relation fields under the Tier 0 recipe.
+
+Each side is encoded into its own discovered factor coordinates -- text
+word classes from distributional induction, vision segment classes from
+feature clustering -- and the two coordinate systems are paired by joint
+structure, never by a declared lexicon. Scene states are the resulting
+sets of segment or phrase codes; relations are moment-kernel similarities
+between scenes.
+
+The field correlation at the hidden pairing is the go/no-go statistic:
+polynomial recovery needs roughly 0.9, and the synthetic world showed
+nothing works below it. Hidden pairs are read only to compute it.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from collections import Counter, defaultdict
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from scipy.optimize import linear_sum_assignment
+from sklearn.cluster import KMeans
+
+from .common import read_json, seed_everything, write_json
+from .natural_families import distributional_families, load_phrases, statistics
+from .synth_cc_battery import moment_field
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--objects", default="artifacts/vg_5k/natural_objects.pt")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--fit-nodes", type=int, default=2000)
+ parser.add_argument("--vision-classes", type=int, default=24)
+ parser.add_argument("--text-classes", type=int, default=24)
+ parser.add_argument("--clusters", type=int, default=24)
+ parser.add_argument("--min-count", type=int, default=200)
+ parser.add_argument("--max-vocabulary", type=int, default=300)
+ parser.add_argument("--context-window", type=int, default=2)
+ parser.add_argument("--min-family", type=int, default=3)
+ parser.add_argument("--components", type=int, default=64)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="artifacts/vg_5k/natural_fields.pt")
+ return parser.parse_args()
+
+
+def text_scene_codes(
+ records: list[dict], word_class: dict[str, int], classes: int
+) -> dict[str, torch.Tensor]:
+ """One code vector per region description: its word-class profile."""
+ codes: dict[str, list[torch.Tensor]] = {}
+ for record in records:
+ vectors = []
+ for phrase in record["region_closed"]:
+ vector = torch.zeros(classes)
+ for token in re.findall(r"[a-z]+", phrase.lower()):
+ if token in word_class:
+ vector[word_class[token]] += 1.0
+ if vector.sum() > 0:
+ vectors.append(vector)
+ if vectors:
+ codes[record["node_id"]] = F.normalize(torch.stack(vectors), dim=-1)
+ return codes
+
+
+def vision_scene_codes(
+ state: dict, model: KMeans, classes: int
+) -> dict[str, torch.Tensor]:
+ """One code vector per segment: its feature-class assignment."""
+ codes: dict[str, torch.Tensor] = {}
+ for node, segments in zip(state["node_ids"], state["segments"]):
+ if not segments:
+ continue
+ labels = model.predict(
+ np.stack([segment["feature"] for segment in segments]).astype(np.float64)
+ )
+ vectors = torch.zeros(len(segments), classes)
+ for row, label in enumerate(labels):
+ vectors[row, int(label)] = 1.0
+ codes[node] = F.normalize(vectors, dim=-1)
+ return codes
+
+
+def align_classes(
+ text_codes: dict[str, torch.Tensor],
+ vision_codes: dict[str, torch.Tensor],
+ classes: int,
+) -> np.ndarray:
+ """Pair vision classes to text classes by marginal frequency rank.
+
+ Within-scene co-occurrence is the stronger signal but needs a second
+ factor to condition on; frequency is the available unimodal statistic
+ at this stage and is reported as the first pass.
+ """
+ text_mass = torch.zeros(classes)
+ for code in text_codes.values():
+ text_mass += code.sum(0)
+ vision_mass = torch.zeros(classes)
+ for code in vision_codes.values():
+ vision_mass += code.sum(0)
+ text_order = torch.argsort(text_mass, descending=True).numpy()
+ vision_order = torch.argsort(vision_mass, descending=True).numpy()
+ mapping = np.empty(classes, dtype=int)
+ mapping[vision_order] = text_order
+ return mapping
+
+
+def moment_states(codes: torch.Tensor) -> torch.Tensor:
+ first = codes.mean(0)
+ second = (codes[:, :, None] * codes[:, None, :]).mean(0).flatten()
+ return torch.cat([first, second])
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ vg_dir = Path(args.vg_dir)
+ text_records = [
+ json.loads(line)
+ for line in (vg_dir / "text_nodes.jsonl").read_text(encoding="utf-8").splitlines()
+ if line.strip()
+ ]
+ truth = [
+ json.loads(line)
+ for line in (vg_dir / "ground_truth.private.jsonl")
+ .read_text(encoding="utf-8")
+ .splitlines()
+ if line.strip()
+ ]
+ state = torch.load(args.objects, map_location="cpu", weights_only=False)
+
+ phrases = load_phrases(vg_dir / "text_nodes.jsonl", "region_closed")
+ counts = Counter(token for tokens in phrases for token in set(tokens))
+ vocabulary = [
+ word
+ for word, count in counts.most_common(args.max_vocabulary)
+ if count >= args.min_count
+ ]
+ _, _, context = statistics(phrases, set(vocabulary), args.context_window)
+ args.clusters = args.text_classes
+ families = distributional_families(vocabulary, context, args)
+ word_class = {
+ word: index for index, family in enumerate(families) for word in family
+ }
+ text_classes = len(families)
+
+ features = np.stack(
+ [
+ segment["feature"]
+ for segments in state["segments"][: args.fit_nodes]
+ for segment in segments
+ ]
+ ).astype(np.float64)
+ vision_model = KMeans(
+ args.vision_classes, n_init=10, random_state=args.seed
+ ).fit(features)
+
+ text_codes = text_scene_codes(text_records, word_class, text_classes)
+ vision_codes = vision_scene_codes(state, vision_model, args.vision_classes)
+
+ shared = min(text_classes, args.vision_classes)
+ mapping = align_classes(
+ {k: v[:, :shared] for k, v in text_codes.items()},
+ {k: v[:, :shared] for k, v in vision_codes.items()},
+ shared,
+ )
+
+ pairs = [
+ pair
+ for pair in truth
+ if pair["vision_node_id"] in vision_codes and pair["text_node_id"] in text_codes
+ ][: args.samples]
+ vision_sets, text_sets = [], []
+ for pair in pairs:
+ vision = vision_codes[pair["vision_node_id"]][:, :shared]
+ remapped = torch.zeros_like(vision)
+ for source in range(shared):
+ remapped[:, mapping[source]] = vision[:, source]
+ vision_sets.append(F.normalize(remapped, dim=-1))
+ text_sets.append(
+ F.normalize(text_codes[pair["text_node_id"]][:, :shared], dim=-1)
+ )
+
+ visual_field = moment_field(vision_sets)
+ text_field = moment_field(text_sets)
+ size = len(pairs)
+ mask = ~np.eye(size, dtype=bool)
+ correlation = float(
+ np.corrcoef(
+ visual_field.double().numpy()[mask], text_field.double().numpy()[mask]
+ )[0, 1]
+ )
+ torch.save(
+ {
+ "visual_field": visual_field,
+ "text_field": text_field,
+ "pairs": [p["vision_node_id"] for p in pairs],
+ },
+ args.output,
+ )
+ summary = {
+ "protocol": (
+ "Text classes from distributional induction, vision classes "
+ "from segment-feature clustering, paired by marginal frequency "
+ "rank. Hidden pairs are read only for the correlation."
+ ),
+ "samples": size,
+ "text_classes": text_classes,
+ "vision_classes": args.vision_classes,
+ "field_correlation_at_truth": correlation,
+ "go_no_go": "recovery needs about 0.9",
+ }
+ print(json.dumps(summary))
+ write_json(str(args.output).replace(".pt", ".json"), summary)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/natural_objects.py b/worldalign/natural_objects.py
new file mode 100644
index 0000000..65e035b
--- /dev/null
+++ b/worldalign/natural_objects.py
@@ -0,0 +1,262 @@
+"""Unsupervised object states for natural images.
+
+The synthetic world's objects came from watershed on a flat background,
+which real photographs do not offer. Self-supervised vision features do:
+DINOv2 patch descriptors segment objects without labels, which is what
+the deep-spectral line of work exploits. Each image is cut into segments
+by spectral clustering of the patch affinity graph, and each segment is
+described by observables that a text corpus also states -- mean colour,
+relative area, position, and its own feature centroid for later category
+clustering.
+
+Nothing here reads text, pairs, or region annotations. Boxes from the
+preprocessing file are not used; segmentation is derived from pixels.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from PIL import Image
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel
+
+from .common import batch_indices, read_json, seed_everything, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images")
+ parser.add_argument("--model", default="facebook/dinov2-base")
+ parser.add_argument("--image-size", type=int, default=448)
+ parser.add_argument("--segments", type=int, default=6)
+ parser.add_argument("--min-patches", type=int, default=8)
+ parser.add_argument(
+ "--oracle-regions",
+ action="store_true",
+ help="Diagnostic upper bound: take segments from the annotated "
+ "region boxes instead of unsupervised segmentation, isolating "
+ "how much of the field deficit segmentation accounts for.",
+ )
+ parser.add_argument(
+ "--views",
+ type=int,
+ default=1,
+ help="Augmentation orbit size. Each view is a random resized crop, "
+ "so content is fixed and framing varies -- the natural-image "
+ "analogue of the synthetic world's re-renders.",
+ )
+ parser.add_argument("--crop-low", type=float, default=0.75)
+ parser.add_argument("--nodes", type=int, default=0)
+ parser.add_argument("--batch-size", type=int, default=16)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="artifacts/vg_5k/natural_objects.pt")
+ return parser.parse_args()
+
+
+def spectral_segments(
+ features: torch.Tensor, grid: int, segments: int, seed: int
+) -> np.ndarray:
+ """Segment a patch grid by clustering the normalised affinity spectrum.
+
+ Feature cosine affinity is combined with a spatial prior so segments
+ stay connected, then the leading eigenvectors of the normalised
+ Laplacian are clustered. This is the standard unsupervised recipe over
+ self-supervised features.
+ """
+ from sklearn.cluster import KMeans
+
+ normalised = F.normalize(features.double(), dim=-1)
+ affinity = (normalised @ normalised.T).clamp_min(0).numpy()
+ coordinates = np.stack(
+ np.meshgrid(np.arange(grid), np.arange(grid), indexing="ij"), -1
+ ).reshape(-1, 2).astype(np.float64)
+ distance = ((coordinates[:, None, :] - coordinates[None, :, :]) ** 2).sum(-1)
+ affinity = affinity * np.exp(-distance / (2 * (grid / 4.0) ** 2))
+ degree = affinity.sum(1)
+ laplacian = affinity / np.sqrt(np.outer(degree, degree) + 1e-9)
+ values, vectors = np.linalg.eigh(laplacian)
+ embedding = vectors[:, -segments:]
+ embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9)
+ return KMeans(segments, n_init=10, random_state=seed).fit_predict(embedding)
+
+
+def describe_segments(
+ labels: np.ndarray,
+ features: torch.Tensor,
+ pixels: torch.Tensor,
+ grid: int,
+ min_patches: int,
+) -> list[dict]:
+ """Observables of each segment: colour, area, position, feature centre."""
+ image = pixels.permute(1, 2, 0).numpy()
+ size = image.shape[0]
+ scale = size // grid
+ described = []
+ for label in np.unique(labels):
+ mask = labels == label
+ if mask.sum() < min_patches:
+ continue
+ rows, cols = np.nonzero(mask.reshape(grid, grid))
+ pixel_mask = np.zeros((size, size), dtype=bool)
+ for row, col in zip(rows, cols):
+ pixel_mask[
+ row * scale : (row + 1) * scale, col * scale : (col + 1) * scale
+ ] = True
+ described.append(
+ {
+ "rgb": image[pixel_mask].mean(0),
+ "area": float(mask.mean()),
+ "centre": np.array([rows.mean() / grid, cols.mean() / grid]),
+ "feature": features[mask].mean(0).numpy(),
+ }
+ )
+ return described
+
+
+def region_segments(
+ record: dict, features: torch.Tensor, pixels: torch.Tensor, grid: int
+) -> list[dict]:
+ """Segments taken from annotated boxes: an oracle for segmentation only."""
+ image = pixels.permute(1, 2, 0).numpy()
+ size = image.shape[0]
+ width, height = record["width"], record["height"]
+ described = []
+ for region in record["regions"]:
+ x0 = max(0.0, min(region["x"] / width, 1.0)) * grid
+ x1 = max(0.0, min((region["x"] + region["width"]) / width, 1.0)) * grid
+ y0 = max(0.0, min(region["y"] / height, 1.0)) * grid
+ y1 = max(0.0, min((region["y"] + region["height"]) / height, 1.0)) * grid
+ columns = np.arange(grid) + 0.5
+ inside_x = (columns >= x0) & (columns <= x1)
+ inside_y = (columns >= y0) & (columns <= y1)
+ mask = (inside_y[:, None] & inside_x[None, :]).flatten()
+ if not mask.any():
+ centre_x = min(grid - 1, max(0, int((x0 + x1) / 2)))
+ centre_y = min(grid - 1, max(0, int((y0 + y1) / 2)))
+ mask = np.zeros(grid * grid, dtype=bool)
+ mask[centre_y * grid + centre_x] = True
+ rows, cols = np.nonzero(mask.reshape(grid, grid))
+ scale = size // grid
+ pixel_mask = np.zeros((size, size), dtype=bool)
+ for row, col in zip(rows, cols):
+ pixel_mask[
+ row * scale : (row + 1) * scale, col * scale : (col + 1) * scale
+ ] = True
+ described.append(
+ {
+ "rgb": image[pixel_mask].mean(0),
+ "area": float(mask.mean()),
+ "centre": np.array([rows.mean() / grid, cols.mean() / grid]),
+ "feature": features[mask].mean(0).numpy(),
+ }
+ )
+ return described
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ records = [
+ json.loads(line)
+ for line in Path(args.vg_dir, "vision_nodes.private.jsonl")
+ .read_text(encoding="utf-8")
+ .splitlines()
+ if line.strip()
+ ]
+ if args.nodes:
+ records = records[: args.nodes]
+ processor = AutoImageProcessor.from_pretrained(args.model)
+ model = AutoModel.from_pretrained(args.model, torch_dtype=torch.float32).to(
+ args.device
+ )
+ model.eval()
+ patch = model.config.patch_size
+ grid = args.image_size // patch
+ mean = torch.tensor(processor.image_mean).view(3, 1, 1)
+ std = torch.tensor(processor.image_std).view(3, 1, 1)
+
+ generator = np.random.default_rng(args.seed)
+
+ def load(record: dict, view: int = 0) -> torch.Tensor:
+ path = Path(args.image_cache, f"{record['source_image_id']}.jpg")
+ with Image.open(path) as image:
+ picture = image.convert("RGB")
+ if view > 0:
+ width, height = picture.size
+ scale = generator.uniform(args.crop_low, 1.0)
+ box_w, box_h = int(width * scale), int(height * scale)
+ left = int(generator.integers(0, max(width - box_w, 1)))
+ top = int(generator.integers(0, max(height - box_h, 1)))
+ picture = picture.crop((left, top, left + box_w, top + box_h))
+ resized = picture.resize(
+ (args.image_size, args.image_size), Image.BILINEAR
+ )
+ return torch.from_numpy(np.asarray(resized).copy()).permute(2, 0, 1).float() / 255.0
+
+ node_ids, all_segments = [], []
+ view_segments: list[list] = [[] for _ in range(args.views)]
+ for indices in tqdm(
+ list(batch_indices(len(records), args.batch_size)), desc="segment"
+ ):
+ for view in range(args.views):
+ batch = [records[index] for index in indices]
+ raw = torch.stack([load(record, view) for record in batch])
+ normalised = ((raw - mean) / std).to(args.device)
+ hidden = model(pixel_values=normalised, return_dict=True).last_hidden_state
+ patches = hidden[:, 1:].float().cpu()
+ for position, record in enumerate(batch):
+ if args.oracle_regions:
+ described = region_segments(
+ record, patches[position], raw[position], grid
+ )
+ else:
+ labels = spectral_segments(
+ patches[position], grid, args.segments, args.seed
+ )
+ described = describe_segments(
+ labels, patches[position], raw[position], grid, args.min_patches
+ )
+ if view == 0:
+ node_ids.append(record["node_id"])
+ all_segments.append(described)
+ view_segments[view].append(described)
+
+ torch.save(
+ {
+ "model": args.model,
+ "node_ids": node_ids,
+ "segments": all_segments,
+ "view_segments": view_segments if args.views > 1 else None,
+ "views": args.views,
+ "grid": grid,
+ "protocol": (
+ "Segments come from spectral clustering of self-supervised "
+ "patch features; no boxes, no text, no pairs."
+ ),
+ },
+ args.output,
+ )
+ counts = [len(item) for item in all_segments]
+ print(
+ json.dumps(
+ {
+ "nodes": len(node_ids),
+ "mean_segments": float(np.mean(counts)),
+ "min_segments": int(np.min(counts)),
+ }
+ )
+ )
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/natural_pipeline.py b/worldalign/natural_pipeline.py
new file mode 100644
index 0000000..c30af4d
--- /dev/null
+++ b/worldalign/natural_pipeline.py
@@ -0,0 +1,204 @@
+"""Natural-data field pipeline: continuous states and content projection.
+
+Produces the field correlation reported for Visual Genome. Two choices
+carry most of it, both measured. States stay continuous -- quantising
+segments and phrases into class codes costs more than half the signal,
+because relation fields need scene similarity and nothing else. And each
+side is projected onto the directions that separate scenes rather than
+parts, fitted per modality on scenes outside the evaluated set.
+
+Vision states come from `natural_objects`; language states are
+positive-PMI context vectors of the corpus, averaged per phrase. Hidden
+pairs are read only to report the correlation.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from collections import Counter
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from scipy.linalg import eigh
+from sklearn.decomposition import TruncatedSVD
+
+from .common import read_json, seed_everything, write_json
+from .natural_families import load_phrases, statistics
+from .synth_cc_battery import moment_field
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--objects", default="artifacts/vg_5k/natural_objects.pt")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--fit-scenes", type=int, default=1200)
+ parser.add_argument("--word-vectors", type=int, default=48)
+ parser.add_argument("--min-count", type=int, default=60)
+ parser.add_argument("--max-vocabulary", type=int, default=600)
+ parser.add_argument("--context-window", type=int, default=2)
+ parser.add_argument("--keep", type=int, default=8)
+ parser.add_argument("--shrinkage", type=float, default=0.05)
+ parser.add_argument("--views", type=int, default=1)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="artifacts/vg_5k/natural_pipeline.pt")
+ return parser.parse_args()
+
+
+def word_vectors(
+ phrases: list[list[str]], args: argparse.Namespace
+) -> dict[str, np.ndarray]:
+ """Positive-PMI context vectors, reduced by truncated SVD."""
+ counts = Counter(token for tokens in phrases for token in set(tokens))
+ vocabulary = [
+ word
+ for word, count in counts.most_common(args.max_vocabulary)
+ if count >= args.min_count
+ ]
+ _, _, context = statistics(phrases, set(vocabulary), args.context_window)
+ features = sorted({key for word in vocabulary for key in context[word]})
+ index = {key: position for position, key in enumerate(features)}
+ matrix = np.zeros((len(vocabulary), len(features)))
+ for row, word in enumerate(vocabulary):
+ for key, value in context[word].items():
+ matrix[row, index[key]] = value
+ total = matrix.sum()
+ expected = matrix.sum(1, keepdims=True) * matrix.sum(0, keepdims=True) / total
+ pmi = np.log(np.maximum(matrix, 1e-9) / np.maximum(expected, 1e-9))
+ pmi[matrix == 0] = 0.0
+ embedding = TruncatedSVD(args.word_vectors, random_state=args.seed).fit_transform(
+ np.maximum(pmi, 0.0)
+ )
+ embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9)
+ return {word: embedding[row] for row, word in enumerate(vocabulary)}
+
+
+def content_directions(
+ sets: list[np.ndarray], shrinkage: float
+) -> tuple[np.ndarray, np.ndarray]:
+ """Directions separating scenes more than they separate parts."""
+ means = np.stack([item.mean(0) for item in sets])
+ centre = means.mean(0)
+ within = np.concatenate([item - item.mean(0, keepdims=True) for item in sets])
+ scatter_within = within.T @ within / max(len(within) - 1, 1)
+ centred = means - centre
+ scatter_between = centred.T @ centred / max(len(centred) - 1, 1)
+ trace = np.trace(scatter_within) / len(scatter_within)
+ values, vectors = eigh(
+ scatter_between, scatter_within + shrinkage * trace * np.eye(len(scatter_within))
+ )
+ return centre, vectors[:, np.argsort(values)[::-1]]
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ vg_dir = Path(args.vg_dir)
+ truth = [
+ json.loads(line)
+ for line in (vg_dir / "ground_truth.private.jsonl")
+ .read_text(encoding="utf-8")
+ .splitlines()
+ if line.strip()
+ ]
+ records = {
+ json.loads(line)["node_id"]: json.loads(line)
+ for line in (vg_dir / "text_nodes.jsonl").read_text(encoding="utf-8").splitlines()
+ if line.strip()
+ }
+ state = torch.load(args.objects, map_location="cpu", weights_only=False)
+ per_view = state.get("view_segments") or [state["segments"]]
+ views = min(args.views, len(per_view))
+ index = {node: position for position, node in enumerate(state["node_ids"])}
+
+ vectors = word_vectors(load_phrases(vg_dir / "text_nodes.jsonl", "region_closed"), args)
+
+ def language_state(node: str) -> np.ndarray | None:
+ rows = []
+ for phrase in records[node]["region_closed"]:
+ hits = [vectors[token] for token in re.findall(r"[a-z]+", phrase.lower())
+ if token in vectors]
+ if hits:
+ rows.append(np.mean(hits, axis=0))
+ return np.stack(rows) if rows else None
+
+ def vision_state(node: str, view: int) -> np.ndarray | None:
+ segments = per_view[view][index[node]]
+ return (
+ np.stack([item["feature"] for item in segments]).astype(np.float64)
+ if segments
+ else None
+ )
+
+ pairs = [
+ pair
+ for pair in truth
+ if pair["vision_node_id"] in index and pair["text_node_id"] in records
+ ]
+ fit = pairs[args.samples : args.samples + args.fit_scenes]
+ language_fit = [language_state(p["text_node_id"]) for p in fit]
+ language_fit = [item for item in language_fit if item is not None and len(item) > 1]
+ vision_fit = [vision_state(p["vision_node_id"], 0) for p in fit]
+ vision_fit = [item for item in vision_fit if item is not None and len(item) > 1]
+ language_centre, language_basis = content_directions(language_fit, args.shrinkage)
+ vision_centre, vision_basis = content_directions(vision_fit, args.shrinkage)
+
+ evaluated = [
+ pair for pair in pairs[: args.samples]
+ if language_state(pair["text_node_id"]) is not None
+ ]
+ keep = args.keep
+
+ def project(raw: np.ndarray, centre: np.ndarray, basis: np.ndarray) -> torch.Tensor:
+ width = min(keep, basis.shape[1])
+ return F.normalize(
+ torch.tensor((raw - centre) @ basis[:, :width], dtype=torch.float32), dim=-1
+ )
+
+ text_sets = [
+ project(language_state(p["text_node_id"]), language_centre, language_basis)
+ for p in evaluated
+ ]
+ view_fields = []
+ for view in range(views):
+ sets = [
+ project(vision_state(p["vision_node_id"], view), vision_centre, vision_basis)
+ for p in evaluated
+ ]
+ view_fields.append(moment_field(sets))
+ visual_field = torch.stack(view_fields).mean(0)
+ text_field = moment_field(text_sets)
+
+ mask = ~np.eye(len(evaluated), dtype=bool)
+ correlation = float(
+ np.corrcoef(
+ visual_field.double().numpy()[mask], text_field.double().numpy()[mask]
+ )[0, 1]
+ )
+ torch.save(
+ {"visual_field": visual_field, "text_field": text_field,
+ "nodes": [p["vision_node_id"] for p in evaluated]},
+ args.output,
+ )
+ summary = {
+ "protocol": (
+ "Continuous states, content projection fitted per modality on "
+ "scenes outside the evaluated set. Hidden pairs are read only "
+ "for the correlation."
+ ),
+ "samples": len(evaluated),
+ "views": views,
+ "kept_directions": keep,
+ "field_correlation_at_truth": correlation,
+ "recovery_threshold": 0.9,
+ }
+ print(json.dumps(summary))
+ write_json(str(args.output).replace(".pt", ".json"), summary)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/precompute_gw.py b/worldalign/precompute_gw.py
new file mode 100644
index 0000000..8eac056
--- /dev/null
+++ b/worldalign/precompute_gw.py
@@ -0,0 +1,51 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import torch
+
+from .common import read_json, seed_everything
+from .gw import gw_pseudo_targets
+from .io import load_feature_pair, select_rows
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--output", default="artifacts/gw.pt")
+ p.add_argument("--clusters", type=int, default=128)
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument("--max-iter", type=int, default=100)
+ return p.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
+ x = select_rows(
+ vision["features"], vlookup, manifest["vision_only_train"]
+ )
+ y = select_rows(text["features"], tlookup, manifest["text_only_train"])
+ result = gw_pseudo_targets(
+ x,
+ y,
+ clusters=args.clusters,
+ seed=args.seed,
+ max_iter=args.max_iter,
+ )
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(result, args.output)
+ print(
+ f"Wrote {args.output}: K={result['clusters']}, "
+ f"GW={result['gw_distance']:.6f}, "
+ f"row_entropy={result['coupling_row_entropy']:.6f}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/prepare.py b/worldalign/prepare.py
new file mode 100644
index 0000000..dbdd101
--- /dev/null
+++ b/worldalign/prepare.py
@@ -0,0 +1,80 @@
+from __future__ import annotations
+
+import argparse
+from collections import Counter
+
+import numpy as np
+from datasets import load_dataset
+
+from .common import DATASET_NAME, DATASET_SPLIT, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--output", default="artifacts/manifest.json")
+ p.add_argument("--dataset", default=DATASET_NAME)
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument("--unpaired-per-modality", type=int, default=12_000)
+ p.add_argument("--paired-train", type=int, default=12_000)
+ p.add_argument("--max-eval", type=int, default=1_000)
+ return p.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ dataset = load_dataset(args.dataset, split=DATASET_SPLIT)
+ by_split: dict[str, list[int]] = {}
+ for i, split in enumerate(dataset["split"]):
+ by_split.setdefault(split, []).append(i)
+
+ counts = Counter(dataset["split"])
+ print(f"Internal Flickr splits: {dict(counts)}")
+ train = np.asarray(by_split["train"], dtype=np.int64)
+ rng = np.random.default_rng(args.seed)
+ rng.shuffle(train)
+
+ n_unpaired = min(args.unpaired_per_modality, len(train) // 2)
+ vision_only = train[:n_unpaired]
+ text_only = train[n_unpaired : 2 * n_unpaired]
+ assert not set(vision_only.tolist()) & set(text_only.tolist())
+
+ paired_n = min(args.paired_train, len(train))
+ paired_train = train[:paired_n]
+
+ val_key = "val" if "val" in by_split else "validation"
+ val = np.asarray(by_split[val_key], dtype=np.int64)[: args.max_eval]
+ test = np.asarray(by_split["test"], dtype=np.int64)[: args.max_eval]
+
+ all_rows = sorted(
+ set(vision_only.tolist())
+ | set(text_only.tolist())
+ | set(paired_train.tolist())
+ | set(val.tolist())
+ | set(test.tolist())
+ )
+ manifest = {
+ "dataset": args.dataset,
+ "dataset_split": DATASET_SPLIT,
+ "seed": args.seed,
+ "vision_only_train": vision_only.tolist(),
+ "text_only_train": text_only.tolist(),
+ "paired_train": paired_train.tolist(),
+ "val": val.tolist(),
+ "test": test.tolist(),
+ "all_rows": all_rows,
+ "protocol": (
+ "vision_only_train and text_only_train contain disjoint image IDs; "
+ "val/test pairs are held out from all bridge training"
+ ),
+ }
+ write_json(args.output, manifest)
+ print(
+ f"Wrote {args.output}: unpaired={n_unpaired}+{n_unpaired}, "
+ f"paired_upper_bound={paired_n}, val={len(val)}, test={len(test)}, "
+ f"features={len(all_rows)} rows"
+ )
+
+
+if __name__ == "__main__":
+ main()
+
diff --git a/worldalign/ricci_control.py b/worldalign/ricci_control.py
new file mode 100644
index 0000000..3091b22
--- /dev/null
+++ b/worldalign/ricci_control.py
@@ -0,0 +1,258 @@
+"""Ricci-flow control for the on-manifold assignment gate.
+
+Hypothesis under test: a discrete Ricci flow that smooths each modality's
+relational geometry before matching could repair the local ordering that
+static relation fields fail. The flow is run independently per modality
+with identical hyperparameters; hidden pairs never touch the flow or the
+energy and only score orderings, as in the base gate.
+
+Two geometries are gated per condition: heat-kernel channels on the raw
+kNN graph (diffusion without flow) and the same construction after
+Ollivier-Ricci weight evolution. Differences between them are attributable
+to the flow itself rather than to the diffusion representation.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import ot
+import torch
+from scipy.sparse import csr_matrix
+from scipy.sparse.csgraph import shortest_path
+
+from .common import seed_everything, write_json
+from .manifold_gate import (
+ gate_a_global_ranking,
+ gate_b_transpositions,
+ load_flickr,
+ load_vg,
+ standardize_relation,
+ steepest_descent,
+)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--dataset", choices=["flickr", "vg"], default="flickr")
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt")
+ parser.add_argument("--text-mode", choices=["single", "orbit_mean"], default="orbit_mean")
+ parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt")
+ parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt")
+ parser.add_argument(
+ "--vg-ground-truth", default="artifacts/vg_5k/ground_truth.private.jsonl"
+ )
+ parser.add_argument("--vg-bundle-channels", action="store_true")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--subset-seed", type=int, default=0)
+ parser.add_argument("--neighbors", type=int, default=16)
+ parser.add_argument("--flow-iterations", type=int, default=10)
+ parser.add_argument("--flow-step", type=float, default=0.4)
+ parser.add_argument("--lazy-alpha", type=float, default=0.5)
+ parser.add_argument("--heat-times", default="1.0,4.0")
+ parser.add_argument("--random-perms", type=int, default=300)
+ parser.add_argument("--derangement-samples", type=int, default=50)
+ parser.add_argument("--descent-restarts", type=int, default=2)
+ 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="cpu")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument("--output", default="artifacts/manifold_gate/ricci_control.json")
+ parser.add_argument("--trajectory-output")
+ return parser.parse_args()
+
+
+def knn_edges(distance: np.ndarray, k: int) -> list[tuple[int, int]]:
+ order = distance.argsort(-1)[:, 1 : k + 1]
+ edges = {
+ (min(i, int(j)), max(i, int(j)))
+ for i in range(len(distance))
+ for j in order[i]
+ }
+ return sorted(edges)
+
+
+def graph_apsp(size: int, weights: dict[tuple[int, int], float]) -> np.ndarray:
+ rows, cols, vals = [], [], []
+ for (i, j), w in weights.items():
+ rows.extend((i, j))
+ cols.extend((j, i))
+ vals.extend((w, w))
+ graph = csr_matrix((vals, (rows, cols)), shape=(size, size))
+ return shortest_path(graph, method="D", directed=False)
+
+
+def ollivier_ricci_apsp(
+ distance: np.ndarray, args: argparse.Namespace, iterations: int
+) -> np.ndarray:
+ """All-pairs geodesics after Ollivier-Ricci weight evolution.
+
+ Lazy uniform neighbor measures, W1 ground costs from current geodesics,
+ multiplicative weight update w <- w * (1 - step * kappa), total edge
+ mass renormalized each iteration. iterations=0 gives the un-flowed
+ graph geometry for the diffusion-only control.
+ """
+ edges = knn_edges(distance, args.neighbors)
+ weights = {edge: max(float(distance[edge]), 1e-9) for edge in edges}
+ total = sum(weights.values())
+ neighbor_map: dict[int, list[int]] = {}
+ for i, j in edges:
+ neighbor_map.setdefault(i, []).append(j)
+ neighbor_map.setdefault(j, []).append(i)
+ apsp = graph_apsp(len(distance), weights)
+ for _ in range(iterations):
+ updated: dict[tuple[int, int], float] = {}
+ for i, j in edges:
+ support_i = [i] + neighbor_map[i]
+ support_j = [j] + neighbor_map[j]
+ mass_i = np.full(len(support_i), (1 - args.lazy_alpha) / len(neighbor_map[i]))
+ mass_i[0] = args.lazy_alpha
+ mass_j = np.full(len(support_j), (1 - args.lazy_alpha) / len(neighbor_map[j]))
+ mass_j[0] = args.lazy_alpha
+ ground = apsp[np.ix_(support_i, support_j)]
+ if not np.isfinite(ground).all():
+ finite_max = apsp[np.isfinite(apsp)].max()
+ ground = np.where(np.isfinite(ground), ground, 2.0 * finite_max)
+ wasserstein = ot.emd2(mass_i, mass_j, ground)
+ geodesic = max(float(apsp[i, j]), 1e-9)
+ curvature = 1.0 - wasserstein / geodesic
+ updated[(i, j)] = max(1e-9, weights[(i, j)] * (1.0 - args.flow_step * curvature))
+ scale = total / sum(updated.values())
+ weights = {edge: w * scale for edge, w in updated.items()}
+ apsp = graph_apsp(len(distance), weights)
+ return apsp
+
+
+def geometry_channels(
+ apsp: np.ndarray, heat_times: tuple[float, ...]
+) -> torch.Tensor:
+ """Standardized relation channels of a flowed geometry.
+
+ Channel 0 is the negative geodesic field; the rest are heat kernels of
+ the normalized Laplacian of a geodesic-scale affinity.
+ """
+ finite = apsp[np.isfinite(apsp) & (apsp > 0)]
+ scale = np.median(finite)
+ capped = np.where(np.isfinite(apsp), apsp, finite.max() * 2.0)
+ affinity = np.exp(-capped / scale)
+ np.fill_diagonal(affinity, 1.0)
+ degree = affinity.sum(-1)
+ normalized = affinity / np.sqrt(degree[:, None] * degree[None, :])
+ values, vectors = np.linalg.eigh((normalized + normalized.T) / 2.0)
+ laplacian_eigen = 1.0 - values
+ channels = [torch.from_numpy(-capped).double()]
+ for t in heat_times:
+ heat = (vectors * np.exp(-t * laplacian_eigen)) @ vectors.T
+ channels.append(torch.from_numpy(heat).double())
+ return torch.stack(
+ [standardize_relation(channel)[0] for channel in channels]
+ )
+
+
+def run_gates(
+ text_channels: torch.Tensor,
+ visual_channels: torch.Tensor,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+) -> dict:
+ text_relation = text_channels[0]
+ visual_relation = visual_channels[0]
+ report = {
+ "gate_a": gate_a_global_ranking(
+ text_channels, visual_channels, text_relation, visual_relation, args, generator
+ ),
+ "gate_b": gate_b_transpositions(
+ text_channels, visual_channels, text_relation, visual_relation, None
+ ),
+ "descent_from_true": steepest_descent(
+ text_channels,
+ visual_channels,
+ text_relation,
+ visual_relation,
+ torch.arange(len(visual_relation)),
+ args,
+ ),
+ }
+ report["descent_from_true"].pop("final_permutation", None)
+ report["descent_from_random"] = []
+ for _ in range(args.descent_restarts):
+ start = torch.argsort(torch.rand(len(visual_relation), generator=generator))
+ result = steepest_descent(
+ text_channels, visual_channels, text_relation, visual_relation, start, args
+ )
+ result.pop("final_permutation", None)
+ result.pop("trajectory", None)
+ report["descent_from_random"].append(result)
+ return report
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ data = load_flickr(args) if args.dataset == "flickr" else load_vg(args)
+ heat_times = tuple(float(t) for t in args.heat_times.split(","))
+ states = {
+ "visual": torch.nn.functional.normalize(
+ data["visual_views"].double().mean(1), dim=-1
+ ),
+ "text": torch.nn.functional.normalize(
+ data["text_views"].double().mean(1), dim=-1
+ ),
+ }
+ distances = {
+ side: (1.0 - state @ state.T).clamp_min(0.0).numpy()
+ for side, state in states.items()
+ }
+ report: dict = {
+ "protocol": (
+ "Each modality's kNN geometry evolves independently under "
+ "Ollivier-Ricci flow; the identical heat-kernel channels are "
+ "gated with and without the flow. Hidden pairs score orderings "
+ "only."
+ ),
+ "meta": {**data["meta"], "flow": vars(args)},
+ "conditions": {},
+ }
+ for label, iterations in (
+ ("diffusion_no_flow", 0),
+ ("ricci_flow", args.flow_iterations),
+ ):
+ channels = {}
+ for side in ("visual", "text"):
+ apsp = ollivier_ricci_apsp(distances[side], args, iterations)
+ channels[side] = geometry_channels(apsp, heat_times)
+ generator = torch.Generator().manual_seed(args.seed)
+ result = run_gates(channels["text"], channels["visual"], args, generator)
+ report["conditions"][label] = result
+ print(
+ json.dumps(
+ {
+ label: {
+ "true_z_mse": result["gate_a"]["random"]["mse"]["true_z"],
+ "improving_fraction": result["gate_b"]["improving_fraction"],
+ "descent_keeps": result["descent_from_true"]["final_accuracy"],
+ "counterfeit_found": any(
+ r["final_objective"]
+ < result["gate_a"]["true"]["mse"]
+ and r["final_accuracy"] < 0.5
+ for r in result["descent_from_random"]
+ ),
+ }
+ }
+ )
+ )
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/spectral_match.py b/worldalign/spectral_match.py
new file mode 100644
index 0000000..9e1a0fe
--- /dev/null
+++ b/worldalign/spectral_match.py
@@ -0,0 +1,172 @@
+"""Spectral matching of relation fields: Umeyama and GRAMPA.
+
+Relation fields are N x N regardless of embedding dimension, so the two
+modalities need no common representation dimension. What they do need is
+comparable spectra: eigenvector matching degrades when effective ranks
+differ or when eigenvalues cluster. GRAMPA is built for that regime -- it
+weights every pair of eigenvectors by 1 / ((lambda_i - mu_j)^2 + eta^2)
+instead of pairing them one to one -- so both solvers are provided along
+with the spectral compatibility diagnostic that predicts whether either
+can work.
+
+No local search: these are polynomial-time solvers that sidestep the
+glassy landscape entirely. Hidden pairs score the output only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+from scipy.optimize import linear_sum_assignment
+
+from .common import read_json, seed_everything, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--fields", help="Saved .pt with visual/text fields.")
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--eta", default="0.05,0.1,0.2,0.5")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v0/spectral_match.json"
+ )
+ parser.add_argument("--fields-output", default="")
+ return parser.parse_args()
+
+
+def spectral_profile(matrix: np.ndarray) -> dict:
+ values = np.linalg.eigvalsh(matrix)[::-1]
+ magnitude = np.abs(values)
+ total = magnitude.sum()
+ cumulative = np.cumsum(magnitude) / max(total, 1e-12)
+ participation = (magnitude.sum() ** 2) / max((magnitude**2).sum(), 1e-12)
+ gaps = np.abs(np.diff(values))
+ return {
+ "top_eigenvalues": values[:12].tolist(),
+ "effective_rank_participation": float(participation),
+ "rank_for_90_percent": int(np.searchsorted(cumulative, 0.90) + 1),
+ "rank_for_99_percent": int(np.searchsorted(cumulative, 0.99) + 1),
+ "median_relative_gap": float(
+ np.median(gaps) / max(magnitude.max(), 1e-12)
+ ),
+ "min_relative_gap_top20": float(
+ gaps[:20].min() / max(magnitude.max(), 1e-12)
+ ),
+ }
+
+
+def umeyama(visual: np.ndarray, text: np.ndarray) -> np.ndarray:
+ """Classic eigenvector-magnitude matching, sign ambiguity absorbed."""
+ _, u = np.linalg.eigh(visual)
+ _, v = np.linalg.eigh(text)
+ score = np.abs(u) @ np.abs(v).T
+ rows, cols = linear_sum_assignment(-score)
+ return cols
+
+
+def grampa(visual: np.ndarray, text: np.ndarray, eta: float) -> np.ndarray:
+ """Pairwise eigen-alignment similarity, robust to clustered spectra."""
+ lam, u = np.linalg.eigh(visual)
+ mu, v = np.linalg.eigh(text)
+ ones = np.ones(len(visual))
+ left = u.T @ ones # [N]
+ right = v.T @ ones
+ weight = np.outer(left, right) / ((lam[:, None] - mu[None, :]) ** 2 + eta**2)
+ similarity = u @ weight @ v.T
+ rows, cols = linear_sum_assignment(-similarity)
+ return cols
+
+
+def score_assignment(
+ assignment: np.ndarray, hidden: np.ndarray, size: int
+) -> dict:
+ """assignment[i] is the shuffled-space index matched to visual row i.
+
+ Shuffled index j denotes original node hidden[j], and visual row i
+ denotes original node i, so the match is correct when
+ hidden[assignment[i]] == i.
+ """
+ return {
+ "accuracy": float((hidden[assignment] == np.arange(size)).mean()),
+ "chance": 1.0 / size,
+ }
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ if args.fields:
+ state = torch.load(args.fields, map_location="cpu", weights_only=False)
+ visual_field = state["visual_field"]
+ text_field = state["text_field"]
+ else:
+ from .synth_triangle_gate import build_fields
+
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ rows = manifest[args.split][: args.samples]
+ visual_field, text_field = build_fields(args, rows, manifest)
+ if args.fields_output:
+ torch.save(
+ {"visual_field": visual_field, "text_field": text_field},
+ args.fields_output,
+ )
+
+ size = len(visual_field)
+ visual = visual_field.double().numpy()
+ text = text_field.double().numpy()
+ # Center and scale: spectral methods compare shapes, not offsets.
+ mask = ~np.eye(size, dtype=bool)
+ for matrix in (visual, text):
+ values = matrix[mask]
+ matrix -= values.mean()
+ matrix /= values.std()
+ np.fill_diagonal(matrix, 0.0)
+
+ generator = np.random.default_rng(args.seed)
+ hidden = generator.permutation(size)
+ text_shuffled = text[np.ix_(hidden, hidden)]
+
+ report = {
+ "protocol": (
+ "Polynomial-time spectral solvers on N x N relation fields; "
+ "embedding dimensions are irrelevant by construction. The "
+ "hidden shuffle is applied to the text field and used only to "
+ "score the returned assignment."
+ ),
+ "samples": size,
+ "spectra": {
+ "visual": spectral_profile(visual),
+ "text": spectral_profile(text_shuffled),
+ },
+ "solvers": {},
+ }
+ # Both fields are built over the same row list, so they are aligned at
+ # the truth without any permutation.
+ correlation = float(np.corrcoef(visual[mask], text[mask])[0, 1])
+ report["field_correlation_at_truth"] = correlation
+
+ assignment = umeyama(visual, text_shuffled)
+ report["solvers"]["umeyama"] = score_assignment(assignment, hidden, size)
+ print(json.dumps({"umeyama": report["solvers"]["umeyama"]}))
+
+ for eta in (float(value) for value in args.eta.split(",")):
+ assignment = grampa(visual, text_shuffled, eta)
+ result = score_assignment(assignment, hidden, size)
+ report["solvers"][f"grampa_eta{eta}"] = result
+ print(json.dumps({f"grampa_eta{eta}": result}))
+
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_cc_battery.py b/worldalign/synth_cc_battery.py
new file mode 100644
index 0000000..be28a01
--- /dev/null
+++ b/worldalign/synth_cc_battery.py
@@ -0,0 +1,377 @@
+"""Upper-bound set battery: connected-component sprites as vision sets.
+
+The learned towers have not yet produced object states, which leaves two
+hypotheses entangled: the set-kernel machinery could be wrong, or only
+the towers could be short. This battery separates them. On this world
+the background is flat, so connected bright components ARE the objects;
+per-component centered sprites are model-free object states of the
+minimal-world-knowledge class (like the color anchors on natural data:
+declared, fixed, no learning). If set-kernel relation fields built from
+these pass the gate and support recovery, the machinery is validated and
+the remaining gap is exactly "an SSL objective that discovers objects".
+
+Text sets are the per-group phrase states of the set battery. Hidden
+pairs score orderings only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from scipy import ndimage
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .manifold_gate import standardize_relation
+from .ricci_control import run_gates
+from .synth_set_battery import (
+ parse_group_phrases,
+ set_similarity_field,
+ text_group_sets,
+)
+from .synth_towers import load_image
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ 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("--sprite-window", type=int, default=56)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument(
+ "--features", choices=["pixels", "descriptors", "onehot"], default="descriptors"
+ )
+ parser.add_argument("--text-mode", choices=["lm", "bow"], default="bow")
+ parser.add_argument(
+ "--kernel", choices=["matching", "moment"], default="matching",
+ help="moment: symmetric-tensor (Fock) set kernel, no matching step",
+ )
+ parser.add_argument("--vision-views", type=int, default=4)
+ 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/cc_battery_gate.json"
+ )
+ return parser.parse_args()
+
+
+def component_sprites(
+ image: torch.Tensor, window: int, merge_distance: float
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Centered sprites of bright connected components, ring-merged.
+
+ Same-group objects sit on a small ring; merging nearby components
+ reassembles groups so a sprite carries multiplicity as pattern.
+ """
+ array = image.permute(1, 2, 0).numpy()
+ background = np.median(array.reshape(-1, 3), axis=0)
+ foreground = (np.abs(array - background).sum(-1) > 0.12)
+ labels, count = ndimage.label(foreground)
+ if count == 0:
+ return torch.zeros(1, 3 * window * window), torch.ones(1)
+ centers = np.array(ndimage.center_of_mass(foreground, labels, range(1, count + 1)))
+ sizes = ndimage.sum(foreground, labels, range(1, count + 1))
+ # Merge components whose centers are close (ring members).
+ parent = list(range(count))
+
+ def find(a: int) -> int:
+ while parent[a] != a:
+ parent[a] = parent[parent[a]]
+ a = parent[a]
+ return a
+
+ for a in range(count):
+ for b in range(a + 1, count):
+ if np.linalg.norm(centers[a] - centers[b]) < merge_distance:
+ parent[find(a)] = find(b)
+ groups: dict[int, list[int]] = {}
+ for a in range(count):
+ groups.setdefault(find(a), []).append(a)
+
+ height, width = foreground.shape
+ half = window // 2
+ padded = np.pad(array, ((half, half), (half, half), (0, 0)))
+ padded_mask = np.pad(foreground, half)
+ sprites, weights = [], []
+ for members in groups.values():
+ member_mask = np.isin(labels, [m + 1 for m in members])
+ mass = float(member_mask.sum())
+ ys, xs = np.nonzero(member_mask)
+ cy, cx = int(ys.mean()), int(xs.mean())
+ patch = padded[cy : cy + window, cx : cx + window].copy()
+ mask_patch = padded_mask[cy : cy + window, cx : cx + window]
+ patch[~mask_patch] = 0.0
+ sprites.append(torch.from_numpy(patch).float().flatten())
+ weights.append(mass)
+ weights = torch.tensor(weights)
+ return torch.stack(sprites), weights / weights.sum().clamp_min(1e-8)
+
+
+HUE_CENTERS = {
+ "red": 0.0, "orange": 30.0, "yellow": 60.0, "green": 120.0,
+ "cyan": 180.0, "blue": 220.0, "purple": 275.0, "pink": 330.0,
+}
+
+
+def component_descriptors(
+ image: torch.Tensor, merge_distance: float
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Rotation-invariant group descriptors: color classes, log-area,
+ member multiplicity, and normalized single-member area."""
+ import colorsys
+
+ array = image.permute(1, 2, 0).numpy()
+ background = np.median(array.reshape(-1, 3), axis=0)
+ foreground = np.abs(array - background).sum(-1) > 0.12
+ labels, count = ndimage.label(foreground)
+ if count == 0:
+ return torch.zeros(1, 19), torch.ones(1)
+ centers = np.array(
+ ndimage.center_of_mass(foreground, labels, range(1, count + 1))
+ )
+ parent = list(range(count))
+
+ def find(a: int) -> int:
+ while parent[a] != a:
+ parent[a] = parent[parent[a]]
+ a = parent[a]
+ return a
+
+ for a in range(count):
+ for b in range(a + 1, count):
+ if np.linalg.norm(centers[a] - centers[b]) < merge_distance:
+ parent[find(a)] = find(b)
+ groups: dict[int, list[int]] = {}
+ for a in range(count):
+ groups.setdefault(find(a), []).append(a)
+ descriptors, weights = [], []
+ total = foreground.size
+ for members in groups.values():
+ member_mask = np.isin(labels, [m + 1 for m in members])
+ pixels = array[member_mask]
+ mean_rgb = pixels.mean(0)
+ h, s_, v = colorsys.rgb_to_hsv(*mean_rgb.tolist())
+ hue = h * 360.0
+ vector = np.zeros(16)
+ if v > 0.75 and s_ < 0.25:
+ vector[8] = 1.0 # white
+ elif v < 0.25:
+ vector[9] = 1.0 # black
+ elif s_ < 0.25:
+ vector[10] = 1.0 # gray
+ elif abs(hue - 30.0) < 25.0 and v < 0.55:
+ vector[11] = 1.0 # brown
+ else:
+ for index, center in enumerate(HUE_CENTERS.values()):
+ distance = min(abs(hue - center), 360.0 - abs(hue - center))
+ if distance < 25.0:
+ vector[index] = 1.0
+ break
+ member_count = len(members)
+ area = float(member_mask.sum()) / total
+ vector[12] = np.log(area + 1e-6) / 6.0
+ vector[13] = (member_count - 1) / 3.0
+ vector[14] = np.log(area / member_count + 1e-6) / 6.0
+ vector[15] = 1.0
+ # Rotation-invariant shape features of the largest single member:
+ # compactness, convexity, and inertia eccentricity separate the
+ # six shapes without orientation.
+ largest = max(members, key=lambda m: (labels == m + 1).sum())
+ single = labels == largest + 1
+ area_px = float(single.sum())
+ eroded = ndimage.binary_erosion(single)
+ perimeter = float((single & ~eroded).sum())
+ compactness = 4.0 * np.pi * area_px / max(perimeter, 1.0) ** 2
+ ys, xs = np.nonzero(single)
+ ys = ys - ys.mean(); xs = xs - xs.mean()
+ cov = np.cov(np.stack([xs, ys])) + 1e-6 * np.eye(2)
+ eigenvalues = np.linalg.eigvalsh(cov)
+ eccentricity = float(1.0 - eigenvalues[0] / eigenvalues[1])
+ hull_span = (xs.max() - xs.min() + 1) * (ys.max() - ys.min() + 1)
+ boxfill = area_px / max(hull_span, 1.0)
+ shape_vector = np.array([compactness, eccentricity, boxfill])
+ vector = np.concatenate([vector, shape_vector])
+ descriptors.append(torch.tensor(vector, dtype=torch.float32))
+ weights.append(float(member_mask.sum()))
+ weights = torch.tensor(weights)
+ return torch.stack(descriptors), weights / weights.sum().clamp_min(1e-8)
+
+
+def onehot_descriptors(
+ raw_sets: list[torch.Tensor],
+) -> list[torch.Tensor]:
+ """Factor-mirrored one-hot recoding of descriptor sets.
+
+ Sizes are binned by corpus terciles of single-member log-area with the
+ middle bin unmarked, mirroring the text side where medium size has no
+ word; shapes are k-means clusters of the rotation-invariant shape
+ features. Both statistics come from the evaluated corpus itself,
+ unimodally. Output channels mirror the bag-of-words support: color
+ (12), small/large (2), count (4), shape cluster (6).
+ """
+ from sklearn.cluster import KMeans
+
+ all_groups = torch.cat(raw_sets)
+ # The renderer shrinks radii with member count, so raw single-member
+ # area confounds size class with multiplicity; residualize log-area on
+ # member count before binning (unimodal statistics).
+ counts_all = (all_groups[:, 13] * 3.0).round().clamp(0, 3)
+ area_all = all_groups[:, 14]
+ count_means = {}
+ for value in (0.0, 1.0, 2.0, 3.0):
+ chosen = counts_all == value
+ count_means[value] = float(area_all[chosen].mean()) if chosen.any() else 0.0
+ adjusted_all = area_all - torch.tensor(
+ [count_means[float(v)] for v in counts_all]
+ )
+ low, high = adjusted_all.quantile(1.0 / 3.0), adjusted_all.quantile(2.0 / 3.0)
+ shape_features = all_groups[:, 16:19].numpy()
+ clusters = KMeans(n_clusters=6, n_init=10, random_state=0).fit(shape_features)
+ recoded = []
+ for groups in raw_sets:
+ vectors = torch.zeros(len(groups), 24)
+ vectors[:, :12] = groups[:, :12]
+ counts_here = (groups[:, 13] * 3.0).round().clamp(0, 3)
+ adjusted = groups[:, 14] - torch.tensor(
+ [count_means[float(v)] for v in counts_here]
+ )
+ vectors[:, 12] = (adjusted <= low).float() # small
+ vectors[:, 13] = (adjusted >= high).float() # large
+ counts = (groups[:, 13] * 3.0).round().long().clamp(0, 3)
+ vectors[torch.arange(len(groups)), 14 + counts] = 1.0
+ labels = clusters.predict(groups[:, 16:19].numpy())
+ vectors[torch.arange(len(groups)), 18 + labels] = 1.0
+ recoded.append(vectors)
+ return recoded
+
+
+def moment_field(sets: list[torch.Tensor]) -> torch.Tensor:
+ """Symmetric-tensor (second-quantized) set kernel, matching-free.
+
+ phi(S) concatenates the degree-1 and degree-2 moments of the set; the
+ field is the Gram matrix of normalized phi. No assignment step, so no
+ matching-value or set-size bias can enter.
+ """
+ phis = []
+ for members in sets:
+ m1 = members.mean(0)
+ m2 = (members[:, :, None] * members[:, None, :]).mean(0).flatten()
+ phi = torch.cat([m1, m2])
+ phis.append(phi / phi.norm().clamp_min(1e-9))
+ stacked = torch.stack(phis)
+ return stacked @ stacked.T
+
+
+def phrase_bow_sets(
+ rows: list[int], captions: list[list[str]], vocabulary: list[str]
+) -> list[torch.Tensor]:
+ index = {word: i for i, word in enumerate(vocabulary)}
+ sets = []
+ for row in rows:
+ phrases = parse_group_phrases(captions[row][0])
+ vectors = torch.zeros(len(phrases), len(index))
+ for p, phrase in enumerate(phrases):
+ for token in phrase.split():
+ if token in index:
+ vectors[p, index[token]] += 1.0
+ sets.append(F.normalize(vectors, dim=-1))
+ return sets
+
+
+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]
+ image_dir = Path(manifest["image_dir"])
+
+ per_view_fields = []
+ set_sizes = []
+ for view in range(args.vision_views):
+ vision_sets = []
+ raw_sets = []
+ for row in tqdm(rows, desc=f"cc sprites v{view}"):
+ image = load_image(image_dir / f"scene{row:06d}_v{view}.png")
+ if args.features in ("descriptors", "onehot"):
+ sprites, _ = component_descriptors(image, args.merge_distance)
+ else:
+ sprites, _ = component_sprites(
+ image, args.sprite_window, args.merge_distance
+ )
+ raw_sets.append(sprites)
+ if view == 0:
+ set_sizes.append(len(sprites))
+ if args.features == "onehot":
+ vision_sets = [
+ F.normalize(v, dim=-1) for v in onehot_descriptors(raw_sets)
+ ]
+ else:
+ vision_sets = [F.normalize(v, dim=-1) for v in raw_sets]
+ if args.kernel == "moment":
+ per_view_fields.append(moment_field(vision_sets))
+ else:
+ per_view_fields.append(set_similarity_field(vision_sets))
+
+ if args.text_mode == "bow":
+ text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"])
+ else:
+ text_sets = text_group_sets(rows, captions, args)
+ if args.kernel == "moment":
+ text_field = moment_field(text_sets)
+ else:
+ text_field = set_similarity_field(text_sets)
+ # Mass weighting inside the matching corrupts similarity grading
+ # (0.16 vs 0.45 against soft truth); match unweighted. Averaging the
+ # per-view fields cancels segmentation errors across resampled layouts.
+ visual_field = torch.stack(per_view_fields).mean(0)
+ 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": (
+ "Vision sets are model-free connected-component sprites "
+ "(declared minimal world knowledge); text sets are per-group "
+ "phrase states. Hidden pairs score orderings only."
+ ),
+ "split": args.split,
+ "samples": len(rows),
+ "mean_vision_set_size": float(np.mean(set_sizes)),
+ **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()
diff --git a/worldalign/synth_decode.py b/worldalign/synth_decode.py
new file mode 100644
index 0000000..183e4a9
--- /dev/null
+++ b/worldalign/synth_decode.py
@@ -0,0 +1,150 @@
+"""Close the loop: recovered pairs to a map to generated descriptions.
+
+Matching is transductive -- it aligns one fixed population. A usable
+system needs a map, so the recovered assignment is treated as pseudo-pair
+supervision for a small vision-to-text-state regressor, and the map is
+then applied to scenes that took no part in the matching. Descriptions
+are read out through the frozen text tower.
+
+Two controls decide whether the output is image-conditioned: the same
+pipeline with the recovered assignment replaced by a random one, and the
+same map applied to a shuffled image. Hidden pairs score; nothing here
+trains on them.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+from .common import read_json, seed_everything, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v1")
+ parser.add_argument("--states", required=True,
+ help=".pt with vision_states, text_states, rows for "
+ "the matched population and the held-out split.")
+ parser.add_argument("--assignment", required=True,
+ help=".pt with the recovered permutation and truth.")
+ parser.add_argument("--ridge", type=float, default=1.0)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="artifacts/synth_v1/decode.json")
+ return parser.parse_args()
+
+
+def fit_map(source: np.ndarray, target: np.ndarray, ridge: float) -> np.ndarray:
+ source = np.concatenate([source, np.ones((len(source), 1))], axis=1)
+ gram = source.T @ source + ridge * np.eye(source.shape[1])
+ return np.linalg.solve(gram, source.T @ target)
+
+
+def apply_map(weights: np.ndarray, source: np.ndarray) -> np.ndarray:
+ source = np.concatenate([source, np.ones((len(source), 1))], axis=1)
+ return source @ weights
+
+
+def nearest_caption(
+ predicted: np.ndarray, bank: np.ndarray
+) -> np.ndarray:
+ predicted = predicted / np.linalg.norm(predicted, axis=1, keepdims=True).clip(1e-9)
+ bank = bank / np.linalg.norm(bank, axis=1, keepdims=True).clip(1e-9)
+ return (predicted @ bank.T).argmax(1)
+
+
+def factor_scores(
+ predicted_rows: list[int], truth_rows: list[int], scenes: list[dict]
+) -> dict:
+ """Do the retrieved descriptions state the right world facts?"""
+ colour_f1, count_ok, group_ok = [], [], []
+ for predicted, truth in zip(predicted_rows, truth_rows):
+ p, t = scenes[predicted], scenes[truth]
+ pc = {g["color"] for g in p["groups"]}
+ tc = {g["color"] for g in t["groups"]}
+ overlap = len(pc & tc)
+ precision = overlap / max(len(pc), 1)
+ recall = overlap / max(len(tc), 1)
+ colour_f1.append(
+ 0.0 if precision + recall == 0 else 2 * precision * recall / (precision + recall)
+ )
+ count_ok.append(
+ sorted(g["count"] for g in p["groups"])
+ == sorted(g["count"] for g in t["groups"])
+ )
+ group_ok.append(len(p["groups"]) == len(t["groups"]))
+ return {
+ "colour_set_f1": float(np.mean(colour_f1)),
+ "count_multiset_exact": float(np.mean(count_ok)),
+ "group_count_exact": float(np.mean(group_ok)),
+ }
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"]
+ states = torch.load(args.states, map_location="cpu", weights_only=False)
+ recovered = torch.load(args.assignment, map_location="cpu", weights_only=False)
+
+ match_vision = states["match_vision"].double().numpy()
+ match_text = states["match_text"].double().numpy()
+ held_vision = states["held_vision"].double().numpy()
+ held_text = states["held_text"].double().numpy()
+ held_rows = states["held_rows"]
+
+ assignment = np.asarray(recovered["assignment"])
+ generator = np.random.default_rng(args.seed)
+ random_assignment = generator.permutation(len(assignment))
+
+ report = {
+ "protocol": (
+ "The recovered assignment supplies pseudo-pairs for a ridge "
+ "map from vision states to text states; the map is applied to "
+ "held-out scenes that took no part in matching. Random "
+ "assignment and shuffled-image controls bound the claim."
+ ),
+ "matched_population": len(assignment),
+ "held_out": len(held_rows),
+ "recovery_accuracy": float(
+ (assignment == np.arange(len(assignment))).mean()
+ ),
+ "conditions": {},
+ }
+
+ for label, permutation in (
+ ("recovered_pairs", assignment),
+ ("random_pairs", random_assignment),
+ ):
+ weights = fit_map(match_vision, match_text[permutation], args.ridge)
+ predicted = apply_map(weights, held_vision)
+ retrieved = nearest_caption(predicted, held_text)
+ exact = float((retrieved == np.arange(len(held_rows))).mean())
+ entry = {
+ "held_out_retrieval_exact": exact,
+ "chance": 1.0 / len(held_rows),
+ **factor_scores(
+ [held_rows[i] for i in retrieved], held_rows, scenes
+ ),
+ }
+ if label == "recovered_pairs":
+ shuffled = generator.permutation(len(held_vision))
+ predicted_shuffled = apply_map(weights, held_vision[shuffled])
+ retrieved_shuffled = nearest_caption(predicted_shuffled, held_text)
+ entry["shuffled_image_control"] = factor_scores(
+ [held_rows[i] for i in retrieved_shuffled], held_rows, scenes
+ )
+ report["conditions"][label] = entry
+ print(json.dumps({label: entry}))
+
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_deep_gate.py b/worldalign/synth_deep_gate.py
new file mode 100644
index 0000000..74bfea9
--- /dev/null
+++ b/worldalign/synth_deep_gate.py
@@ -0,0 +1,203 @@
+"""Corrected gate: the deepest reachable minimum, not descent retention.
+
+Today's lesson: descent retention measures the search operator, not the
+energy. Sampled-proposal descent keeps the truth for every candidate
+energy, while long tempering on the same energy reaches states well below
+it. The only decision-relevant question is therefore
+
+ E(truth) <= E(deepest state a strong searcher reaches) ?
+
+This module answers it uniformly for the candidate energies (pairwise
+moment kernel, triangle-only, and their sum) with one strong searcher:
+long parallel tempering with large proposal batches from random starts,
+plus a truth-initialized tempering arm that reports whether the truth
+itself survives thermal agitation. Hidden pairs score only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+
+from .common import read_json, seed_everything, write_json
+from .synth_triangle_gate import TriangleEnergy, build_fields, standardized
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--triples", type=int, default=200000)
+ parser.add_argument(
+ "--energies",
+ default="pair,triangle,both",
+ help="Comma list from pair, triangle, both.",
+ )
+ parser.add_argument("--replicas", type=int, default=6)
+ parser.add_argument("--rounds", type=int, default=1500)
+ parser.add_argument("--proposals", type=int, default=64)
+ parser.add_argument("--temp-high", type=float, default=3e-2)
+ parser.add_argument("--temp-low", type=float, default=1e-4)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument("--output", default="artifacts/synth_v0/deep_gate.json")
+ return parser.parse_args()
+
+
+def temper(
+ energy: TriangleEnergy,
+ starts: list[torch.Tensor],
+ truth: torch.Tensor,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+ label: str,
+) -> dict:
+ size = energy.size
+ temperatures = torch.logspace(
+ torch.log10(torch.tensor(args.temp_low)),
+ torch.log10(torch.tensor(args.temp_high)),
+ len(starts),
+ )
+ states = [start.clone() for start in starts]
+ energies = [energy.total(state) for state in states]
+ best = {"energy": min(energies), "accuracy": 0.0}
+ for round_index in range(args.rounds):
+ for replica in range(len(states)):
+ temperature = float(temperatures[replica])
+ for _ in range(args.proposals):
+ p = int(torch.randint(0, size, (1,), generator=generator))
+ q = int(torch.randint(0, size, (1,), generator=generator))
+ if p == q:
+ continue
+ delta = energy.swap_delta(states[replica], p, q)
+ threshold = -temperature * float(
+ torch.rand(1, generator=generator).clamp_min(1e-12).log()
+ )
+ if delta < threshold:
+ states[replica][[p, q]] = states[replica][[q, p]]
+ energies[replica] += delta
+ if round_index % args.exchange_every == 0:
+ for replica in range(len(states) - 1):
+ gap = (energies[replica] - energies[replica + 1]) * (
+ 1.0 / float(temperatures[replica])
+ - 1.0 / float(temperatures[replica + 1])
+ )
+ accept = gap > 0 or float(
+ torch.rand(1, generator=generator)
+ ) < min(1.0, float(torch.tensor(gap).exp()))
+ if accept:
+ states[replica], states[replica + 1] = (
+ states[replica + 1],
+ states[replica],
+ )
+ energies[replica], energies[replica + 1] = (
+ energies[replica + 1],
+ energies[replica],
+ )
+ cold = min(range(len(states)), key=lambda r: energies[r])
+ if energies[cold] < best["energy"]:
+ best = {
+ "energy": energies[cold],
+ "accuracy": float(
+ (states[cold].cpu() == truth.cpu()).float().mean()
+ ),
+ "round": round_index,
+ }
+ exact = [energy.total(state) for state in states]
+ cold = min(range(len(states)), key=lambda r: exact[r])
+ return {
+ "arm": label,
+ "best_seen": best,
+ "final_cold_energy": exact[cold],
+ "final_cold_accuracy": float(
+ (states[cold].cpu() == truth.cpu()).float().mean()
+ ),
+ "final_accuracies": [
+ float((state.cpu() == truth.cpu()).float().mean()) for state in states
+ ],
+ }
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ rows = manifest[args.split][: args.samples]
+ visual_field, text_field = build_fields(args, rows, manifest)
+
+ device = torch.device(args.device)
+ size = len(rows)
+ generator = torch.Generator().manual_seed(args.seed)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden).to(device)
+ text = standardized(text_field[hidden][:, hidden].double().to(device)).float()
+ visual = standardized(visual_field.double().to(device)).float()
+
+ triples = torch.randint(0, size, (args.triples, 3), generator=generator)
+ triples = triples[
+ (triples[:, 0] != triples[:, 1])
+ & (triples[:, 1] != triples[:, 2])
+ & (triples[:, 0] != triples[:, 2])
+ ].to(device)
+
+ weights = {
+ "pair": (1.0, 0.0),
+ "triangle": (0.0, 1.0),
+ "both": (1.0, 1.0),
+ }
+ report = {
+ "protocol": (
+ "The decision statistic is the deepest energy a strong "
+ "searcher reaches versus the energy of the truth. Descent "
+ "retention is reported but not used: it measures the search "
+ "operator. Hidden pairs score only."
+ ),
+ "samples": size,
+ "triples": len(triples),
+ "energies": {},
+ }
+ for name in (item.strip() for item in args.energies.split(",")):
+ pair_weight, triangle_weight = weights[name]
+ energy = TriangleEnergy(text, visual, triples, pair_weight, triangle_weight)
+ true_energy = energy.total(truth)
+ random_starts = [
+ torch.argsort(torch.rand(size, generator=generator)).to(device)
+ for _ in range(args.replicas)
+ ]
+ from_random = temper(energy, random_starts, truth, args, generator, "random")
+ from_truth = temper(
+ energy,
+ [truth.clone() for _ in range(args.replicas)],
+ truth,
+ args,
+ generator,
+ "truth",
+ )
+ deepest = min(from_random["best_seen"]["energy"], from_truth["best_seen"]["energy"])
+ entry = {
+ "true_energy": true_energy,
+ "from_random": from_random,
+ "from_truth": from_truth,
+ "deepest_seen": deepest,
+ "margin_over_true": deepest / true_energy - 1.0,
+ "passes": bool(deepest >= true_energy - 1e-9),
+ "recovery_accuracy": max(
+ from_random["best_seen"]["accuracy"],
+ from_random["final_cold_accuracy"],
+ ),
+ }
+ report["energies"][name] = entry
+ print(json.dumps({name: {k: entry[k] for k in ("true_energy", "deepest_seen", "margin_over_true", "passes", "recovery_accuracy")}}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_extract.py b/worldalign/synth_extract.py
new file mode 100644
index 0000000..3186c3e
--- /dev/null
+++ b/worldalign/synth_extract.py
@@ -0,0 +1,144 @@
+"""Feature extraction for the synthetic world, in main-pipeline schema.
+
+Emits vision.pt, text.pt, and text_orbits.pt files with the same fields
+the Flickr loaders read, so the diagnostic, gate, projection, and recovery
+stack runs on the synthetic world unchanged.
+"""
+
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from tqdm import tqdm
+
+from .common import batch_indices, read_json
+from .synth_towers import TextTower, VisionTower, load_image, tokenize
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--vision-tower", default="artifacts/synth_v0/vision_tower.pt")
+ parser.add_argument("--text-tower", default="artifacts/synth_v0/text_tower.pt")
+ parser.add_argument("--batch-size", type=int, default=512)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--vision-output", default="artifacts/synth_v0/vision.pt")
+ parser.add_argument("--text-output", default="artifacts/synth_v0/text.pt")
+ parser.add_argument(
+ "--orbits-output", default="artifacts/synth_v0/text_orbits.pt"
+ )
+ return parser.parse_args()
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ captions = read_json(Path(args.data_dir, "captions.json"))["captions"]
+ image_dir = Path(manifest["image_dir"])
+ views = manifest["visual_views"]
+
+ vision_state = torch.load(args.vision_tower, map_location="cpu", weights_only=False)
+ vision_args = vision_state["args"]
+ vision = VisionTower(
+ manifest["image_size"],
+ vision_args["patch"],
+ vision_args["dim"],
+ vision_args["depth"],
+ vision_args["heads"],
+ ).to(args.device)
+ vision.load_state_dict(vision_state["model"])
+ vision.eval()
+
+ 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)]
+ features = []
+ objective = vision_args.get("objective", "infonce")
+ for indices in tqdm(list(batch_indices(len(jobs), args.batch_size)), desc="vision"):
+ 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)
+ tokens = vision.encode(pixels)
+ state = tokens[:, 1:].mean(1) if objective in ("simmim", "data2vec") else tokens[:, 0]
+ features.append(state.float().cpu())
+ view_features = torch.cat(features).reshape(len(rendered_rows), views, -1)
+ torch.save(
+ {
+ "model": "synth_vision_tower",
+ "rows": rendered_rows,
+ "features": F.normalize(view_features.mean(1), dim=-1),
+ "view_features": view_features,
+ "views_per_scene": views,
+ },
+ args.vision_output,
+ )
+
+ text_state = torch.load(args.text_tower, map_location="cpu", weights_only=False)
+ text_args = text_state["args"]
+ vocab = text_state["vocab"]
+ text = TextTower(
+ len(vocab),
+ text_args["text_dim"],
+ text_args["depth"],
+ text_args.get("text_heads", 4),
+ text_args["context"],
+ ).to(args.device)
+ text.load_state_dict(text_state["model"])
+ text.eval()
+
+ rows = list(range(manifest["all_rows"]))
+ orbit_size = len(captions[0])
+ jobs = [(row, k) for row in rows for k in range(orbit_size)]
+ pooled = []
+ for indices in tqdm(list(batch_indices(len(jobs), args.batch_size)), desc="text"):
+ batch = [tokenize(captions[jobs[i][0]][jobs[i][1]], vocab) for i in indices]
+ longest = min(text_args["context"], max(len(s) for s in batch))
+ tokens = torch.zeros(len(batch), longest, dtype=torch.long)
+ for index, sentence in enumerate(batch):
+ clipped = sentence[:longest]
+ tokens[index, : len(clipped)] = torch.tensor(clipped)
+ tokens = tokens.to(args.device)
+ hidden = text(tokens)
+ mask = (tokens != 0).float()[..., None]
+ pooled.append(
+ ((hidden * mask).sum(1) / mask.sum(1).clamp_min(1.0)).float().cpu()
+ )
+ orbit_features = torch.cat(pooled).reshape(len(rows), orbit_size, -1)
+ torch.save(
+ {
+ "model": "synth_text_tower",
+ "layer": -1,
+ "rows": rows,
+ "features": orbit_features[:, 0],
+ "captions": [captions[row][0] for row in rows],
+ "all_captions": captions,
+ },
+ args.text_output,
+ )
+ torch.save(
+ {
+ "model": "synth_text_tower",
+ "layer": -1,
+ "rows": rows,
+ "features": orbit_features,
+ "captions": captions,
+ "views_per_orbit": orbit_size,
+ "row_groups": ["all"],
+ },
+ args.orbits_output,
+ )
+ print(
+ f"Wrote {args.vision_output}, {args.text_output}, {args.orbits_output}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_fast_gate.py b/worldalign/synth_fast_gate.py
new file mode 100644
index 0000000..54b6891
--- /dev/null
+++ b/worldalign/synth_fast_gate.py
@@ -0,0 +1,345 @@
+"""Closed-form batched gate: exact all-triple triangle energy.
+
+The sampled triangle energy cost 5.7 ms per proposal because each
+evaluation gathered a variable-length triple set in Python. It has a
+closed form. Writing M(sigma) for the elementwise product of the
+sigma-permuted text field with the visual field,
+
+ pairwise = const - 2 * sum(M)
+ all-triple = const - 2 * trace(M^3) / 6
+
+because trace of a power is invariant under simultaneous row-column
+permutation, so the text-only and vision-only terms do not move. One
+batched matrix product evaluates trace(M^3) = sum(M * (M @ M)) for
+hundreds of candidate permutations at once, over every C(N,3) triple
+rather than a sample.
+
+This buys a genuinely stronger searcher: exact steepest descent over all
+N(N-1)/2 transpositions per step, plus batched-proposal tempering.
+Hidden pairs score orderings only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+
+from .common import read_json, seed_everything, write_json
+from .synth_triangle_gate import build_fields, standardized
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--fields", default="", help="Saved .pt with visual_field/text_field.")
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--energies", default="pair,triangle,both")
+ parser.add_argument("--chunk", type=int, default=512)
+ parser.add_argument("--descent-restarts", type=int, default=5)
+ parser.add_argument("--descent-max-steps", type=int, default=4000)
+ parser.add_argument("--replicas", type=int, default=8)
+ parser.add_argument("--rounds", type=int, default=3000)
+ parser.add_argument("--proposals", type=int, default=64)
+ parser.add_argument("--temp-high", type=float, default=3e-2)
+ parser.add_argument("--temp-low", type=float, default=1e-4)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--polish-steps", type=int, default=2000)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument("--output", default="artifacts/synth_v0/fast_gate.json")
+ return parser.parse_args()
+
+
+class ClosedFormEnergy:
+ """Energy as a function of M = permuted-text * visual, batched."""
+
+ def __init__(
+ self,
+ text: torch.Tensor,
+ visual: torch.Tensor,
+ pair_weight: float,
+ triangle_weight: float,
+ chunk: int,
+ ) -> None:
+ self.text = text
+ self.visual = visual
+ self.pair_weight = pair_weight
+ self.triangle_weight = triangle_weight
+ self.chunk = chunk
+ self.size = len(visual)
+ self.pair_count = self.size * (self.size - 1)
+ self.triple_count = self.size * (self.size - 1) * (self.size - 2)
+ # Permutation-invariant constants, kept so reported values match
+ # the direct definitions of the two energies.
+ square_text = text * text
+ square_visual = visual * visual
+ self.pair_constant = float(square_text.sum() + square_visual.sum())
+ self.triangle_constant = float(
+ self._trace_cube(square_text[None])[0]
+ + self._trace_cube(square_visual[None])[0]
+ )
+
+ @staticmethod
+ def _trace_cube(matrices: torch.Tensor) -> torch.Tensor:
+ return (matrices * torch.bmm(matrices, matrices)).sum((-2, -1))
+
+ def energy(self, permutations: torch.Tensor) -> torch.Tensor:
+ """Exact energy for a batch of permutations [B, N]."""
+ values = []
+ for start in range(0, len(permutations), self.chunk):
+ block = permutations[start : start + self.chunk]
+ permuted = self.text[block[:, :, None], block[:, None, :]]
+ product = permuted * self.visual
+ total = torch.zeros(len(block), device=permuted.device)
+ if self.pair_weight:
+ pair = (self.pair_constant - 2.0 * product.sum((-2, -1))) / (
+ self.pair_count
+ )
+ total = total + self.pair_weight * pair
+ if self.triangle_weight:
+ triangle = (
+ self.triangle_constant - 2.0 * self._trace_cube(product)
+ ) / self.triple_count
+ total = total + self.triangle_weight * triangle
+ values.append(total)
+ return torch.cat(values)
+
+
+def all_swaps(size: int, device: torch.device) -> torch.Tensor:
+ rows, cols = torch.triu_indices(size, size, offset=1, device=device)
+ return torch.stack([rows, cols], dim=1)
+
+
+def apply_swaps(permutation: torch.Tensor, swaps: torch.Tensor) -> torch.Tensor:
+ batch = permutation[None].repeat(len(swaps), 1)
+ index = torch.arange(len(swaps), device=permutation.device)
+ p, q = swaps[:, 0], swaps[:, 1]
+ values_p = batch[index, p].clone()
+ batch[index, p] = batch[index, q]
+ batch[index, q] = values_p
+ return batch
+
+
+def steepest_descent(
+ energy: ClosedFormEnergy,
+ start: torch.Tensor,
+ swaps: torch.Tensor,
+ max_steps: int,
+) -> tuple[torch.Tensor, float]:
+ """Exact steepest descent over every transposition each step."""
+ current = start.clone()
+ value = float(energy.energy(current[None])[0])
+ for _ in range(max_steps):
+ candidates = apply_swaps(current, swaps)
+ values = energy.energy(candidates)
+ best = int(values.argmin())
+ if float(values[best]) >= value - 1e-12:
+ break
+ current = candidates[best]
+ value = float(values[best])
+ return current, value
+
+
+def temper(
+ energy: ClosedFormEnergy,
+ starts: torch.Tensor,
+ args: argparse.Namespace,
+ generator: torch.Generator,
+ device: torch.device,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Batched-proposal parallel tempering; returns states and energies."""
+ size = energy.size
+ replicas = len(starts)
+ temperatures = torch.logspace(
+ torch.log10(torch.tensor(args.temp_low)),
+ torch.log10(torch.tensor(args.temp_high)),
+ replicas,
+ ).to(device)
+ states = starts.clone()
+ values = energy.energy(states)
+ for round_index in range(args.rounds):
+ p = torch.randint(
+ 0, size, (replicas, args.proposals), generator=generator
+ ).to(device)
+ q = torch.randint(
+ 0, size, (replicas, args.proposals), generator=generator
+ ).to(device)
+ valid = p != q
+ batch = states[:, None, :].repeat(1, args.proposals, 1)
+ index_r = torch.arange(replicas, device=device)[:, None]
+ original_p = batch.gather(2, p[..., None]).squeeze(-1)
+ original_q = batch.gather(2, q[..., None]).squeeze(-1)
+ batch.scatter_(2, p[..., None], original_q[..., None])
+ batch.scatter_(2, q[..., None], original_p[..., None])
+ flat = batch.reshape(replicas * args.proposals, size)
+ proposal_values = energy.energy(flat).reshape(replicas, args.proposals)
+ deltas = proposal_values - values[:, None]
+ noise = torch.rand(
+ replicas, args.proposals, generator=generator
+ ).to(device)
+ threshold = -temperatures[:, None] * noise.clamp_min(1e-12).log()
+ accept = (deltas < threshold) & valid
+ first = torch.where(
+ accept.any(-1),
+ accept.float().argmax(-1),
+ torch.zeros(replicas, dtype=torch.long, device=device),
+ )
+ taken = accept.any(-1)
+ chosen = batch[index_r.squeeze(-1), first]
+ states = torch.where(taken[:, None], chosen, states)
+ values = torch.where(
+ taken, proposal_values[index_r.squeeze(-1), first], values
+ )
+ if round_index % args.exchange_every == 0:
+ for replica in range(replicas - 1):
+ gap = (values[replica] - values[replica + 1]) * (
+ 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1]
+ )
+ accept_swap = gap > 0 or float(
+ torch.rand(1, generator=generator)
+ ) < float(gap.exp().clamp(max=1.0))
+ if accept_swap:
+ states[[replica, replica + 1]] = states[
+ [replica + 1, replica]
+ ]
+ values[[replica, replica + 1]] = values[
+ [replica + 1, replica]
+ ]
+ return states, values
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ if args.fields:
+ state = torch.load(args.fields, map_location="cpu", weights_only=False)
+ visual_field, text_field = state["visual_field"], state["text_field"]
+ else:
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ rows = manifest[args.split][: args.samples]
+ visual_field, text_field = build_fields(args, rows, manifest)
+
+ device = torch.device(args.device)
+ size = len(visual_field)
+ generator = torch.Generator().manual_seed(args.seed)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden).to(device) # scored, never optimized against
+ text = standardized(text_field[hidden][:, hidden].double().to(device)).float()
+ visual = standardized(visual_field.double().to(device)).float()
+ swaps = all_swaps(size, device)
+
+ weights = {"pair": (1.0, 0.0), "triangle": (0.0, 1.0), "both": (1.0, 1.0)}
+ report = {
+ "protocol": (
+ "Closed-form energies over all triples; the decision statistic "
+ "is the energy of the truth against the deepest state reached "
+ "by exact steepest descent and batched tempering. Hidden pairs "
+ "score only."
+ ),
+ "samples": size,
+ "energies": {},
+ }
+ for name in (item.strip() for item in args.energies.split(",")):
+ pair_weight, triangle_weight = weights[name]
+ energy = ClosedFormEnergy(
+ text, visual, pair_weight, triangle_weight, args.chunk
+ )
+ true_energy = float(energy.energy(truth[None])[0])
+ random_batch = torch.stack(
+ [
+ torch.argsort(torch.rand(size, generator=generator))
+ for _ in range(200)
+ ]
+ ).to(device)
+ random_values = energy.energy(random_batch)
+
+ kept, kept_value = steepest_descent(
+ energy, truth, swaps, args.descent_max_steps
+ )
+ retention = float((kept.cpu() == truth.cpu()).float().mean())
+
+ quenches = []
+ for restart in range(args.descent_restarts):
+ start = torch.argsort(torch.rand(size, generator=generator)).to(device)
+ final, value = steepest_descent(
+ energy, start, swaps, args.descent_max_steps
+ )
+ quenches.append(
+ {
+ "energy": value,
+ "accuracy": float((final.cpu() == truth.cpu()).float().mean()),
+ }
+ )
+
+ starts = torch.stack(
+ [
+ torch.argsort(torch.rand(size, generator=generator))
+ for _ in range(args.replicas)
+ ]
+ ).to(device)
+ states, values = temper(energy, starts, args, generator, device)
+ cold = int(values.argmin())
+ polished, polished_value = steepest_descent(
+ energy, states[cold], swaps, args.polish_steps
+ )
+ tempering = {
+ "cold_energy": float(values[cold]),
+ "polished_energy": polished_value,
+ "polished_accuracy": float(
+ (polished.cpu() == truth.cpu()).float().mean()
+ ),
+ "best_replica_accuracy": max(
+ float((state.cpu() == truth.cpu()).float().mean())
+ for state in states
+ ),
+ }
+ deepest = min(
+ [polished_value, float(values.min())]
+ + [item["energy"] for item in quenches]
+ )
+ entry = {
+ "true_energy": true_energy,
+ "random_mean": float(random_values.mean()),
+ "true_z": float(
+ (random_values.mean() - true_energy)
+ / random_values.std().clamp_min(1e-12)
+ ),
+ "descent_retention_diagnostic": retention,
+ "descent_from_truth_energy": kept_value,
+ "quenches": quenches,
+ "tempering": tempering,
+ "deepest_seen": deepest,
+ "margin_over_true": deepest / abs(true_energy) - true_energy / abs(true_energy),
+ "passes": bool(deepest >= true_energy - 1e-9),
+ "recovery_accuracy": tempering["polished_accuracy"],
+ }
+ report["energies"][name] = entry
+ print(
+ json.dumps(
+ {
+ name: {
+ key: entry[key]
+ for key in (
+ "true_energy",
+ "deepest_seen",
+ "passes",
+ "recovery_accuracy",
+ "descent_retention_diagnostic",
+ "true_z",
+ )
+ }
+ }
+ )
+ )
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_precision_gate.py b/worldalign/synth_precision_gate.py
new file mode 100644
index 0000000..46887fa
--- /dev/null
+++ b/worldalign/synth_precision_gate.py
@@ -0,0 +1,279 @@
+"""Channel-matched (precision-weighted) relational energy for the synth world.
+
+Planted-problem theory: search is glassy when the energy mismatches the
+generative channel. The orbit provides the channel unimodally -- the
+variance of each vision relation entry across re-rendered views measures
+exactly how corrupted that entry is (segmentation errors are the dominant
+noise and are view-dependent), so entries are weighted by their orbit
+precision. Text fields are parse-deterministic and enter unweighted.
+
+The weighted energy loses the all-pairs closed form (the quadratic terms
+no longer cancel), but each swap delta is still O(N) and vectorizes over
+sampled pairs; gates and tempering below use that form. Hidden pairs
+score orderings only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .synth_cc_battery import (
+ component_descriptors,
+ moment_field,
+ onehot_descriptors,
+ phrase_bow_sets,
+)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--precision-floor", type=float, default=1e-4)
+ parser.add_argument("--random-perms", type=int, default=300)
+ parser.add_argument("--transposition-samples", type=int, default=100000)
+ parser.add_argument("--descent-restarts", type=int, default=5)
+ parser.add_argument("--descent-max-steps", type=int, default=5000)
+ parser.add_argument("--replicas", type=int, default=8)
+ parser.add_argument("--tempering-rounds", type=int, default=60000)
+ parser.add_argument("--temp-high", type=float, default=3e-3)
+ parser.add_argument("--temp-low", type=float, default=1e-5)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v0/precision_gate.json"
+ )
+ return parser.parse_args()
+
+
+def standardized(field: torch.Tensor) -> torch.Tensor:
+ mask = ~torch.eye(len(field), dtype=torch.bool)
+ values = field[mask]
+ out = (field - values.mean()) / values.std().clamp_min(1e-9)
+ return out.masked_fill(~mask, 0.0)
+
+
+def weighted_energy(
+ text: torch.Tensor, visual: torch.Tensor, weights: torch.Tensor,
+ permutations: torch.Tensor,
+) -> torch.Tensor:
+ fields = text[permutations[:, :, None], permutations[:, None, :]]
+ difference = (fields - visual) ** 2 * weights
+ mask = ~torch.eye(text.shape[-1], dtype=torch.bool, device=text.device)
+ return difference[:, mask].sum(-1) / weights[mask].sum()
+
+
+def swap_deltas(
+ permuted_text: torch.Tensor,
+ visual: torch.Tensor,
+ weights: torch.Tensor,
+ pairs_p: torch.Tensor,
+ pairs_q: torch.Tensor,
+) -> torch.Tensor:
+ """Weighted swap deltas, O(N) per proposal, batched over replicas.
+
+ permuted_text: [R, N, N]; pairs: [R, P]. Row and column contributions
+ are equal by symmetry of all three matrices; the k in {p, q} terms are
+ excluded because the (p, q) entry itself is unchanged by the swap.
+ """
+ replica_index = torch.arange(len(permuted_text), device=permuted_text.device)
+ rows_p = permuted_text[replica_index[:, None], pairs_p]
+ rows_q = permuted_text[replica_index[:, None], pairs_q]
+ visual_p, visual_q = visual[pairs_p], visual[pairs_q]
+ weight_p, weight_q = weights[pairs_p], weights[pairs_q]
+ new_p = (rows_q - visual_p) ** 2 * weight_p
+ old_p = (rows_p - visual_p) ** 2 * weight_p
+ new_q = (rows_p - visual_q) ** 2 * weight_q
+ old_q = (rows_q - visual_q) ** 2 * weight_q
+ total = (new_p - old_p + new_q - old_q).sum(-1)
+ columns = torch.stack([pairs_p, pairs_q], -1)
+ correction = torch.zeros_like(total)
+ for slot in range(2):
+ chosen = columns[..., slot]
+ correction = correction + (
+ (rows_q.gather(-1, chosen[..., None]) - visual_p.gather(-1, chosen[..., None])) ** 2
+ - (rows_p.gather(-1, chosen[..., None]) - visual_p.gather(-1, chosen[..., None])) ** 2
+ ).squeeze(-1) * weight_p.gather(-1, chosen[..., None]).squeeze(-1)
+ correction = correction + (
+ (rows_p.gather(-1, chosen[..., None]) - visual_q.gather(-1, chosen[..., None])) ** 2
+ - (rows_q.gather(-1, chosen[..., None]) - visual_q.gather(-1, chosen[..., None])) ** 2
+ ).squeeze(-1) * weight_q.gather(-1, chosen[..., None]).squeeze(-1)
+ mask_count = weights[~torch.eye(len(visual), dtype=torch.bool, device=visual.device)].sum()
+ return 2.0 * (total - correction) / mask_count
+
+
+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]
+ image_dir = Path(manifest["image_dir"])
+ from .synth_towers import load_image
+
+ per_view = []
+ for view in range(args.vision_views):
+ raw = []
+ for row in tqdm(rows, desc=f"cc v{view}"):
+ sprites, _ = component_descriptors(
+ load_image(image_dir / f"scene{row:06d}_v{view}.png"),
+ args.merge_distance,
+ )
+ raw.append(sprites)
+ sets = [F.normalize(v, dim=-1) for v in onehot_descriptors(raw)]
+ per_view.append(moment_field(sets))
+ stack = torch.stack(per_view)
+ visual_field = stack.mean(0)
+ variance = stack.var(0)
+ weights = 1.0 / (variance + args.precision_floor)
+ weights = weights / weights.mean()
+
+ text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"])
+ text_field = moment_field(text_sets)
+
+ device = torch.device(args.device)
+ size = len(rows)
+ generator = torch.Generator().manual_seed(args.seed)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden)
+ text_input = standardized(text_field[hidden][:, hidden].double()).float().to(device)
+ visual = standardized(visual_field.double()).float().to(device)
+ weight_map = weights.float().to(device)
+
+ identity = torch.arange(size, device=device)
+ true_energy = float(
+ weighted_energy(text_input, visual, weight_map, truth[None].to(device))[0]
+ )
+ random_perms = torch.stack(
+ [torch.argsort(torch.rand(size, generator=generator)) for _ in range(args.random_perms)]
+ ).to(device)
+ random_energies = weighted_energy(text_input, visual, weight_map, random_perms)
+ gate_a = {
+ "true": true_energy,
+ "random_mean": float(random_energies.mean()),
+ "random_std": float(random_energies.std()),
+ "true_z": float((random_energies.mean() - true_energy) / random_energies.std().clamp_min(1e-12)),
+ }
+
+ samples = min(args.transposition_samples, size * (size - 1) // 2)
+ pairs_p = torch.randint(0, size, (1, samples), generator=generator).to(device)
+ pairs_q = torch.randint(0, size, (1, samples), generator=generator).to(device)
+ valid = (pairs_p != pairs_q).squeeze(0)
+ fields_true = text_input[truth.to(device)][:, truth.to(device)][None]
+ deltas = swap_deltas(fields_true, visual, weight_map, pairs_p, pairs_q).squeeze(0)[valid]
+ gate_b = {
+ "sampled": int(valid.sum()),
+ "improving_fraction": float((deltas < 0).float().mean()),
+ }
+
+ def descent(start: torch.Tensor) -> dict:
+ current = start.clone()
+ energy = float(weighted_energy(text_input, visual, weight_map, current[None])[0])
+ for _ in range(args.descent_max_steps):
+ fields = text_input[current][:, current][None]
+ cp = torch.randint(0, size, (1, 4096), generator=generator).to(device)
+ cq = torch.randint(0, size, (1, 4096), generator=generator).to(device)
+ dd = swap_deltas(fields, visual, weight_map, cp, cq).squeeze(0)
+ best = int(dd.argmin())
+ if float(dd[best]) >= -1e-12:
+ break
+ p, q = int(cp[0, best]), int(cq[0, best])
+ current[[p, q]] = current[[q, p]]
+ energy += float(dd[best])
+ return {
+ "accuracy": float((current.cpu() == truth).float().mean()),
+ "energy": float(weighted_energy(text_input, visual, weight_map, current[None])[0]),
+ }
+
+ from_true = descent(truth.to(device))
+ restarts = [
+ descent(torch.argsort(torch.rand(size, generator=generator)).to(device))
+ for _ in range(args.descent_restarts)
+ ]
+ best_random = min(r["energy"] for r in restarts)
+
+ # Tempering recovery with the weighted deltas.
+ temperatures = torch.logspace(
+ torch.log10(torch.tensor(args.temp_low)),
+ torch.log10(torch.tensor(args.temp_high)),
+ args.replicas,
+ ).to(device)
+ permutations = torch.stack(
+ [torch.randperm(size, generator=generator).to(device) for _ in range(args.replicas)]
+ )
+ energies = weighted_energy(text_input, visual, weight_map, permutations)
+ for round_index in range(args.tempering_rounds):
+ fields = text_input[permutations[:, :, None], permutations[:, None, :]]
+ cp = torch.randint(0, size, (args.replicas, 24), generator=generator).to(device)
+ cq = torch.randint(0, size, (args.replicas, 24), generator=generator).to(device)
+ dd = swap_deltas(fields, visual, weight_map, cp, cq)
+ noise = torch.rand(args.replicas, 24, generator=generator).to(device)
+ ok = (dd < -temperatures[:, None] * noise.clamp_min(1e-12).log()) & (cp != cq)
+ for replica in range(args.replicas):
+ hits = torch.nonzero(ok[replica])
+ if not len(hits):
+ continue
+ first = int(hits[0, 0])
+ p, q = int(cp[replica, first]), int(cq[replica, first])
+ permutations[replica][[p, q]] = permutations[replica][[q, p]]
+ energies[replica] = energies[replica] + dd[replica, first]
+ if round_index % args.exchange_every == 0:
+ for replica in range(args.replicas - 1):
+ gap = (energies[replica] - energies[replica + 1]) * (
+ 1.0 / temperatures[replica] - 1.0 / temperatures[replica + 1]
+ )
+ if gap > 0 or torch.rand(1, generator=generator).item() < float(gap.exp()):
+ permutations[[replica, replica + 1]] = permutations[[replica + 1, replica]]
+ energies[[replica, replica + 1]] = energies[[replica + 1, replica]]
+ if round_index % 10000 == 0:
+ energies = weighted_energy(text_input, visual, weight_map, permutations)
+ energies = weighted_energy(text_input, visual, weight_map, permutations)
+ accuracies = (permutations.cpu() == truth[None]).float().mean(-1)
+ cold = int(energies.argmin())
+
+ report = {
+ "protocol": (
+ "Relation entries are weighted by orbit-derived precision "
+ "(variance of the vision field across re-rendered views); "
+ "weights are unimodal statistics. Hidden pairs score only."
+ ),
+ "samples": size,
+ "weight_stats": {
+ "min": float(weights.min()),
+ "median": float(weights.median()),
+ "max": float(weights.max()),
+ },
+ "gate_a": gate_a,
+ "gate_b": gate_b,
+ "descent_from_true": from_true,
+ "descent_from_random": restarts,
+ "tempering": {
+ "cold_energy": float(energies[cold]),
+ "cold_accuracy": float(accuracies[cold]),
+ "best_accuracy": float(accuracies.max()),
+ },
+ "verdict": {
+ "true_energy": true_energy,
+ "best_random_descent": best_random,
+ "counterfeit_found": bool(best_random < true_energy and min(r["accuracy"] for r in restarts) < 0.5),
+ "recovery_accuracy": float(accuracies.max()),
+ },
+ }
+ print(json.dumps({"verdict": report["verdict"], "gate_b": gate_b, "from_true": from_true}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_probes.py b/worldalign/synth_probes.py
new file mode 100644
index 0000000..180b96a
--- /dev/null
+++ b/worldalign/synth_probes.py
@@ -0,0 +1,129 @@
+"""Ground-truth factor probes for synthetic-world representations.
+
+Linear probes from a representation to the discrete scene factors measure
+which world variables survive the encoder and its pooling. Probes are
+fitted per modality on that modality's own training split and evaluated
+on held-out scenes; scene truth is generator metadata, so this is a
+within-modality diagnostic, not a cross-modal signal.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+
+from .common import read_json, write_json
+from .synth_world import COLORS, SHAPES
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--features", required=True)
+ parser.add_argument(
+ "--side", choices=["vision", "text"], required=True,
+ help="Selects the training split whose scenes fit the probes.",
+ )
+ parser.add_argument("--ridge", type=float, default=1.0)
+ parser.add_argument("--output", required=True)
+ return parser.parse_args()
+
+
+def scene_targets(scene: dict) -> dict[str, np.ndarray | float]:
+ color_presence = np.zeros(len(COLORS))
+ shape_presence = np.zeros(len(SHAPES))
+ for group in scene["groups"]:
+ color_presence[list(COLORS).index(group["color"])] = 1.0
+ shape_presence[SHAPES.index(group["shape"])] = 1.0
+ return {
+ "color_presence": color_presence,
+ "shape_presence": shape_presence,
+ "group_count": float(len(scene["groups"])),
+ "object_total": float(sum(g["count"] for g in scene["groups"])),
+ "relation_count": float(len(scene["relations"])),
+ }
+
+
+def ridge_fit(
+ x: np.ndarray, y: np.ndarray, ridge: float
+) -> tuple[np.ndarray, np.ndarray]:
+ x = np.concatenate([x, np.ones((len(x), 1))], axis=1)
+ gram = x.T @ x + ridge * np.eye(x.shape[1])
+ weights = np.linalg.solve(gram, x.T @ y)
+ return weights, x @ weights
+
+
+def evaluate(
+ features_train: np.ndarray,
+ features_test: np.ndarray,
+ train_targets: np.ndarray,
+ test_targets: np.ndarray,
+ ridge: float,
+ binary: bool,
+) -> dict:
+ weights, _ = ridge_fit(features_train, train_targets, ridge)
+ prediction = (
+ np.concatenate([features_test, np.ones((len(features_test), 1))], axis=1)
+ @ weights
+ )
+ if binary:
+ accuracy = float(((prediction > 0.5) == (test_targets > 0.5)).mean())
+ balanced = []
+ for column in range(test_targets.shape[1]):
+ truth = test_targets[:, column] > 0.5
+ if truth.any() and (~truth).any():
+ hit = (prediction[:, column] > 0.5) == truth
+ balanced.append(
+ (hit[truth].mean() + hit[~truth].mean()) / 2.0
+ )
+ return {
+ "accuracy": accuracy,
+ "balanced_accuracy": float(np.mean(balanced)),
+ }
+ residual = prediction[:, 0] - test_targets[:, 0]
+ variance = test_targets[:, 0].var()
+ return {
+ "r2": float(1.0 - residual.var() / max(variance, 1e-9)),
+ "mae": float(np.abs(residual).mean()),
+ }
+
+
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"]
+ state = torch.load(args.features, map_location="cpu", weights_only=False)
+ lookup = {int(row): i for i, row in enumerate(state["rows"])}
+ features = state["features"].float().numpy()
+
+ split_key = "vision_only_train" if args.side == "vision" else "text_only_train"
+ train_rows = [row for row in manifest[split_key] if row in lookup][:8000]
+ test_rows = [row for row in manifest["test"] if row in lookup]
+
+ x_train = features[[lookup[row] for row in train_rows]]
+ x_test = features[[lookup[row] for row in test_rows]]
+ report = {"features": args.features, "side": args.side, "targets": {}}
+ for name in ("color_presence", "shape_presence"):
+ y_train = np.stack([scene_targets(scenes[row])[name] for row in train_rows])
+ y_test = np.stack([scene_targets(scenes[row])[name] for row in test_rows])
+ report["targets"][name] = evaluate(
+ x_train, x_test, y_train, y_test, args.ridge, binary=True
+ )
+ for name in ("group_count", "object_total", "relation_count"):
+ y_train = np.array(
+ [[scene_targets(scenes[row])[name]] for row in train_rows]
+ )
+ y_test = np.array([[scene_targets(scenes[row])[name]] for row in test_rows])
+ report["targets"][name] = evaluate(
+ x_train, x_test, y_train, y_test, args.ridge, binary=False
+ )
+ write_json(args.output, report)
+ print(json.dumps(report, indent=2))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_recovery.py b/worldalign/synth_recovery.py
new file mode 100644
index 0000000..e045a72
--- /dev/null
+++ b/worldalign/synth_recovery.py
@@ -0,0 +1,184 @@
+"""Blind recovery on synthetic set-kernel fields: the end-to-end test.
+
+Builds the connected-component descriptor field and the phrase
+bag-of-words field, hides the text order behind a shuffle, and runs
+parallel tempering on the relational energy alone. Recovery accuracy
+against the hidden truth is the first end-to-end measurement of world
+matching in the closed world.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+from tqdm import tqdm
+
+from .blind_recovery import arm_tempering, permutation_energy_batch
+from .common import read_json, seed_everything, write_json
+from .manifold_gate import standardize_relation
+from .synth_cc_battery import (
+ component_descriptors,
+ moment_field,
+ onehot_descriptors,
+ phrase_bow_sets,
+)
+from .synth_set_battery import set_similarity_field
+from .synth_towers import load_image
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--replicas", type=int, default=8)
+ parser.add_argument("--tempering-rounds", type=int, default=60000)
+ parser.add_argument("--temp-high", type=float, default=3e-3)
+ parser.add_argument("--temp-low", type=float, default=1e-5)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--unary-weight", type=float, default=0.0)
+ parser.add_argument("--init", default="random")
+ parser.add_argument("--residualize-size", action="store_true", default=False)
+ parser.add_argument(
+ "--features", choices=["descriptors", "onehot"], default="onehot"
+ )
+ parser.add_argument("--kernel", choices=["matching", "moment"], default="moment")
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v0/recovery_end_to_end.json"
+ )
+ return parser.parse_args()
+
+
+def residualize(field: torch.Tensor, sizes: torch.Tensor) -> torch.Tensor:
+ """Regress the set-size nuisance out of a matching-value field.
+
+ Set sizes are unimodal observables; their sum, difference, and product
+ explain a size-driven component that differs between modalities and
+ is exploitable by counterfeit assignments.
+ """
+ n = len(field)
+ mask = ~torch.eye(n, dtype=torch.bool)
+ features = torch.stack(
+ [
+ (sizes[:, None] + sizes[None, :])[mask],
+ (sizes[:, None] - sizes[None, :]).abs()[mask],
+ (sizes[:, None] * sizes[None, :])[mask],
+ torch.ones(int(mask.sum()), dtype=torch.float64),
+ ],
+ dim=1,
+ )
+ values = field.double()[mask]
+ solution = torch.linalg.lstsq(features, values[:, None]).solution
+ residual = values - (features @ solution).squeeze(1)
+ output = field.double().clone()
+ output[mask] = residual
+ output.fill_diagonal_(0.0)
+ return output.float()
+
+
+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]
+ image_dir = Path(manifest["image_dir"])
+
+ import torch.nn.functional as F
+
+ per_view_fields = []
+ for view in range(args.vision_views):
+ raw_sets = []
+ for row in tqdm(rows, desc=f"cc v{view}"):
+ image = load_image(image_dir / f"scene{row:06d}_v{view}.png")
+ sprites, _ = component_descriptors(image, args.merge_distance)
+ raw_sets.append(sprites)
+ if args.features == "onehot":
+ vision_sets = [
+ F.normalize(v, dim=-1) for v in onehot_descriptors(raw_sets)
+ ]
+ else:
+ vision_sets = [F.normalize(v, dim=-1) for v in raw_sets]
+ if args.kernel == "moment":
+ per_view_fields.append(moment_field(vision_sets))
+ else:
+ per_view_fields.append(set_similarity_field(vision_sets))
+ visual_field = torch.stack(per_view_fields).mean(0)
+ text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"])
+ if args.kernel == "moment":
+ text_field = moment_field(text_sets)
+ else:
+ text_field = set_similarity_field(text_sets)
+
+ if args.residualize_size:
+ vision_sizes = torch.tensor(
+ [len(s) for s in vision_sets], dtype=torch.float64
+ )
+ text_sizes = torch.tensor([len(s) for s in text_sets], dtype=torch.float64)
+ visual_field = residualize(visual_field, vision_sizes)
+ text_field = residualize(text_field, text_sizes)
+
+ device = torch.device(args.device)
+ size = len(rows)
+ generator = torch.Generator().manual_seed(args.seed)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden)
+ text_input = text_field[hidden][:, hidden]
+
+ text_standardized = standardize_relation(text_input.double())[0].float().to(device)
+ visual_standardized = (
+ standardize_relation(visual_field.double())[0].float().to(device)
+ )
+ true_energy = float(
+ permutation_energy_batch(
+ text_standardized, visual_standardized, truth[None].to(device)
+ )[0]
+ )
+ report = {
+ "protocol": (
+ "Set-kernel fields from released renders and captions; text "
+ "order hidden behind a shuffle; tempering sees no truth. "
+ "Hidden truth scores the outcome only."
+ ),
+ "split": args.split,
+ "samples": size,
+ "true_energy": true_energy,
+ "chance_accuracy": 1.0 / size,
+ "tempering": arm_tempering(
+ text_standardized,
+ visual_standardized,
+ truth,
+ true_energy,
+ args,
+ generator,
+ unary=None,
+ ),
+ }
+ best = max(
+ report["tempering"]["replicas"] + [report["tempering"]["best"]],
+ key=lambda c: c["accuracy"],
+ )
+ by_energy = min(
+ report["tempering"]["replicas"] + [report["tempering"]["best"]],
+ key=lambda c: c["energy"],
+ )
+ report["summary"] = {
+ "best_accuracy": best["accuracy"],
+ "best_accuracy_energy_over_true": best["energy_over_true"],
+ "lowest_energy_accuracy": by_energy["accuracy"],
+ "lowest_energy_over_true": by_energy["energy_over_true"],
+ }
+ print(json.dumps({"summary": report["summary"]}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
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()
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()
diff --git a/worldalign/synth_towers.py b/worldalign/synth_towers.py
new file mode 100644
index 0000000..15221ac
--- /dev/null
+++ b/worldalign/synth_towers.py
@@ -0,0 +1,348 @@
+"""From-scratch unimodal towers for the synthetic closed world.
+
+Vision: a compact ViT trained with InfoNCE over natural orbit positives --
+two renders of the same scene, which differ exactly by the world's
+continuous nuisance. No flips (they would erase left/right semantics) and
+no color jitter (color is world content). Scene identity within the
+vision-only split is unimodal metadata.
+
+Text: a small word-level causal LM trained on the captions of the
+text-only split. Both towers therefore learn the same world from disjoint
+scenes and separate modalities, with dialable capacity.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import math
+import random
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from PIL import Image
+from torch import nn
+
+from .common import read_json, seed_everything
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--side", choices=["vision", "text"], required=True)
+ parser.add_argument(
+ "--objective", choices=["infonce", "simmim", "hybrid", "data2vec"], default="infonce"
+ )
+ parser.add_argument("--mask-ratio", type=float, default=0.6)
+ parser.add_argument("--recon-weight", type=float, default=25.0)
+ parser.add_argument("--ema-decay", type=float, default=0.999)
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--dim", type=int, default=192)
+ parser.add_argument("--depth", type=int, default=6)
+ parser.add_argument("--heads", type=int, default=3)
+ parser.add_argument("--text-dim", type=int, default=256)
+ parser.add_argument("--text-heads", type=int, default=4)
+ parser.add_argument("--patch", type=int, default=16)
+ parser.add_argument("--context", type=int, default=80)
+ parser.add_argument("--epochs", type=int, default=80)
+ parser.add_argument("--batch-size", type=int, default=256)
+ parser.add_argument("--lr", type=float, default=1e-3)
+ parser.add_argument("--temperature", type=float, default=0.2)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260730)
+ parser.add_argument("--output", required=True)
+ return parser.parse_args()
+
+
+class Block(nn.Module):
+ def __init__(self, dim: int, heads: int) -> None:
+ super().__init__()
+ self.norm1 = nn.LayerNorm(dim)
+ self.attention = nn.MultiheadAttention(dim, heads, batch_first=True)
+ self.norm2 = nn.LayerNorm(dim)
+ self.mlp = nn.Sequential(
+ nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim)
+ )
+
+ def forward(
+ self, x: torch.Tensor, mask: torch.Tensor | None = None
+ ) -> torch.Tensor:
+ normed = self.norm1(x)
+ attended, _ = self.attention(
+ normed, normed, normed, attn_mask=mask, need_weights=False
+ )
+ x = x + attended
+ return x + self.mlp(self.norm2(x))
+
+
+class VisionTower(nn.Module):
+ def __init__(self, image_size: int, patch: int, dim: int, depth: int, heads: int) -> None:
+ super().__init__()
+ self.patch = patch
+ self.patch_embed = nn.Conv2d(3, dim, kernel_size=patch, stride=patch)
+ tokens = (image_size // patch) ** 2
+ self.cls = nn.Parameter(torch.zeros(1, 1, dim))
+ self.mask_token = nn.Parameter(torch.zeros(1, 1, dim))
+ self.positions = nn.Parameter(torch.zeros(1, tokens + 1, dim))
+ nn.init.trunc_normal_(self.positions, std=0.02)
+ nn.init.trunc_normal_(self.cls, std=0.02)
+ nn.init.trunc_normal_(self.mask_token, std=0.02)
+ self.blocks = nn.ModuleList(Block(dim, heads) for _ in range(depth))
+ self.norm = nn.LayerNorm(dim)
+ self.head = nn.Sequential(
+ nn.Linear(dim, dim), nn.GELU(), nn.Linear(dim, 128)
+ )
+ self.reconstruction = nn.Linear(dim, patch * patch * 3)
+ self.feature_head = nn.Linear(dim, dim)
+
+ def encode(
+ self, pixels: torch.Tensor, mask: torch.Tensor | None = None
+ ) -> torch.Tensor:
+ x = self.patch_embed(pixels).flatten(2).transpose(1, 2)
+ if mask is not None:
+ x = torch.where(mask[..., None], self.mask_token.expand_as(x), x)
+ x = torch.cat([self.cls.expand(len(x), -1, -1), x], dim=1)
+ x = x + self.positions
+ for block in self.blocks:
+ x = block(x)
+ return self.norm(x)
+
+ def forward(self, pixels: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ tokens = self.encode(pixels)
+ return tokens[:, 0], self.head(tokens[:, 0])
+
+
+class TextTower(nn.Module):
+ def __init__(self, vocab: int, dim: int, depth: int, heads: int, context: int) -> None:
+ super().__init__()
+ self.embed = nn.Embedding(vocab, dim)
+ self.positions = nn.Parameter(torch.zeros(1, context, dim))
+ nn.init.trunc_normal_(self.positions, std=0.02)
+ self.blocks = nn.ModuleList(Block(dim, heads) for _ in range(depth))
+ self.norm = nn.LayerNorm(dim)
+ self.context = context
+
+ def forward(self, tokens: torch.Tensor) -> torch.Tensor:
+ length = tokens.shape[1]
+ x = self.embed(tokens) + self.positions[:, :length]
+ mask = torch.triu(
+ torch.full((length, length), float("-inf"), device=tokens.device), 1
+ )
+ for block in self.blocks:
+ x = block(x, mask)
+ return self.norm(x)
+
+ def logits(self, hidden: torch.Tensor) -> torch.Tensor:
+ return hidden @ self.embed.weight.T
+
+
+def load_image(path: Path) -> torch.Tensor:
+ with Image.open(path) as image:
+ array = np.asarray(image.convert("RGB"), dtype=np.float32) / 255.0
+ return torch.from_numpy(array).permute(2, 0, 1)
+
+
+def patchify(pixels: torch.Tensor, patch: int) -> torch.Tensor:
+ batch, channels, height, width = pixels.shape
+ grid = height // patch
+ x = pixels.reshape(batch, channels, grid, patch, grid, patch)
+ return x.permute(0, 2, 4, 3, 5, 1).reshape(batch, grid * grid, -1)
+
+
+def train_vision(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 = VisionTower(
+ manifest["image_size"], args.patch, args.dim, args.depth, args.heads
+ ).to(args.device)
+ optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.05)
+ tokens = (manifest["image_size"] // args.patch) ** 2
+ if args.objective in ("infonce", "hybrid"):
+ samples_per_epoch = len(rows)
+ else:
+ samples_per_epoch = len(rows) * views
+ teacher = None
+ if args.objective == "data2vec":
+ import copy
+
+ teacher = copy.deepcopy(model)
+ for parameter in teacher.parameters():
+ parameter.requires_grad_(False)
+ steps_per_epoch = max(1, samples_per_epoch // args.batch_size)
+ schedule = torch.optim.lr_scheduler.CosineAnnealingLR(
+ optimizer, T_max=args.epochs * steps_per_epoch
+ )
+ rng = random.Random(args.seed)
+ for epoch in range(args.epochs):
+ total, count = 0.0, 0
+ if args.objective in ("infonce", "hybrid"):
+ order = rows.copy()
+ rng.shuffle(order)
+ for start in range(0, len(order) - args.batch_size + 1, args.batch_size):
+ batch_rows = order[start : start + args.batch_size]
+ pairs = []
+ for row in batch_rows:
+ first, second = rng.sample(range(views), 2)
+ pairs.append(
+ load_image(image_dir / f"scene{row:06d}_v{first}.png")
+ )
+ pairs.append(
+ load_image(image_dir / f"scene{row:06d}_v{second}.png")
+ )
+ pixels = torch.stack(pairs).to(args.device)
+ _, projected = model(pixels)
+ projected = F.normalize(projected, dim=-1)
+ logits = projected @ projected.T / args.temperature
+ logits.fill_diagonal_(float("-inf"))
+ targets = torch.arange(len(projected), device=args.device) ^ 1
+ loss = F.cross_entropy(logits, targets)
+ if args.objective == "hybrid":
+ mask = (
+ torch.rand(len(pixels), tokens, device=args.device)
+ < args.mask_ratio
+ )
+ encoded = model.encode(pixels, mask=mask)
+ predicted = model.reconstruction(encoded[:, 1:][mask])
+ target_patches = patchify(pixels, args.patch)[mask]
+ loss = loss + args.recon_weight * F.mse_loss(
+ predicted, target_patches
+ )
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ optimizer.step()
+ schedule.step()
+ total += float(loss)
+ count += 1
+ else:
+ jobs = [(row, view) for row in rows for view in range(views)]
+ rng.shuffle(jobs)
+ 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)
+ mask = (
+ torch.rand(len(batch), tokens, device=args.device)
+ < args.mask_ratio
+ )
+ encoded = model.encode(pixels, mask=mask)
+ if args.objective == "simmim":
+ predicted = model.reconstruction(encoded[:, 1:][mask])
+ target = patchify(pixels, args.patch)[mask]
+ loss = F.mse_loss(predicted, target)
+ else:
+ with torch.no_grad():
+ reference = teacher.encode(pixels)[:, 1:]
+ reference = F.layer_norm(
+ reference, reference.shape[-1:]
+ )
+ predicted = model.feature_head(encoded[:, 1:][mask])
+ loss = F.smooth_l1_loss(predicted, reference[mask])
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ optimizer.step()
+ schedule.step()
+ if teacher is not None:
+ with torch.no_grad():
+ for student_p, teacher_p in zip(
+ model.parameters(), teacher.parameters()
+ ):
+ teacher_p.lerp_(student_p, 1.0 - args.ema_decay)
+ 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), "side": "vision"},
+ args.output,
+ )
+ print(f"Wrote {args.output}")
+
+
+def tokenize(caption: str, vocab: dict[str, int]) -> list[int]:
+ tokens = caption.replace(",", " ").replace(".", " .").split()
+ return [vocab["<bos>"]] + [vocab[token] for token in tokens]
+
+
+def build_vocab(manifest: dict) -> dict[str, int]:
+ words = list(manifest["vocabulary"]) + ["."]
+ vocab = {"<pad>": 0, "<bos>": 1}
+ for word in sorted(set(words)):
+ vocab.setdefault(word, len(vocab))
+ return vocab
+
+
+def train_text(args: argparse.Namespace) -> None:
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ captions = read_json(Path(args.data_dir, "captions.json"))["captions"]
+ vocab = build_vocab(manifest)
+ sentences = [
+ tokenize(caption, vocab)
+ for row in manifest["text_only_train"]
+ for caption in captions[row]
+ ]
+ model = TextTower(
+ len(vocab), args.text_dim, args.depth, args.text_heads, args.context
+ ).to(args.device)
+ optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=0.01)
+ rng = random.Random(args.seed)
+ steps_per_epoch = max(1, len(sentences) // args.batch_size)
+ schedule = torch.optim.lr_scheduler.CosineAnnealingLR(
+ optimizer, T_max=args.epochs * steps_per_epoch
+ )
+ for epoch in range(args.epochs):
+ rng.shuffle(sentences)
+ total, count = 0.0, 0
+ for start in range(0, len(sentences) - args.batch_size + 1, args.batch_size):
+ batch = sentences[start : start + args.batch_size]
+ longest = min(args.context, max(len(s) for s in batch))
+ tokens = torch.zeros(len(batch), longest, dtype=torch.long)
+ for index, sentence in enumerate(batch):
+ clipped = sentence[:longest]
+ tokens[index, : len(clipped)] = torch.tensor(clipped)
+ tokens = tokens.to(args.device)
+ hidden = model(tokens)
+ logits = model.logits(hidden[:, :-1])
+ targets = tokens[:, 1:]
+ loss = F.cross_entropy(
+ logits.reshape(-1, logits.shape[-1]),
+ targets.reshape(-1),
+ ignore_index=0,
+ )
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ optimizer.step()
+ schedule.step()
+ 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(),
+ "vocab": vocab,
+ "args": vars(args),
+ "side": "text",
+ },
+ args.output,
+ )
+ print(f"Wrote {args.output}")
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ if args.side == "vision":
+ train_vision(args)
+ else:
+ train_text(args)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_transfer_energy.py b/worldalign/synth_transfer_energy.py
new file mode 100644
index 0000000..67a081f
--- /dev/null
+++ b/worldalign/synth_transfer_energy.py
@@ -0,0 +1,276 @@
+"""Prediction-transfer alignment energy for the synthetic world.
+
+The third rung of the energy ladder: no hand-designed geometry and no
+pair-trained EBM. Each modality's own predictor induces a substitution
+kernel over scenes -- which other scenes the world model finds compatible
+-- and the alignment energy is the conjugacy defect of the two kernels
+under a candidate coupling.
+
+- Vision kernel: the masked-reconstruction tower inpaints a half-masked
+ render; pixel distance between the inpainting and other scenes' renders
+ gives K_V[i, i'].
+- Text kernel: the causal LM scores scene j's caption sentences as a
+ continuation of scene i's caption prefix; referential consistency makes
+ same-content continuations likely, giving K_T.
+
+Both kernels are unimodal-predictor functionals. Hidden pairs enter only
+through the gate that scores orderings of the energy.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+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_towers import TextTower, VisionTower, load_image, tokenize
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument(
+ "--vision-tower", default="artifacts/synth_v0/vision_tower_simmim.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("--mask-ratio", type=float, default=0.5)
+ parser.add_argument("--kernel-temperature", type=float, default=0.03)
+ parser.add_argument("--batch-size", type=int, default=256)
+ 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/transfer_energy_gate.json"
+ )
+ parser.add_argument(
+ "--kernels-output", default="artifacts/synth_v0/transfer_kernels.pt"
+ )
+ return parser.parse_args()
+
+
+@torch.inference_mode()
+def vision_kernel(
+ rows: list[int], manifest: dict, args: argparse.Namespace
+) -> torch.Tensor:
+ state = torch.load(args.vision_tower, map_location="cpu", weights_only=False)
+ saved = state["args"]
+ model = VisionTower(
+ manifest["image_size"], saved["patch"], saved["dim"], saved["depth"], saved["heads"]
+ ).to(args.device)
+ model.load_state_dict(state["model"])
+ model.eval()
+ image_dir = Path(manifest["image_dir"])
+ tokens = (manifest["image_size"] // saved["patch"]) ** 2
+ generator = torch.Generator(device=args.device).manual_seed(args.seed)
+
+ references = []
+ for indices in batch_indices(len(rows), args.batch_size):
+ pixels = torch.stack(
+ [load_image(image_dir / f"scene{rows[i]:06d}_v0.png") for i in indices]
+ )
+ references.append(pixels)
+ references = torch.cat(references) # [N, 3, H, W] view-0 renders on cpu
+
+ inpainted = []
+ for indices in tqdm(
+ list(batch_indices(len(rows), args.batch_size)), desc="vision kernel"
+ ):
+ # Inpaint view 1 under a random half mask; compare to view-0 renders.
+ pixels = torch.stack(
+ [load_image(image_dir / f"scene{rows[i]:06d}_v1.png") for i in indices]
+ ).to(args.device)
+ mask = (
+ torch.rand(len(pixels), tokens, generator=generator, device=args.device)
+ < args.mask_ratio
+ )
+ encoded = model.encode(pixels, mask=mask)
+ patches = model.reconstruction(encoded[:, 1:]) # [B, T, p*p*3]
+ grid = manifest["image_size"] // saved["patch"]
+ patch = saved["patch"]
+ image = patches.reshape(len(pixels), grid, grid, patch, patch, 3)
+ image = image.permute(0, 5, 1, 3, 2, 4).reshape(
+ len(pixels), 3, manifest["image_size"], manifest["image_size"]
+ )
+ blended = torch.where(
+ mask.reshape(len(pixels), 1, grid, 1, grid, 1)
+ .expand(-1, 3, -1, patch, -1, patch)
+ .reshape_as(image),
+ image,
+ pixels,
+ )
+ inpainted.append(blended.float().cpu())
+ inpainted = torch.cat(inpainted)
+
+ # Layout-invariant content distance: Chamfer between foreground patch
+ # sets in pixel space. Renders of the same scene share patch CONTENT
+ # (colors, local shape fragments) while positions are resampled, so
+ # per-pixel comparison would measure layout instead.
+ patch = saved["patch"]
+ grid = manifest["image_size"] // patch
+
+ def patch_sets(images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ pieces = images.reshape(len(images), 3, grid, patch, grid, patch)
+ pieces = pieces.permute(0, 2, 4, 1, 3, 5).reshape(len(images), grid * grid, -1)
+ background = pieces.mean(-1)
+ foreground = background > background.min(1, keepdim=True).values + 0.02
+ return pieces, foreground
+
+ inpaint_pieces, inpaint_fg = patch_sets(inpainted)
+ reference_pieces, reference_fg = patch_sets(references)
+ distances = torch.zeros(len(rows), len(rows))
+ device = args.device
+ reference_pieces = reference_pieces.to(device)
+ reference_fg = reference_fg.to(device)
+ for start in range(0, len(rows), 32):
+ stop = min(start + 32, len(rows))
+ block = inpaint_pieces[start:stop].to(device)
+ block_fg = inpaint_fg[start:stop].to(device)
+ cross = torch.cdist(block, reference_pieces.reshape(-1, block.shape[-1]))
+ cross = cross.reshape(stop - start, block.shape[1], len(rows), -1)
+ masked_ab = cross.masked_fill(
+ ~reference_fg[None, None, :, :], float("inf")
+ ).amin(-1)
+ forward = (
+ (masked_ab * block_fg[:, :, None]).sum(1)
+ / block_fg.sum(1).clamp_min(1)[:, None]
+ )
+ masked_ba = cross.masked_fill(
+ ~block_fg[:, :, None, None].expand_as(cross), float("inf")
+ ).amin(1)
+ backward = (
+ (masked_ba * reference_fg[None, :, :]).sum(-1)
+ / reference_fg.sum(-1).clamp_min(1)[None, :]
+ )
+ distances[start:stop] = (0.5 * (forward + backward)).float().cpu()
+ distances[torch.isinf(distances) | torch.isnan(distances)] = distances[
+ torch.isfinite(distances)
+ ].max()
+ return distances
+
+
+@torch.inference_mode()
+def text_kernel(
+ rows: list[int], captions: list[list[str]], args: argparse.Namespace
+) -> 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()
+
+ prefixes = [tokenize(captions[row][0], vocab) for row in rows]
+ # Continuations are the RELATION sentences only: they refer back to
+ # groups the prefix must have introduced, so their likelihood grades
+ # shared content instead of rewarding verbatim repetition.
+ continuations = []
+ for row in rows:
+ sentences = captions[row][1].split(". ")
+ relational = ". ".join(sentences[1:]) if len(sentences) > 1 else sentences[0]
+ continuations.append(tokenize(relational, vocab)[1:]) # drop bos
+ scores = torch.zeros(len(rows), len(rows))
+ jobs = [(i, j) for i in range(len(rows)) for j in range(len(rows))]
+ for indices in tqdm(
+ list(batch_indices(len(jobs), args.batch_size)), desc="text kernel"
+ ):
+ batch = [jobs[k] for k in indices]
+ sequences = [
+ (prefixes[i] + continuations[j])[: saved["context"]] for i, j 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)
+ logits = model.logits(hidden[:, :-1])
+ log_probs = F.log_softmax(logits, dim=-1)
+ for row, (i, j) in enumerate(batch):
+ begin = len(prefixes[i]) - 1
+ end = min(len(prefixes[i]) + len(continuations[j]), longest) - 1
+ if end <= begin:
+ scores[i, j] = 0.0
+ continue
+ targets = tokens[row, begin + 1 : end + 1]
+ picked = log_probs[row, begin:end].gather(1, targets[:, None])
+ scores[i, j] = float(picked.mean())
+ return -scores # negative mean log-likelihood as a distance
+
+
+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]
+
+ visual_distance = vision_kernel(rows, manifest, args)
+ text_distance = text_kernel(rows, captions, args)
+ torch.save(
+ {
+ "rows": rows,
+ "visual_distance": visual_distance,
+ "text_distance": text_distance,
+ },
+ args.kernels_output,
+ )
+
+ # Substitution kernels as standardized relation channels: negative
+ # distances, symmetrized, standardized off-diagonal. The gate machinery
+ # then scores orderings exactly as for any relation field.
+ visual_similarity = -visual_distance
+ text_similarity = -0.5 * (text_distance + text_distance.T)
+ visual_channels = standardize_relation(
+ 0.5 * (visual_similarity + visual_similarity.T).double()
+ )[0][None]
+ text_channels = standardize_relation(text_similarity.double())[0][None]
+ generator = torch.Generator().manual_seed(args.seed)
+ report = {
+ "protocol": (
+ "Substitution kernels are unimodal-predictor functionals "
+ "(masked inpainting distance; caption continuation "
+ "likelihood). Hidden pairs score orderings only."
+ ),
+ "split": args.split,
+ "samples": len(rows),
+ **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()
diff --git a/worldalign/synth_triangle_gate.py b/worldalign/synth_triangle_gate.py
new file mode 100644
index 0000000..35bc2b2
--- /dev/null
+++ b/worldalign/synth_triangle_gate.py
@@ -0,0 +1,296 @@
+"""Third-order (triangle) relational invariants for the synthetic world.
+
+Counterfeits so far satisfy pairwise statistics: they are wrong sections
+that look flat when probed along edges. The holonomy analogue on a
+relation field is a triangle: for a node triple the product-like
+statistic T[a,b,c] combines three edges at once, so preserving pairwise
+marginals no longer suffices. Each node participates in O(N^2) triangles
+rather than O(N) edges, which multiplies the constraint count without
+touching the states.
+
+The triangle energy sums squared cross-modal differences of standardized
+triangle tensors over a fixed random triple sample; the sample is drawn
+once from node indices only. Swap deltas are evaluated exactly on the
+triples touching the swapped nodes. Hidden pairs score orderings only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .synth_cc_battery import (
+ component_descriptors,
+ moment_field,
+ onehot_descriptors,
+ phrase_bow_sets,
+)
+from .synth_towers import load_image
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--triples", type=int, default=200000)
+ parser.add_argument("--pair-weight", type=float, default=1.0)
+ parser.add_argument("--triangle-weight", type=float, default=1.0)
+ parser.add_argument("--random-perms", type=int, default=200)
+ parser.add_argument("--transposition-samples", type=int, default=20000)
+ parser.add_argument("--descent-restarts", type=int, default=3)
+ parser.add_argument("--descent-max-steps", type=int, default=3000)
+ parser.add_argument("--descent-proposals", type=int, default=2048)
+ parser.add_argument("--replicas", type=int, default=8)
+ parser.add_argument("--tempering-rounds", type=int, default=20000)
+ parser.add_argument("--temp-high", type=float, default=3e-3)
+ parser.add_argument("--temp-low", type=float, default=1e-5)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v0/triangle_gate.json"
+ )
+ return parser.parse_args()
+
+
+def standardized(field: torch.Tensor) -> torch.Tensor:
+ mask = ~torch.eye(len(field), dtype=torch.bool, device=field.device)
+ values = field[mask]
+ out = (field - values.mean()) / values.std().clamp_min(1e-9)
+ return out.masked_fill(~mask, 0.0)
+
+
+class TriangleEnergy:
+ """Pairwise-plus-triangle energy over a fixed triple sample."""
+
+ def __init__(
+ self,
+ text: torch.Tensor,
+ visual: torch.Tensor,
+ triples: torch.Tensor,
+ pair_weight: float,
+ triangle_weight: float,
+ ) -> None:
+ self.text = text
+ self.visual = visual
+ self.triples = triples
+ self.pair_weight = pair_weight
+ self.triangle_weight = triangle_weight
+ self.size = len(visual)
+ self.mask = ~torch.eye(self.size, dtype=torch.bool, device=visual.device)
+ self.pair_count = int(self.mask.sum())
+ a, b, c = triples[:, 0], triples[:, 1], triples[:, 2]
+ self.visual_triangle = (
+ visual[a, b] * visual[b, c] * visual[a, c]
+ )
+ # Triples touching each node, for local delta evaluation.
+ self.touch = [[] for _ in range(self.size)]
+ for index, (x, y, z) in enumerate(triples.tolist()):
+ self.touch[x].append(index)
+ self.touch[y].append(index)
+ self.touch[z].append(index)
+ self.touch = [
+ torch.tensor(items, device=visual.device, dtype=torch.long)
+ for items in self.touch
+ ]
+
+ def pairwise(self, permutation: torch.Tensor) -> torch.Tensor:
+ fields = self.text[permutation][:, permutation]
+ return ((fields - self.visual) ** 2)[self.mask].sum() / self.pair_count
+
+ def triangle(self, permutation: torch.Tensor) -> torch.Tensor:
+ fields = self.text[permutation][:, permutation]
+ a, b, c = self.triples[:, 0], self.triples[:, 1], self.triples[:, 2]
+ text_triangle = fields[a, b] * fields[b, c] * fields[a, c]
+ return ((text_triangle - self.visual_triangle) ** 2).mean()
+
+ def total(self, permutation: torch.Tensor) -> float:
+ return float(
+ self.pair_weight * self.pairwise(permutation)
+ + self.triangle_weight * self.triangle(permutation)
+ )
+
+ def swap_delta(self, permutation: torch.Tensor, p: int, q: int) -> float:
+ """Exact delta for swapping positions p and q."""
+ before_pair = self._local_pair(permutation, p, q)
+ indices = torch.cat([self.touch[p], self.touch[q]]).unique()
+ before_triangle = self._local_triangle(permutation, indices)
+ trial = permutation.clone()
+ trial[[p, q]] = trial[[q, p]]
+ after_pair = self._local_pair(trial, p, q)
+ after_triangle = self._local_triangle(trial, indices)
+ return float(
+ self.pair_weight * (after_pair - before_pair) / self.pair_count
+ + self.triangle_weight
+ * (after_triangle - before_triangle)
+ / len(self.triples)
+ )
+
+ def _local_pair(
+ self, permutation: torch.Tensor, p: int, q: int
+ ) -> torch.Tensor:
+ rows = permutation[[p, q]]
+ fields = self.text[rows][:, permutation]
+ visual_rows = self.visual[[p, q]]
+ difference = (fields - visual_rows) ** 2
+ difference[0, p] = 0.0
+ difference[1, q] = 0.0
+ # Rows and columns are symmetric; count row contributions twice and
+ # subtract the doubly counted (p, q) pair once.
+ total = 2.0 * difference.sum()
+ return total - 2.0 * difference[0, q]
+
+ def _local_triangle(
+ self, permutation: torch.Tensor, indices: torch.Tensor
+ ) -> torch.Tensor:
+ triples = self.triples[indices]
+ a, b, c = triples[:, 0], triples[:, 1], triples[:, 2]
+ pa, pb, pc = permutation[a], permutation[b], permutation[c]
+ text_triangle = (
+ self.text[pa, pb] * self.text[pb, pc] * self.text[pa, pc]
+ )
+ return ((text_triangle - self.visual_triangle[indices]) ** 2).sum()
+
+
+def build_fields(args: argparse.Namespace, rows: list[int], manifest: dict) -> tuple:
+ image_dir = Path(manifest["image_dir"])
+ per_view = []
+ for view in range(args.vision_views):
+ raw = []
+ for row in tqdm(rows, desc=f"cc v{view}"):
+ sprites, _ = component_descriptors(
+ load_image(image_dir / f"scene{row:06d}_v{view}.png"),
+ args.merge_distance,
+ )
+ raw.append(sprites)
+ sets = [F.normalize(v, dim=-1) for v in onehot_descriptors(raw)]
+ per_view.append(moment_field(sets))
+ visual_field = torch.stack(per_view).mean(0)
+ captions = read_json(Path(args.data_dir, "captions.json"))["captions"]
+ text_sets = phrase_bow_sets(rows, captions, manifest["vocabulary"])
+ return visual_field, moment_field(text_sets)
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ rows = manifest[args.split][: args.samples]
+ visual_field, text_field = build_fields(args, rows, manifest)
+
+ device = torch.device(args.device)
+ size = len(rows)
+ generator = torch.Generator().manual_seed(args.seed)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden).to(device)
+ text = standardized(text_field[hidden][:, hidden].double().to(device)).float()
+ visual = standardized(visual_field.double().to(device)).float()
+
+ triples = torch.randint(0, size, (args.triples, 3), generator=generator)
+ triples = triples[
+ (triples[:, 0] != triples[:, 1])
+ & (triples[:, 1] != triples[:, 2])
+ & (triples[:, 0] != triples[:, 2])
+ ].to(device)
+ energy = TriangleEnergy(
+ text, visual, triples, args.pair_weight, args.triangle_weight
+ )
+
+ true_total = energy.total(truth)
+ randoms = []
+ for _ in range(args.random_perms):
+ permutation = torch.argsort(torch.rand(size, generator=generator)).to(device)
+ randoms.append(energy.total(permutation))
+ randoms = torch.tensor(randoms)
+ gate_a = {
+ "true": true_total,
+ "random_mean": float(randoms.mean()),
+ "true_z": float((randoms.mean() - true_total) / randoms.std().clamp_min(1e-12)),
+ }
+
+ improving = 0
+ checked = 0
+ for _ in range(min(args.transposition_samples, 4000)):
+ p = int(torch.randint(0, size, (1,), generator=generator))
+ q = int(torch.randint(0, size, (1,), generator=generator))
+ if p == q:
+ continue
+ checked += 1
+ improving += energy.swap_delta(truth, p, q) < 0
+ gate_b = {
+ "sampled": checked,
+ "improving_fraction": improving / max(checked, 1),
+ }
+
+ def descent(start: torch.Tensor) -> dict:
+ current = start.clone()
+ for _ in range(args.descent_max_steps):
+ best_delta, best_pair = 0.0, None
+ for _ in range(64):
+ p = int(torch.randint(0, size, (1,), generator=generator))
+ q = int(torch.randint(0, size, (1,), generator=generator))
+ if p == q:
+ continue
+ delta = energy.swap_delta(current, p, q)
+ if delta < best_delta:
+ best_delta, best_pair = delta, (p, q)
+ if best_pair is None:
+ break
+ p, q = best_pair
+ current[[p, q]] = current[[q, p]]
+ return {
+ "accuracy": float((current.cpu() == truth.cpu()).float().mean()),
+ "energy": energy.total(current),
+ }
+
+ from_true = descent(truth)
+ restarts = [
+ descent(torch.argsort(torch.rand(size, generator=generator)).to(device))
+ for _ in range(args.descent_restarts)
+ ]
+ best_random = min(r["energy"] for r in restarts)
+
+ report = {
+ "protocol": (
+ "Pairwise-plus-triangle energy on moment set-kernel fields; "
+ "triples sampled from node indices only. Hidden pairs score "
+ "orderings."
+ ),
+ "samples": size,
+ "triples": len(triples),
+ "weights": {
+ "pair": args.pair_weight,
+ "triangle": args.triangle_weight,
+ },
+ "gate_a": gate_a,
+ "gate_b": gate_b,
+ "descent_from_true": from_true,
+ "descent_from_random": restarts,
+ "verdict": {
+ "true_energy": true_total,
+ "best_random_descent": best_random,
+ "counterfeit_found": bool(
+ best_random < true_total
+ and min(r["accuracy"] for r in restarts) < 0.5
+ ),
+ "descent_keeps": from_true["accuracy"],
+ "true_z": gate_a["true_z"],
+ "improving_fraction": gate_b["improving_fraction"],
+ },
+ }
+ print(json.dumps({"verdict": report["verdict"]}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_triangle_recovery.py b/worldalign/synth_triangle_recovery.py
new file mode 100644
index 0000000..a9bca61
--- /dev/null
+++ b/worldalign/synth_triangle_recovery.py
@@ -0,0 +1,173 @@
+"""Blind recovery under the triangle energy.
+
+The triangle gate holds the true assignment at N=512 with no counterfeit
+below it. The search question is separate and now worth asking again:
+tempering with third-order swap deltas, from random starts, with the
+hidden order behind a shuffle. Recovery accuracy is the end-to-end world
+matching measurement.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+
+from .common import read_json, seed_everything, write_json
+from .synth_triangle_gate import TriangleEnergy, build_fields, standardized
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v0")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--vision-views", type=int, default=4)
+ parser.add_argument("--triples", type=int, default=400000)
+ parser.add_argument("--pair-weight", type=float, default=1.0)
+ parser.add_argument("--triangle-weight", type=float, default=1.0)
+ parser.add_argument("--replicas", type=int, default=6)
+ parser.add_argument("--rounds", type=int, default=4000)
+ parser.add_argument("--proposals", type=int, default=48)
+ parser.add_argument("--temp-high", type=float, default=2e-2)
+ parser.add_argument("--temp-low", type=float, default=1e-4)
+ parser.add_argument("--exchange-every", type=int, default=20)
+ parser.add_argument("--greedy-rounds", type=int, default=3000)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v0/triangle_recovery.json"
+ )
+ return parser.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(Path(args.data_dir, "manifest.json"))
+ rows = manifest[args.split][: args.samples]
+ visual_field, text_field = build_fields(args, rows, manifest)
+
+ device = torch.device(args.device)
+ size = len(rows)
+ generator = torch.Generator().manual_seed(args.seed)
+ hidden = torch.randperm(size, generator=generator)
+ truth = torch.argsort(hidden).to(device)
+ text = standardized(text_field[hidden][:, hidden].double().to(device)).float()
+ visual = standardized(visual_field.double().to(device)).float()
+
+ triples = torch.randint(0, size, (args.triples, 3), generator=generator)
+ triples = triples[
+ (triples[:, 0] != triples[:, 1])
+ & (triples[:, 1] != triples[:, 2])
+ & (triples[:, 0] != triples[:, 2])
+ ].to(device)
+ energy = TriangleEnergy(
+ text, visual, triples, args.pair_weight, args.triangle_weight
+ )
+ true_energy = energy.total(truth)
+
+ temperatures = torch.logspace(
+ torch.log10(torch.tensor(args.temp_low)),
+ torch.log10(torch.tensor(args.temp_high)),
+ args.replicas,
+ )
+ states = [
+ torch.argsort(torch.rand(size, generator=generator)).to(device)
+ for _ in range(args.replicas)
+ ]
+ energies = [energy.total(state) for state in states]
+ history = []
+ for round_index in range(args.rounds):
+ for replica in range(args.replicas):
+ temperature = float(temperatures[replica])
+ for _ in range(args.proposals):
+ p = int(torch.randint(0, size, (1,), generator=generator))
+ q = int(torch.randint(0, size, (1,), generator=generator))
+ if p == q:
+ continue
+ delta = energy.swap_delta(states[replica], p, q)
+ threshold = -temperature * float(
+ torch.rand(1, generator=generator).clamp_min(1e-12).log()
+ )
+ if delta < threshold:
+ states[replica][[p, q]] = states[replica][[q, p]]
+ energies[replica] += delta
+ if round_index % args.exchange_every == 0:
+ for replica in range(args.replicas - 1):
+ gap = (energies[replica] - energies[replica + 1]) * (
+ 1.0 / float(temperatures[replica])
+ - 1.0 / float(temperatures[replica + 1])
+ )
+ if gap > 0 or float(torch.rand(1, generator=generator)) < min(
+ 1.0, float(torch.tensor(gap).exp())
+ ):
+ states[replica], states[replica + 1] = (
+ states[replica + 1],
+ states[replica],
+ )
+ energies[replica], energies[replica + 1] = (
+ energies[replica + 1],
+ energies[replica],
+ )
+ if round_index % 200 == 0:
+ cold = min(range(args.replicas), key=lambda r: energies[r])
+ record = {
+ "round": round_index,
+ "cold_energy": energies[cold],
+ "cold_accuracy": float(
+ (states[cold].cpu() == truth.cpu()).float().mean()
+ ),
+ "energy_over_true": energies[cold] / true_energy - 1.0,
+ }
+ history.append(record)
+ print(json.dumps(record))
+
+ # Greedy polish of the coldest replica.
+ cold = min(range(args.replicas), key=lambda r: energies[r])
+ current = states[cold].clone()
+ for _ in range(args.greedy_rounds):
+ best_delta, best_pair = 0.0, None
+ for _ in range(96):
+ p = int(torch.randint(0, size, (1,), generator=generator))
+ q = int(torch.randint(0, size, (1,), generator=generator))
+ if p == q:
+ continue
+ delta = energy.swap_delta(current, p, q)
+ if delta < best_delta:
+ best_delta, best_pair = delta, (p, q)
+ if best_pair is None:
+ break
+ p, q = best_pair
+ current[[p, q]] = current[[q, p]]
+
+ final = {
+ "accuracy": float((current.cpu() == truth.cpu()).float().mean()),
+ "energy": energy.total(current),
+ }
+ final["energy_over_true"] = final["energy"] / true_energy - 1.0
+ report = {
+ "protocol": (
+ "Tempering on the pairwise-plus-triangle energy from random "
+ "starts; text order hidden behind a shuffle; hidden truth "
+ "scores the outcome only."
+ ),
+ "samples": size,
+ "true_energy": true_energy,
+ "chance_accuracy": 1.0 / size,
+ "history": history,
+ "final_polished": final,
+ "replica_accuracies": [
+ float((state.cpu() == truth.cpu()).float().mean()) for state in states
+ ],
+ }
+ print(json.dumps({"final": final}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/synth_world.py b/worldalign/synth_world.py
new file mode 100644
index 0000000..8875fde
--- /dev/null
+++ b/worldalign/synth_world.py
@@ -0,0 +1,506 @@
+"""Synthetic closed world v0: procedural scenes with exact closure.
+
+A scene's discrete state is a set of object groups (count, size, color,
+shape) plus a set of spatial relation constraints between groups. Captions
+mention exactly the discrete state; renders realize it with continuous
+nuisance (positions, jitter) resampled per view. Shared content and
+modality-private variation are therefore separated by construction, and
+every dial -- vocabulary, ontology size, groups per scene, scene count,
+orbit multiplicity, intervention density -- is a generator argument.
+
+Outputs mirror the Flickr pipeline layout: a manifest with disjoint
+vision-only and text-only scene rows plus held-out val/test, an image
+directory with `sceneNNNNNN_vK.png` views, and per-scene caption lists.
+Intervention variants with edit metadata are stored for the response
+battery; nothing downstream reads them yet.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import math
+import random
+from pathlib import Path
+
+from PIL import Image, ImageDraw
+
+from .common import write_json
+
+COLORS = {
+ "red": (205, 49, 49),
+ "orange": (224, 133, 44),
+ "yellow": (229, 213, 74),
+ "green": (64, 168, 75),
+ "blue": (59, 104, 214),
+ "purple": (139, 72, 190),
+ "pink": (228, 136, 179),
+ "white": (238, 238, 238),
+ "gray": (140, 140, 140),
+ "brown": (125, 84, 48),
+}
+SHAPES = ("circle", "square", "triangle", "star", "diamond", "cross")
+SIZES = {"small": (7, 10), "medium": (13, 17), "large": (21, 26)}
+COUNT_WORDS = {1: "one", 2: "two", 3: "three", 4: "four"}
+RELATIONS = ("left of", "right of", "above", "below")
+BACKGROUND = (24, 24, 28)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--output-dir", default="artifacts/synth_v0")
+ parser.add_argument("--image-size", type=int, default=128)
+ parser.add_argument("--vision-only", type=int, default=12000)
+ parser.add_argument("--text-only", type=int, default=12000)
+ parser.add_argument("--val", type=int, default=1000)
+ parser.add_argument("--test", type=int, default=1000)
+ parser.add_argument("--min-groups", type=int, default=2)
+ parser.add_argument("--max-groups", type=int, default=4)
+ parser.add_argument("--visual-views", type=int, default=4)
+ parser.add_argument("--captions", type=int, default=4)
+ parser.add_argument("--interventions", type=int, default=2,
+ help="Edited variants per eval scene.")
+ parser.add_argument("--seed", type=int, default=20260730)
+ parser.add_argument(
+ "--skew",
+ type=float,
+ default=0.0,
+ help="Zipf exponent for colour and shape sampling. Non-uniform "
+ "marginals make the cross-modal value correspondence recoverable "
+ "from frequency alone, with no declared dictionary.",
+ )
+ parser.add_argument(
+ "--correlate",
+ type=float,
+ default=0.0,
+ help="Colour-shape coupling strength. Real worlds pair values "
+ "(bananas are yellow), which makes co-occurrence structure "
+ "informative where bare marginals are not.",
+ )
+ parser.add_argument(
+ "--report-bias",
+ type=float,
+ default=0.0,
+ help="Reporting bias strength. Captions mention rarer colours "
+ "preferentially, so text frequency stops tracking pixel "
+ "frequency -- the situation in real corpora, where nobody "
+ "describes the grey road.",
+ )
+ parser.add_argument(
+ "--texture",
+ action="store_true",
+ help="Dot-textured object fills: within-object structure that "
+ "gives binding objectives a purchase, without changing the "
+ "discrete state or the captions.",
+ )
+ return parser.parse_args()
+
+
+_PROFILE_CACHE: dict[tuple[str, float], list[float]] = {}
+
+
+def colour_shape_profile(colour: str, concentration: float) -> list[float]:
+ """Fixed per-colour shape distribution, deterministic in the colour."""
+ key = (colour, concentration)
+ if key not in _PROFILE_CACHE:
+ local = random.Random(hash(key) & 0xFFFFFFFF)
+ raw = [local.gammavariate(1.0 / concentration, 1.0) for _ in SHAPES]
+ total = sum(raw) or 1.0
+ _PROFILE_CACHE[key] = [value / total for value in raw]
+ return _PROFILE_CACHE[key]
+
+
+def zipf_weights(count: int, exponent: float) -> list[float]:
+ raw = [1.0 / (rank + 1) ** exponent for rank in range(count)]
+ total = sum(raw)
+ return [value / total for value in raw]
+
+
+def sample_group(
+ rng: random.Random, skew: float = 0.0, correlate: float = 0.0
+) -> dict:
+ if skew > 0.0:
+ colors = rng.choices(
+ list(COLORS), weights=zipf_weights(len(COLORS), skew)
+ )[0]
+ else:
+ colors = rng.choice(list(COLORS))
+ if correlate > 0.0:
+ # Each colour carries its own shape distribution, drawn once from a
+ # Dirichlet with concentration 1/correlate. Distinct profiles make
+ # the joint table identifying, as in a real world where objects of
+ # a kind take characteristic forms.
+ weights = colour_shape_profile(colors, correlate)
+ if skew > 0.0:
+ base = zipf_weights(len(SHAPES), skew)
+ weights = [w * b for w, b in zip(weights, base)]
+ shapes = rng.choices(SHAPES, weights=weights)[0]
+ elif skew > 0.0:
+ shapes = rng.choices(
+ SHAPES, weights=zipf_weights(len(SHAPES), skew)
+ )[0]
+ else:
+ shapes = rng.choice(SHAPES)
+ return {
+ "count": rng.randint(1, 4),
+ "size": rng.choice(list(SIZES)),
+ "color": colors,
+ "shape": shapes,
+ }
+
+
+def sample_scene(rng: random.Random, args: argparse.Namespace) -> dict:
+ skew = getattr(args, "skew", 0.0)
+ correlate = getattr(args, "correlate", 0.0)
+ groups = [
+ sample_group(rng, skew, correlate)
+ for _ in range(rng.randint(args.min_groups, args.max_groups))
+ ]
+ # Reject referential ambiguity: color-shape pairs are unique per scene,
+ # so every relational mention has exactly one referent.
+ signatures = [(group["color"], group["shape"]) for group in groups]
+ while len(set(signatures)) < len(signatures):
+ groups = [sample_group(rng, skew, correlate) for _ in range(len(groups))]
+ signatures = [(group["color"], group["shape"]) for group in groups]
+ relation_count = rng.randint(1, min(3, len(groups) * (len(groups) - 1) // 2))
+ pairs = [(a, b) for a in range(len(groups)) for b in range(len(groups)) if a < b]
+ rng.shuffle(pairs)
+ relations = [
+ {"a": a, "b": b, "relation": rng.choice(RELATIONS)}
+ for a, b in pairs[:relation_count]
+ ]
+ return {"groups": groups, "relations": relations}
+
+
+def relation_holds(relation: str, pa: tuple[float, float], pb: tuple[float, float], margin: float) -> bool:
+ if relation == "left of":
+ return pa[0] < pb[0] - margin
+ if relation == "right of":
+ return pa[0] > pb[0] + margin
+ if relation == "above":
+ return pa[1] < pb[1] - margin
+ return pa[1] > pb[1] + margin
+
+
+def group_geometry(group: dict, rng: random.Random, size: int) -> dict:
+ low, high = SIZES[group["size"]]
+ shrink = (1.0, 0.95, 0.85, 0.75)[group["count"] - 1]
+ radius = rng.uniform(low, high) * size / 128.0 * shrink
+ count = group["count"]
+ if count == 1:
+ offsets = [(0.0, 0.0)]
+ extent = radius
+ else:
+ # Ring placement guarantees exact visible multiplicity: adjacent
+ # spacing 2 s sin(pi/k) stays above 2.15 r by construction.
+ ring = 1.10 * 1.075 * radius / math.sin(math.pi / count)
+ phase = rng.uniform(0.0, 2.0 * math.pi)
+ offsets = [
+ (
+ ring * math.cos(phase + 2.0 * math.pi * k / count),
+ ring * math.sin(phase + 2.0 * math.pi * k / count),
+ )
+ for k in range(count)
+ ]
+ extent = ring + radius
+ return {"radius": radius, "offsets": offsets, "extent": extent}
+
+
+def worst_case_extent(group: dict, size: int) -> float:
+ high = SIZES[group["size"]][1] * size / 128.0
+ shrink = (1.0, 0.95, 0.85, 0.75)[group["count"] - 1]
+ radius = high * shrink
+ if group["count"] == 1:
+ return radius
+ return 1.10 * 1.075 * radius / math.sin(math.pi / group["count"]) + radius
+
+
+def place_groups(
+ scene: dict,
+ rng: random.Random,
+ size: int,
+ extents: list[float] | None = None,
+) -> list[tuple[float, float]] | None:
+ margin = size * 0.08
+ if extents is None:
+ extents = [worst_case_extent(group, size) for group in scene["groups"]]
+ for _ in range(300):
+ centers = []
+ feasible = True
+ for extent in extents:
+ low, high = extent + 2.0, size - extent - 2.0
+ if low >= high:
+ feasible = False
+ break
+ centers.append((rng.uniform(low, high), rng.uniform(low, high)))
+ if not feasible:
+ return None
+ if any(
+ math.dist(centers[a], centers[b]) < extents[a] + extents[b] + 4.0
+ for a in range(len(centers))
+ for b in range(a + 1, len(centers))
+ ):
+ continue
+ if all(
+ relation_holds(r["relation"], centers[r["a"]], centers[r["b"]], margin)
+ for r in scene["relations"]
+ ):
+ return centers
+ return None
+
+
+def draw_shape(draw: ImageDraw.ImageDraw, shape: str, x: float, y: float, radius: float, fill: tuple) -> None:
+ if shape == "circle":
+ draw.ellipse([x - radius, y - radius, x + radius, y + radius], fill=fill)
+ elif shape == "square":
+ draw.rectangle([x - radius, y - radius, x + radius, y + radius], fill=fill)
+ elif shape == "triangle":
+ draw.polygon(
+ [(x, y - radius), (x - radius, y + radius), (x + radius, y + radius)],
+ fill=fill,
+ )
+ elif shape == "diamond":
+ draw.polygon(
+ [(x, y - radius), (x + radius, y), (x, y + radius), (x - radius, y)],
+ fill=fill,
+ )
+ elif shape == "cross":
+ arm = radius * 0.42
+ draw.rectangle([x - arm, y - radius, x + arm, y + radius], fill=fill)
+ draw.rectangle([x - radius, y - arm, x + radius, y + arm], fill=fill)
+ else: # star
+ points = []
+ for k in range(10):
+ r = radius if k % 2 == 0 else radius * 0.45
+ angle = -math.pi / 2 + k * math.pi / 5
+ points.append((x + r * math.cos(angle), y + r * math.sin(angle)))
+ draw.polygon(points, fill=fill)
+
+
+def render_scene(scene: dict, rng: random.Random, size: int, texture: bool = False) -> Image.Image | None:
+ geometries = [group_geometry(group, rng, size) for group in scene["groups"]]
+ centers = place_groups(
+ scene, rng, size, extents=[g["extent"] for g in geometries]
+ )
+ if centers is None:
+ return None
+ shade = rng.randint(-6, 6)
+ image = Image.new("RGB", (size, size), tuple(c + shade for c in BACKGROUND))
+ draw = ImageDraw.Draw(image)
+ for group, geometry, center in zip(scene["groups"], geometries, centers):
+ base = tuple(
+ min(255, max(0, channel + rng.randint(-10, 10)))
+ for channel in COLORS[group["color"]]
+ )
+ for off in geometry["offsets"]:
+ draw_shape(
+ draw,
+ group["shape"],
+ center[0] + off[0],
+ center[1] + off[1],
+ geometry["radius"],
+ base,
+ )
+ return image
+
+
+def group_phrase(group: dict, rng: random.Random) -> str:
+ size_word = "" if group["size"] == "medium" else group["size"] + " "
+ plural = "es" if group["shape"] == "cross" else "s"
+ noun = group["shape"] + (plural if group["count"] > 1 else "")
+ count_word = COUNT_WORDS[group["count"]] if group["count"] > 1 else (
+ "a" if size_word == "" or size_word[0] not in "aeiou" else "an"
+ )
+ return f"{count_word} {size_word}{group['color']} {noun}"
+
+
+def caption_scene(
+ scene: dict, rng: random.Random, report_bias: float = 0.0
+) -> str:
+ order = list(range(len(scene["groups"])))
+ rng.shuffle(order)
+ if report_bias > 0.0 and len(order) > 1:
+ # Mention probability falls with the colour's population frequency,
+ # so the text marginal inverts the pixel marginal.
+ ranks = {name: index for index, name in enumerate(COLORS)}
+ keep = []
+ for index in order:
+ rank = ranks[scene["groups"][index]["color"]]
+ salience = ((rank + 1) / len(COLORS)) ** report_bias
+ if rng.random() < 0.25 + 0.75 * salience:
+ keep.append(index)
+ order = keep or order[:1]
+ phrases = [group_phrase(scene["groups"][index], rng) for index in order]
+ if len(phrases) > 1:
+ listed = ", ".join(phrases[:-1]) + " and " + phrases[-1]
+ else:
+ listed = phrases[0]
+ opener = rng.choice(["there are", "the picture shows", "you can see"])
+ sentences = [f"{opener} {listed}."]
+ relations = list(scene["relations"])
+ rng.shuffle(relations)
+ mentioned = set(order)
+ relations = [
+ relation
+ for relation in relations
+ if relation["a"] in mentioned and relation["b"] in mentioned
+ ]
+ for relation in relations:
+ subject = group_phrase(scene["groups"][relation["a"]], rng)
+ target = group_phrase(scene["groups"][relation["b"]], rng)
+ sentences.append(f"the {subject.split(' ', 1)[1]} {'are' if scene['groups'][relation['a']]['count'] > 1 else 'is'} {relation['relation']} the {target.split(' ', 1)[1]}.")
+ return " ".join(sentences)
+
+
+def intervene(scene: dict, rng: random.Random) -> tuple[dict, dict]:
+ edited = json.loads(json.dumps(scene))
+ kinds = ["recolor", "count", "remove", "relation"]
+ if len(edited["groups"]) <= 2:
+ kinds.remove("remove")
+ if not edited["relations"]:
+ kinds.remove("relation")
+ kind = rng.choice(kinds)
+ if kind == "recolor":
+ index = rng.randrange(len(edited["groups"]))
+ old = edited["groups"][index]["color"]
+ shape = edited["groups"][index]["shape"]
+ taken = {
+ g["color"]
+ for i, g in enumerate(edited["groups"])
+ if i != index and g["shape"] == shape
+ }
+ edited["groups"][index]["color"] = rng.choice(
+ [c for c in COLORS if c != old and c not in taken]
+ )
+ detail = {"kind": kind, "group": index, "from": old, "to": edited["groups"][index]["color"]}
+ elif kind == "count":
+ index = rng.randrange(len(edited["groups"]))
+ old = edited["groups"][index]["count"]
+ edited["groups"][index]["count"] = old % 4 + 1
+ detail = {"kind": kind, "group": index, "from": old, "to": edited["groups"][index]["count"]}
+ elif kind == "remove":
+ index = rng.randrange(len(edited["groups"]))
+ edited["groups"].pop(index)
+ edited["relations"] = [
+ r for r in edited["relations"] if r["a"] != index and r["b"] != index
+ ]
+ for relation in edited["relations"]:
+ relation["a"] -= relation["a"] > index
+ relation["b"] -= relation["b"] > index
+ detail = {"kind": kind, "group": index}
+ else:
+ index = rng.randrange(len(edited["relations"]))
+ old = edited["relations"][index]["relation"]
+ flip = {"left of": "right of", "right of": "left of", "above": "below", "below": "above"}
+ edited["relations"][index]["relation"] = flip[old]
+ detail = {"kind": kind, "relation_index": index, "from": old, "to": flip[old]}
+ return edited, detail
+
+
+def main() -> None:
+ args = parse_args()
+ rng = random.Random(args.seed)
+ output = Path(args.output_dir)
+ (output / "images").mkdir(parents=True, exist_ok=True)
+
+ total = args.vision_only + args.text_only + args.val + args.test
+ scenes, captions = [], []
+ while len(scenes) < total:
+ scene = sample_scene(rng, args)
+ if place_groups(scene, rng, args.image_size) is None:
+ continue
+ scenes.append(scene)
+ captions.append(
+ [
+ caption_scene(scene, rng, args.report_bias)
+ for _ in range(args.captions)
+ ]
+ )
+
+ rows = list(range(total))
+ vision_only = rows[: args.vision_only]
+ text_only = rows[args.vision_only : args.vision_only + args.text_only]
+ val = rows[args.vision_only + args.text_only : args.vision_only + args.text_only + args.val]
+ test = rows[-args.test :]
+
+ needs_render = set(vision_only) | set(val) | set(test)
+ rendered = 0
+ for row in sorted(needs_render):
+ for view in range(args.visual_views):
+ image = None
+ while image is None:
+ image = render_scene(
+ scenes[row], rng, args.image_size, texture=args.texture
+ )
+ image.save(output / "images" / f"scene{row:06d}_v{view}.png")
+ rendered += 1
+
+ interventions = []
+ for row in val + test:
+ for _ in range(args.interventions):
+ edited, detail = intervene(scenes[row], rng)
+ if place_groups(edited, rng, args.image_size) is None:
+ continue
+ index = len(interventions)
+ image = None
+ while image is None:
+ image = render_scene(
+ edited, rng, args.image_size, texture=args.texture
+ )
+ image.save(output / "images" / f"edit{index:06d}.png")
+ interventions.append(
+ {
+ "index": index,
+ "row": row,
+ "detail": detail,
+ "caption": caption_scene(edited, rng, args.report_bias),
+ }
+ )
+
+ vocabulary = sorted(
+ {
+ token
+ for caption_list in captions
+ for caption in caption_list
+ for token in caption.replace(",", " ").replace(".", " ").split()
+ }
+ )
+ manifest = {
+ "dataset": "synth_v0",
+ "image_dir": str(output / "images"),
+ "image_size": args.image_size,
+ "seed": args.seed,
+ "visual_views": args.visual_views,
+ "vision_only_train": vision_only,
+ "text_only_train": text_only,
+ "paired_train": [],
+ "val": val,
+ "test": test,
+ "all_rows": total,
+ "vocabulary_size": len(vocabulary),
+ "vocabulary": vocabulary,
+ "protocol": (
+ "vision_only_train and text_only_train are disjoint scene rows; "
+ "val/test pairs are held out for evaluation only. Captions "
+ "mention exactly the discrete scene state."
+ ),
+ }
+ write_json(output / "manifest.json", manifest)
+ write_json(output / "scenes.private.json", {"scenes": scenes})
+ write_json(output / "captions.json", {"captions": captions})
+ write_json(output / "interventions.private.json", {"interventions": interventions})
+ print(
+ json.dumps(
+ {
+ "scenes": total,
+ "rendered_images": rendered,
+ "interventions": len(interventions),
+ "vocabulary_size": len(vocabulary),
+ "example_caption": captions[test[0]][0],
+ }
+ )
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/tier0_dictionary.py b/worldalign/tier0_dictionary.py
new file mode 100644
index 0000000..96fcbac
--- /dev/null
+++ b/worldalign/tier0_dictionary.py
@@ -0,0 +1,324 @@
+"""Tier 0: derive the cross-modal value correspondence, never declare it.
+
+A declared lexicon ("red" means hue 0) is a hand-supplied cross-modal
+prior. Tier 0 forbids it. Every factor value correspondence is instead
+recovered from unimodal statistics of the two disjoint training splits:
+
+- ordered factors (count, size) match by their intrinsic order, with the
+ small residual ambiguity enumerated and settled by the alignment
+ criterion rather than by assertion;
+- unordered factors (colour, shape) match by marginal frequency rank,
+ which is a unimodal observable on both sides. This works exactly when
+ the world's factor marginals are non-uniform -- true of real corpora
+ and of the skewed synthetic world, false of the uniform one, where the
+ correspondence is information-theoretically unrecoverable.
+
+The recovered dictionary is then applied to build comparable object
+descriptors. Hidden pairs are used only to report how many entries the
+derivation got right; nothing here reads them.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from collections import Counter
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .synth_cc_battery import component_descriptors
+from .synth_set_battery import parse_group_phrases
+from .synth_towers import load_image
+
+SINGULAR_ARTICLES = ("a", "an")
+STOPWORDS = {
+ "there", "are", "the", "picture", "shows", "you", "can", "see",
+ "is", "of", "left", "right", "above", "below", "and",
+}
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v1")
+ parser.add_argument("--text-scenes", type=int, default=6000)
+ parser.add_argument("--vision-scenes", type=int, default=3000)
+ parser.add_argument("--merge-distance", type=float, default=30.0)
+ parser.add_argument("--colour-classes", type=int, default=10)
+ parser.add_argument("--shape-classes", type=int, default=6)
+ parser.add_argument("--seed", type=int, default=20260731)
+ parser.add_argument(
+ "--output", default="artifacts/synth_v1/tier0_dictionary.json"
+ )
+ return parser.parse_args()
+
+
+def text_token_statistics(
+ captions: list[list[str]], rows: list[int]
+) -> dict:
+ """Group-slot token frequencies and singular/plural association.
+
+ Everything here reads captions only. Number words are identified by
+ their association with singular noun forms and their mutual exclusion
+ inside a group phrase; the remaining modifier vocabulary splits into
+ frequency-ranked classes.
+ """
+ phrase_total = 0
+ slot_counts: Counter = Counter()
+ with_singular: Counter = Counter()
+ total_singular = 0
+ cooccurrence: Counter = Counter()
+ for row in rows:
+ for phrase in parse_group_phrases(captions[row][0]):
+ tokens = [
+ token
+ for token in re.findall(r"[a-z]+|\d+", phrase.lower())
+ if token not in STOPWORDS
+ ]
+ if not tokens:
+ continue
+ phrase_total += 1
+ head = tokens[-1]
+ singular = not head.endswith("s")
+ total_singular += singular
+ modifiers = tokens[:-1]
+ for token in modifiers:
+ slot_counts[token] += 1
+ with_singular[token] += singular
+ for first in modifiers:
+ for second in modifiers:
+ if first != second:
+ cooccurrence[(first, second)] += 1
+ return {
+ "slot_counts": slot_counts,
+ "with_singular": with_singular,
+ "cooccurrence": cooccurrence,
+ "total_singular": total_singular,
+ "phrase_total": phrase_total,
+ }
+
+
+def partition_text_vocabulary(stats: dict) -> dict:
+ """Discover modifier families by mutual exclusivity, then name them.
+
+ Tokens of one factor never co-occur inside a group phrase, so families
+ are maximal mutually exclusive sets: greedily place each token in the
+ first family none of whose members it ever co-occurs with. Families are
+ then identified by two unimodal signals -- coverage (how many phrases
+ carry a member) and morphological association (whether the choice
+ predicts the head noun's plural suffix). No token list is declared.
+ """
+ counts = stats["slot_counts"]
+ cooccurrence = stats["cooccurrence"]
+ phrases = stats["phrase_total"]
+ ordered = sorted(counts, key=lambda token: -counts[token])
+ families: list[list[str]] = []
+ for token in ordered:
+ for family in families:
+ if all(
+ cooccurrence[(token, member)] == 0
+ and cooccurrence[(member, token)] == 0
+ for member in family
+ ):
+ family.append(token)
+ break
+ else:
+ families.append([token])
+ described = []
+ for family in families:
+ coverage = sum(counts[token] for token in family) / max(phrases, 1)
+ rates = [
+ stats["with_singular"][token] / max(counts[token], 1)
+ for token in family
+ ]
+ described.append(
+ {
+ "tokens": sorted(family, key=lambda token: -counts[token]),
+ "coverage": coverage,
+ "morphology_spread": float(max(rates) - min(rates)),
+ }
+ )
+ described.sort(key=lambda item: -item["coverage"])
+ # The count family is the near-complete family whose choice predicts the
+ # plural suffix; the other near-complete family is the dominant
+ # unordered attribute; partial families are optional modifiers.
+ complete = [item for item in described if item["coverage"] > 0.8]
+ partial = [item for item in described if item["coverage"] <= 0.8]
+ complete.sort(key=lambda item: -item["morphology_spread"])
+ count_family = complete[0]["tokens"] if complete else []
+ attribute_families = [item["tokens"] for item in complete[1:]]
+ return {
+ "count_words": count_family,
+ "colour_words": attribute_families[0] if attribute_families else [],
+ "other_attribute_words": attribute_families[1:],
+ "size_words": [item["tokens"] for item in partial],
+ "families_detail": described,
+ }
+
+
+def component_raw(image: torch.Tensor, merge_distance: float) -> dict:
+ """Connected-component groups with raw appearance, no colour rules.
+
+ Returns mean RGB, area fraction, and member count per group. Nothing
+ here quantises colour, so no declared hue boundary enters Tier 0.
+ """
+ from scipy import ndimage
+
+ array = image.permute(1, 2, 0).numpy()
+ background = np.median(array.reshape(-1, 3), axis=0)
+ foreground = np.abs(array - background).sum(-1) > 0.12
+ labels, count = ndimage.label(foreground)
+ if count == 0:
+ return {"rgb": np.zeros((0, 3)), "area": np.zeros(0), "members": np.zeros(0)}
+ centers = np.array(
+ ndimage.center_of_mass(foreground, labels, range(1, count + 1))
+ )
+ parent = list(range(count))
+
+ def find(a: int) -> int:
+ while parent[a] != a:
+ parent[a] = parent[parent[a]]
+ a = parent[a]
+ return a
+
+ for a in range(count):
+ for b in range(a + 1, count):
+ if np.linalg.norm(centers[a] - centers[b]) < merge_distance:
+ parent[find(a)] = find(b)
+ groups: dict[int, list[int]] = {}
+ for a in range(count):
+ groups.setdefault(find(a), []).append(a)
+ rgb, area, members = [], [], []
+ total = foreground.size
+ for group in groups.values():
+ mask = np.isin(labels, [m + 1 for m in group])
+ rgb.append(array[mask].mean(0))
+ area.append(float(mask.sum()) / total)
+ members.append(len(group))
+ return {
+ "rgb": np.stack(rgb),
+ "area": np.array(area),
+ "members": np.array(members, dtype=np.int64),
+ }
+
+
+def vision_value_statistics(
+ rows: list[int], manifest: dict, args: argparse.Namespace
+) -> dict:
+ """Colour-class and size-class frequencies from pixels alone.
+
+ Object colours are clustered in hue-saturation-value space with the
+ requested number of classes; class identity is arbitrary, only the
+ frequency ranking is used downstream.
+ """
+ from sklearn.cluster import KMeans
+
+ image_dir = Path(manifest["image_dir"])
+ raw = [
+ component_raw(
+ load_image(image_dir / f"scene{row:06d}_v0.png"), args.merge_distance
+ )
+ for row in tqdm(rows, desc="vision values")
+ ]
+ rgb = np.concatenate([item["rgb"] for item in raw])
+ # Cluster raw appearance: class boundaries come from the data, not from
+ # a declared hue table.
+ clusters = KMeans(
+ n_clusters=args.colour_classes, n_init=10, random_state=args.seed
+ ).fit(rgb)
+ labels = clusters.labels_
+ return {
+ "colour_frequency": Counter(labels.tolist()),
+ "colour_labels": labels,
+ "cluster_centres": clusters.cluster_centers_,
+ "raw": raw,
+ "clusters": clusters,
+ }
+
+
+def frequency_rank_map(
+ text_words: list[str], vision_frequency: Counter, classes: int
+) -> dict:
+ """Match unordered values by descending marginal frequency."""
+ vision_ranked = [
+ label for label, _ in vision_frequency.most_common(classes)
+ ]
+ return {
+ word: vision_ranked[index]
+ for index, word in enumerate(text_words[:classes])
+ }
+
+
+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"]
+
+ text_rows = manifest["text_only_train"][: args.text_scenes]
+ stats = text_token_statistics(captions, text_rows)
+ families = partition_text_vocabulary(stats)
+
+ vision_rows = manifest["vision_only_train"][: args.vision_scenes]
+ vision = vision_value_statistics(vision_rows, manifest, args)
+
+ dictionary = frequency_rank_map(
+ families["colour_words"], vision["colour_frequency"], args.colour_classes
+ )
+
+ report = {
+ "protocol": (
+ "Factor families and their value correspondence are derived "
+ "from unimodal statistics of disjoint splits: noun-number "
+ "association separates count words, coverage separates colour "
+ "from size words, and marginal frequency rank pairs colour "
+ "values across modalities. No declared lexicon."
+ ),
+ "text_families": {
+ key: families[key]
+ for key in ("count_words", "colour_words", "size_words")
+ },
+ "families_detail": families["families_detail"],
+ "vision_colour_frequency": [
+ [int(label), int(count)]
+ for label, count in vision["colour_frequency"].most_common()
+ ],
+ "derived_colour_map": {
+ word: int(label) for word, label in dictionary.items()
+ },
+ }
+
+ # Evaluation only: how many derived entries are semantically right?
+ scenes = read_json(Path(args.data_dir, "scenes.private.json"))["scenes"]
+ truth_frequency = Counter(
+ group["color"] for row in vision_rows for group in scenes[row]["groups"]
+ )
+ truth_rank = [name for name, _ in truth_frequency.most_common()]
+ text_frequency = Counter(
+ group["color"] for row in text_rows for group in scenes[row]["groups"]
+ )
+ text_rank = [name for name, _ in text_frequency.most_common()]
+ report["evaluation_only"] = {
+ "vision_side_truth_frequency_rank": truth_rank,
+ "text_side_truth_frequency_rank": text_rank,
+ "rank_agreement": float(
+ np.mean([a == b for a, b in zip(truth_rank, text_rank)])
+ ),
+ "text_colour_words_recovered": sorted(
+ set(families["colour_words"][: args.colour_classes])
+ & set(truth_rank)
+ ),
+ }
+ print(json.dumps(report["text_families"], indent=1))
+ print(json.dumps(report["evaluation_only"], indent=1))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/tier0_pipeline.py b/worldalign/tier0_pipeline.py
new file mode 100644
index 0000000..778e20f
--- /dev/null
+++ b/worldalign/tier0_pipeline.py
@@ -0,0 +1,310 @@
+"""Tier 0 pipeline: unpaired corpora to aligned relation fields.
+
+Consolidates the components that produced the synthetic world's
+end-to-end result, each of which was developed and measured separately:
+
+- watershed object extraction, which splits touching ring members that
+ connected components merge (exact object count 74.2% to 100%);
+- Gestalt appearance grouping, which assembles members into groups by
+ shared colour, size, and shape rather than by a distance threshold
+ (exact group count 71.9% to 90.2%);
+- size classes from radial extent under one-dimensional k-means per
+ member count, because area confounds size with shape and the classes
+ are gap-separated rather than equally populated (52.7% to 87.9%);
+- text factor families from mutual exclusivity within a group phrase;
+- the cross-modal value correspondence from marginal frequency rank.
+
+Nothing crosses modalities except the frequency ranking, and the two
+corpora it reads are disjoint: no instance appears on both sides.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from collections import Counter, defaultdict
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from scipy import ndimage
+from scipy.cluster.hierarchy import fcluster, linkage
+from sklearn.cluster import KMeans
+from tqdm import tqdm
+
+from .common import read_json, seed_everything, write_json
+from .synth_cc_battery import moment_field
+from .synth_set_battery import parse_group_phrases
+from .synth_towers import load_image
+from .tier0_dictionary import partition_text_vocabulary, text_token_statistics
+
+NUMBER_TO_COUNT = {"two": 2, "three": 3, "four": 4}
+SINGULAR = ("a", "an")
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--data-dir", default="artifacts/synth_v1")
+ parser.add_argument("--split", choices=["val", "test"], default="test")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--offset", type=int, default=0)
+ parser.add_argument("--fit-scenes", type=int, default=1500)
+ parser.add_argument("--peak-distance", type=int, default=5)
+ parser.add_argument("--group-threshold", type=float, default=0.35)
+ parser.add_argument("--views", type=int, default=0, help="0 uses manifest.")
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", required=True)
+ parser.add_argument("--states-output", default="")
+ return parser.parse_args()
+
+
+def extract_objects(image: torch.Tensor, peak_distance: int) -> list[dict]:
+ """Foreground objects, splitting touching ones by distance watershed."""
+ array = image.permute(1, 2, 0).numpy()
+ background = np.median(array.reshape(-1, 3), axis=0)
+ foreground = np.abs(array - background).sum(-1) > 0.12
+ distance = ndimage.distance_transform_edt(foreground)
+ window = 2 * peak_distance + 1
+ peaks = (
+ distance >= ndimage.maximum_filter(distance, size=window) - 1e-9
+ ) & (distance > 1.0)
+ labels, count = ndimage.label(peaks)
+ if count == 0:
+ labels, count = ndimage.label(foreground)
+ else:
+ while True:
+ grown = ndimage.grey_dilation(labels, size=3)
+ take = (labels == 0) & foreground & (grown > 0)
+ if not take.any():
+ break
+ labels = np.where(take, grown, labels)
+ count = labels.max()
+ objects = []
+ for index in range(1, count + 1):
+ mask = labels == index
+ area = float(mask.sum())
+ if area < 6:
+ continue
+ ys, xs = np.nonzero(mask)
+ centred_y, centred_x = ys - ys.mean(), xs - xs.mean()
+ covariance = np.cov(np.stack([centred_x, centred_y])) + 1e-6 * np.eye(2)
+ eigenvalues = np.linalg.eigvalsh(covariance)
+ eroded = ndimage.binary_erosion(mask)
+ perimeter = max(float((mask & ~eroded).sum()), 1.0)
+ box = (xs.max() - xs.min() + 1) * (ys.max() - ys.min() + 1)
+ objects.append(
+ {
+ "rgb": array[mask].mean(0),
+ "area": area,
+ "extent": float(np.hypot(centred_y, centred_x).max()),
+ "shape": np.array(
+ [
+ 4 * np.pi * area / perimeter**2,
+ 1 - eigenvalues[0] / eigenvalues[1],
+ area / max(box, 1),
+ ]
+ ),
+ }
+ )
+ return objects
+
+
+def appearance(item: dict) -> np.ndarray:
+ return np.concatenate(
+ [item["rgb"] * 3.0, [np.log(item["area"] + 1e-6) * 0.6], item["shape"]]
+ )
+
+
+def group_objects(objects: list[dict], threshold: float) -> list[dict]:
+ """Members of one group share appearance; group by similarity."""
+ if not objects:
+ return []
+ if len(objects) == 1:
+ labels = np.array([0])
+ else:
+ features = np.stack([appearance(item) for item in objects])
+ labels = fcluster(linkage(features, "complete"), threshold, "distance")
+ buckets: dict[int, list[dict]] = {}
+ for item, label in zip(objects, labels):
+ buckets.setdefault(int(label), []).append(item)
+ return [
+ {
+ "rgb": np.mean([item["rgb"] for item in members], axis=0),
+ "members": len(members),
+ "extent": float(np.mean([item["extent"] for item in members])),
+ }
+ for members in buckets.values()
+ ]
+
+
+class VisionCoder:
+ """Colour classes and size classes fitted on the vision corpus alone."""
+
+ def __init__(self, groups: list[dict], classes: int, seed: int) -> None:
+ colours = np.stack([group["rgb"] for group in groups]).astype(np.float64)
+ self.colour_model = KMeans(classes, n_init=10, random_state=seed).fit(colours)
+ frequency = Counter(self.colour_model.labels_.tolist())
+ self.colour_rank = {
+ label: rank for rank, (label, _) in enumerate(frequency.most_common())
+ }
+ by_count: dict[int, list[float]] = defaultdict(list)
+ for group in groups:
+ by_count[min(group["members"], 4)].append(group["extent"])
+ self.size_models = {}
+ for count, extents in by_count.items():
+ model = KMeans(3, n_init=10, random_state=seed).fit(
+ np.asarray(extents, dtype=np.float64)[:, None]
+ )
+ order = np.argsort(model.cluster_centers_[:, 0])
+ self.size_models[count] = (
+ model,
+ {int(label): rank for rank, label in enumerate(order)},
+ )
+
+ def encode(self, groups: list[dict], classes: int) -> torch.Tensor:
+ if not groups:
+ groups = [{"rgb": np.zeros(3), "members": 1, "extent": 1.0}]
+ colours = self.colour_model.predict(
+ np.stack([group["rgb"] for group in groups]).astype(np.float64)
+ )
+ vectors = []
+ for group, colour in zip(groups, colours):
+ model, order = self.size_models[min(group["members"], 4)]
+ size = order[
+ int(model.predict(np.array([[group["extent"]]], dtype=np.float64))[0])
+ ]
+ vectors.append(
+ factor_vector(self.colour_rank[int(colour)], group["members"], size, classes)
+ )
+ return F.normalize(torch.stack(vectors), dim=-1)
+
+
+def factor_vector(colour: int, count: int, size: int, classes: int) -> torch.Tensor:
+ vector = torch.zeros(classes + 4 + 3)
+ vector[colour] = 1.0
+ vector[classes + min(count - 1, 3)] = 1.0
+ vector[classes + 4 + size] = 1.0
+ return vector
+
+
+def encode_caption(caption: str, colour_rank: dict[str, int], classes: int) -> torch.Tensor:
+ vectors = []
+ for phrase in parse_group_phrases(caption):
+ tokens = phrase.split()
+ colour = next((colour_rank[t] for t in tokens if t in colour_rank), 0)
+ count = (
+ 1
+ if any(token in SINGULAR for token in tokens)
+ else next(
+ (NUMBER_TO_COUNT[t] for t in tokens if t in NUMBER_TO_COUNT), 1
+ )
+ )
+ size = 0 if "small" in tokens else (2 if "large" in tokens else 1)
+ vectors.append(factor_vector(colour, count, size, classes))
+ if not vectors:
+ vectors = [factor_vector(0, 1, 1, classes)]
+ return F.normalize(torch.stack(vectors), dim=-1)
+
+
+def moment_state(states: torch.Tensor) -> torch.Tensor:
+ first = states.mean(0)
+ second = (states[:, :, None] * states[:, None, :]).mean(0).flatten()
+ return torch.cat([first, second])
+
+
+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"]
+ image_dir = Path(manifest["image_dir"])
+ views = args.views or manifest["visual_views"]
+ rows = manifest[args.split][args.offset : args.offset + args.samples]
+
+ families = partition_text_vocabulary(
+ text_token_statistics(captions, manifest["text_only_train"])
+ )
+ colour_words = families["colour_words"]
+ classes = len(colour_words)
+ colour_rank = {word: rank for rank, word in enumerate(colour_words)}
+
+ fit_groups = [
+ group
+ for row in tqdm(
+ manifest["vision_only_train"][: args.fit_scenes], desc="fit vision"
+ )
+ for group in group_objects(
+ extract_objects(
+ load_image(image_dir / f"scene{row:06d}_v0.png"), args.peak_distance
+ ),
+ args.group_threshold,
+ )
+ ]
+ coder = VisionCoder(fit_groups, classes, args.seed)
+
+ per_view_fields = []
+ view_states = []
+ for view in range(views):
+ sets = [
+ coder.encode(
+ group_objects(
+ extract_objects(
+ load_image(image_dir / f"scene{row:06d}_v{view}.png"),
+ args.peak_distance,
+ ),
+ args.group_threshold,
+ ),
+ classes,
+ )
+ for row in tqdm(rows, desc=f"encode view {view}")
+ ]
+ per_view_fields.append(moment_field(sets))
+ if view == 0:
+ view_states = [moment_state(item) for item in sets]
+ visual_field = torch.stack(per_view_fields).mean(0)
+
+ text_sets = [encode_caption(captions[row][0], colour_rank, classes) for row in rows]
+ text_field = moment_field(text_sets)
+
+ mask = ~np.eye(len(rows), dtype=bool)
+ correlation = float(
+ np.corrcoef(
+ visual_field.double().numpy()[mask], text_field.double().numpy()[mask]
+ )[0, 1]
+ )
+ torch.save(
+ {"visual_field": visual_field, "text_field": text_field, "rows": rows},
+ args.output,
+ )
+ if args.states_output:
+ torch.save(
+ {
+ "vision_states": torch.stack(view_states),
+ "text_states": torch.stack([moment_state(item) for item in text_sets]),
+ "rows": rows,
+ },
+ args.states_output,
+ )
+ summary = {
+ "data_dir": args.data_dir,
+ "split": args.split,
+ "samples": len(rows),
+ "colour_classes": classes,
+ "text_families": {
+ key: families[key] for key in ("count_words", "colour_words", "size_words")
+ },
+ "field_correlation_at_truth": correlation,
+ "note": (
+ "The dictionary is derived from disjoint corpora; the "
+ "correlation is a diagnostic computed with hidden pairs and "
+ "never used by the pipeline."
+ ),
+ }
+ print(json.dumps({"field_correlation_at_truth": correlation}))
+ write_json(args.output.replace(".pt", ".json"), summary)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/train_bridge.py b/worldalign/train_bridge.py
new file mode 100644
index 0000000..da29194
--- /dev/null
+++ b/worldalign/train_bridge.py
@@ -0,0 +1,255 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import numpy as np
+import torch
+from torch.optim import AdamW
+from tqdm import tqdm
+from scipy.optimize import linear_sum_assignment
+
+from .common import (
+ cosine_isometry_loss,
+ cosine_loss,
+ cosine_schedule,
+ parameter_count,
+ read_json,
+ retrieval_metrics,
+ seed_everything,
+ sliced_wasserstein,
+)
+from .gw import gw_pseudo_targets
+from .io import load_feature_pair, select_rows
+from .models import Bridge
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--mode", choices=["paired", "unpaired_swd", "unpaired_gw"])
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--output", required=True)
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--steps", type=int, default=4_000)
+ p.add_argument("--batch-size", type=int, default=256)
+ p.add_argument("--hidden-dim", type=int, default=1536)
+ p.add_argument("--linear", action="store_true")
+ p.add_argument("--lr", type=float, default=3e-4)
+ p.add_argument("--warmup", type=int, default=200)
+ p.add_argument("--clusters", type=int, default=128)
+ p.add_argument(
+ "--gw-cache",
+ help="Optional path for reusable GW prototype coupling and assignments.",
+ )
+ p.add_argument("--swd-weight", type=float, default=10.0)
+ p.add_argument("--isometry-weight", type=float, default=1.0)
+ p.add_argument("--gw-weight", type=float, default=1.0)
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument("--eval-every", type=int, default=200)
+ return p.parse_args()
+
+
+def evaluate(
+ bridge: Bridge,
+ x: torch.Tensor,
+ y: torch.Tensor,
+ device: str,
+) -> dict[str, float]:
+ bridge.eval()
+ mapped = []
+ with torch.inference_mode():
+ for chunk in x.split(512):
+ mapped.append(bridge(chunk.to(device)).cpu())
+ bridge.train()
+ return retrieval_metrics(torch.cat(mapped), y)
+
+
+def evaluate_gw_cluster_mapping(
+ gw: dict,
+ val_x: torch.Tensor,
+ val_y: torch.Tensor,
+) -> dict[str, float]:
+ vx = torch.nn.functional.normalize(gw["vision_centers"], dim=-1)
+ ty = torch.nn.functional.normalize(gw["text_centers"], dim=-1)
+ x_labels = (torch.nn.functional.normalize(val_x, dim=-1) @ vx.T).argmax(1)
+ y_labels = (torch.nn.functional.normalize(val_y, dim=-1) @ ty.T).argmax(1)
+ predicted_map = gw["coupling"].argmax(1)
+ predicted = predicted_map[x_labels]
+ accuracy = (predicted == y_labels).float().mean().item()
+
+ k = vx.shape[0]
+ contingency = torch.zeros(k, k, dtype=torch.float64)
+ for i, j in zip(x_labels.tolist(), y_labels.tolist()):
+ contingency[i, j] += 1
+ row, col = linear_sum_assignment(-contingency.numpy())
+ oracle_correct = contingency[row, col].sum().item()
+ return {
+ "paired_val_cluster_accuracy": float(accuracy),
+ "paired_val_cluster_chance": 1.0 / k,
+ "paired_val_cluster_oracle_permutation": float(
+ oracle_correct / max(len(val_x), 1)
+ ),
+ }
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
+
+ if args.mode == "paired":
+ rows = manifest["paired_train"]
+ train_x = select_rows(vision["features"], vlookup, rows)
+ train_y = select_rows(text["features"], tlookup, rows)
+ target_by_cluster = None
+ assignments = None
+ gw_meta = {}
+ gw_result = None
+ else:
+ train_x = select_rows(
+ vision["features"], vlookup, manifest["vision_only_train"]
+ )
+ train_y = select_rows(
+ text["features"], tlookup, manifest["text_only_train"]
+ )
+ target_by_cluster = None
+ assignments = None
+ gw_meta = {}
+ gw_result = None
+ if args.mode == "unpaired_gw":
+ if args.gw_cache and Path(args.gw_cache).exists():
+ print(f"Loading GW coupling from {args.gw_cache}")
+ gw = torch.load(
+ args.gw_cache, map_location="cpu", weights_only=False
+ )
+ else:
+ print("Computing unpaired GW prototype coupling...")
+ gw = gw_pseudo_targets(
+ train_x, train_y, args.clusters, args.seed
+ )
+ if args.gw_cache:
+ Path(args.gw_cache).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(gw, args.gw_cache)
+ print(f"Wrote GW coupling to {args.gw_cache}")
+ gw_result = gw
+ target_by_cluster = gw["target_centers"]
+ assignments = gw["vision_assignments"]
+ gw_meta = {
+ "gw_distance": gw["gw_distance"],
+ "coupling_row_entropy": gw["coupling_row_entropy"],
+ "clusters": gw["clusters"],
+ }
+ print(f"GW diagnostics: {gw_meta}")
+
+ val_rows = manifest["val"]
+ val_x = select_rows(vision["features"], vlookup, val_rows)
+ val_y = select_rows(text["features"], tlookup, val_rows)
+ if gw_result is not None:
+ gw_meta.update(evaluate_gw_cluster_mapping(gw_result, val_x, val_y))
+ print(f"GW paired-eval diagnostics (not used for training): {gw_meta}")
+
+ bridge = Bridge(
+ train_x.shape[-1],
+ train_y.shape[-1],
+ hidden_dim=args.hidden_dim,
+ linear=args.linear,
+ ).to(args.device)
+ print(f"Bridge parameters: {parameter_count(bridge):,}")
+ optimizer = AdamW(bridge.parameters(), lr=args.lr, weight_decay=1e-4)
+ generator = torch.Generator().manual_seed(args.seed)
+
+ history = []
+ best_score = -1.0
+ best_state = None
+ progress = tqdm(range(args.steps), desc=f"bridge:{args.mode}")
+ for step in progress:
+ ix = torch.randint(
+ len(train_x), (args.batch_size,), generator=generator
+ )
+ if args.mode == "paired":
+ iy = ix
+ else:
+ iy = torch.randint(
+ len(train_y), (args.batch_size,), generator=generator
+ )
+ x = train_x[ix].to(args.device)
+ y = train_y[iy].to(args.device)
+ mapped = bridge(x)
+
+ if args.mode == "paired":
+ alignment = cosine_loss(mapped, y)
+ swd = mapped.new_zeros(())
+ gw_loss = mapped.new_zeros(())
+ else:
+ alignment = mapped.new_zeros(())
+ swd = sliced_wasserstein(mapped, y, num_projections=64)
+ if target_by_cluster is not None and assignments is not None:
+ pseudo = target_by_cluster[assignments[ix]].to(args.device)
+ gw_loss = cosine_loss(mapped, pseudo)
+ else:
+ gw_loss = mapped.new_zeros(())
+ isometry = cosine_isometry_loss(x, mapped)
+ loss = (
+ alignment
+ + args.swd_weight * swd
+ + args.gw_weight * gw_loss
+ + args.isometry_weight * isometry
+ )
+
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ torch.nn.utils.clip_grad_norm_(bridge.parameters(), 1.0)
+ optimizer.step()
+ scale = cosine_schedule(step, args.steps, args.warmup)
+ for group in optimizer.param_groups:
+ group["lr"] = args.lr * scale
+
+ if step % 20 == 0:
+ progress.set_postfix(
+ loss=f"{loss.item():.3f}",
+ align=f"{alignment.item():.3f}",
+ swd=f"{swd.item():.3g}",
+ gw=f"{gw_loss.item():.3f}",
+ iso=f"{isometry.item():.3f}",
+ )
+ if step % args.eval_every == 0 or step == args.steps - 1:
+ metrics = evaluate(bridge, val_x, val_y, args.device)
+ record = {"step": step, "loss": float(loss.item()), **metrics}
+ history.append(record)
+ # Paired validation is recorded for scientific evaluation, not used to
+ # select the unsupervised checkpoint. Save final for unpaired modes.
+ if args.mode == "paired" and metrics["i2t_r@1"] > best_score:
+ best_score = metrics["i2t_r@1"]
+ best_state = {
+ k: v.detach().cpu().clone()
+ for k, v in bridge.state_dict().items()
+ }
+ print(record)
+
+ if args.mode == "paired" and best_state is not None:
+ bridge.load_state_dict(best_state)
+ checkpoint_selection = "best paired validation R@1 (upper bound only)"
+ else:
+ checkpoint_selection = "final step; no paired validation selection"
+
+ state = {
+ "config": bridge.config(),
+ "state_dict": bridge.state_dict(),
+ "mode": args.mode,
+ "args": vars(args),
+ "vision_model": vision["model"],
+ "text_model": text["model"],
+ "history": history,
+ "gw": gw_meta,
+ "checkpoint_selection": checkpoint_selection,
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/train_prefix.py b/worldalign/train_prefix.py
new file mode 100644
index 0000000..76607e6
--- /dev/null
+++ b/worldalign/train_prefix.py
@@ -0,0 +1,141 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import torch
+from torch.optim import AdamW
+from tqdm import tqdm
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import (
+ cosine_schedule,
+ dtype_for_device,
+ parameter_count,
+ read_json,
+ seed_everything,
+)
+from .io import load_feature_pair, select_rows
+from .models import PrefixAdapter
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--output", default="artifacts/prefix.pt")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--steps", type=int, default=3_000)
+ p.add_argument("--batch-size", type=int, default=32)
+ p.add_argument("--prefix-length", type=int, default=8)
+ p.add_argument("--hidden-dim", type=int, default=2048)
+ p.add_argument("--max-length", type=int, default=48)
+ p.add_argument("--lr", type=float, default=3e-4)
+ p.add_argument("--warmup", type=int, default=200)
+ p.add_argument("--seed", type=int, default=20260728)
+ return p.parse_args()
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ _, text, _, tlookup = load_feature_pair(args.vision, args.text)
+ train_rows = manifest["text_only_train"]
+ semantic = select_rows(text["features"], tlookup, train_rows)
+ captions_by_row = {
+ int(row): caption for row, caption in zip(text["rows"], text["captions"])
+ }
+ captions = [captions_by_row[int(row)] for row in train_rows]
+
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ tokenizer.padding_side = "right"
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+ for parameter in lm.parameters():
+ parameter.requires_grad_(False)
+ lm_dim = lm.get_input_embeddings().embedding_dim
+
+ adapter = PrefixAdapter(
+ semantic_dim=semantic.shape[-1],
+ lm_dim=lm_dim,
+ prefix_length=args.prefix_length,
+ hidden_dim=args.hidden_dim,
+ ).to(args.device)
+ print(f"Prefix adapter parameters: {parameter_count(adapter):,}")
+ optimizer = AdamW(adapter.parameters(), lr=args.lr, weight_decay=1e-4)
+ generator = torch.Generator().manual_seed(args.seed)
+
+ history = []
+ progress = tqdm(range(args.steps), desc="text-only prefix")
+ for step in progress:
+ ids = torch.randint(
+ len(semantic), (args.batch_size,), generator=generator
+ )
+ batch_captions = [captions[int(i)] for i in ids]
+ tokens = tokenizer(
+ batch_captions,
+ padding=True,
+ truncation=True,
+ max_length=args.max_length,
+ return_tensors="pt",
+ )
+ input_ids = tokens["input_ids"].to(args.device)
+ attention = tokens["attention_mask"].to(args.device)
+ prefix = adapter(semantic[ids].to(args.device)).to(dtype)
+ token_embeddings = lm.get_input_embeddings()(input_ids)
+ inputs_embeds = torch.cat([prefix, token_embeddings], dim=1)
+ prefix_attention = torch.ones(
+ prefix.shape[:2], dtype=attention.dtype, device=args.device
+ )
+ full_attention = torch.cat([prefix_attention, attention], dim=1)
+ labels = input_ids.clone()
+ labels[attention == 0] = -100
+ prefix_labels = torch.full(
+ prefix.shape[:2], -100, dtype=labels.dtype, device=args.device
+ )
+ full_labels = torch.cat([prefix_labels, labels], dim=1)
+ result = lm(
+ inputs_embeds=inputs_embeds,
+ attention_mask=full_attention,
+ labels=full_labels,
+ use_cache=False,
+ return_dict=True,
+ )
+ loss = result.loss
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ torch.nn.utils.clip_grad_norm_(adapter.parameters(), 1.0)
+ optimizer.step()
+ scale = cosine_schedule(step, args.steps, args.warmup)
+ for group in optimizer.param_groups:
+ group["lr"] = args.lr * scale
+
+ if step % 20 == 0:
+ progress.set_postfix(loss=f"{loss.item():.3f}")
+ if step % 100 == 0 or step == args.steps - 1:
+ history.append({"step": step, "loss": float(loss.item())})
+
+ state = {
+ "config": adapter.config(),
+ "state_dict": adapter.state_dict(),
+ "text_model": text["model"],
+ "text_layer": text["layer"],
+ "args": vars(args),
+ "history": history,
+ "training": "text only; no image features or image-text pairs",
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
+
diff --git a/worldalign/vg_attention_extract.py b/worldalign/vg_attention_extract.py
new file mode 100644
index 0000000..0ceb4ce
--- /dev/null
+++ b/worldalign/vg_attention_extract.py
@@ -0,0 +1,269 @@
+"""R3 battery extraction: model-layer relational readout for VG nodes.
+
+The hypothesis under test is that relations live in the model's
+computation rather than in output embedding geometry. For each node this
+extracts, per modality, a stack of view-pair relation channels read from
+inside the frozen models:
+
+- text: the sixteen region phrases are encoded jointly in one context;
+ cross-phrase attention mass, pooled over layer groups and averaged over
+ several phrase orders (position and causality artifacts cancel), plus
+ in-context phrase states from the final layer.
+- vision: the full image is encoded once; patch-patch attention pooled
+ over region boxes per layer group, plus in-context region states pooled
+ from patch tokens.
+
+Block means are computed as P A P^T with span-indicator matrices P, and
+attention layers are reduced into group accumulators one layer at a time,
+so neither the full layer stack nor per-pair loops materialize.
+
+No pairs, node identities, or view correspondences are used. Outputs are
+keyed by released node IDs in released view order.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from concurrent.futures import ThreadPoolExecutor
+from pathlib import Path
+
+import numpy as np
+import torch
+from PIL import Image
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel, AutoTokenizer
+
+from .common import batch_indices
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--side", choices=["text", "vision"], required=True)
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images")
+ parser.add_argument("--text-model", default="Qwen/Qwen2.5-0.5B")
+ parser.add_argument("--vision-model", default="facebook/dinov2-small")
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--orders", type=int, default=4, help="Text phrase orders.")
+ parser.add_argument("--image-size", type=int, default=224)
+ parser.add_argument(
+ "--batch-size", type=int, default=16, help="Sequences or images per forward."
+ )
+ parser.add_argument("--limit", type=int)
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument("--output", required=True)
+ return parser.parse_args()
+
+
+def read_jsonl(path: Path) -> list[dict]:
+ return [
+ json.loads(line)
+ for line in path.read_text(encoding="utf-8").splitlines()
+ if line.strip()
+ ]
+
+
+def layer_groups(count: int, groups: int) -> list[list[int]]:
+ bounds = torch.linspace(0, count, groups + 1).long().tolist()
+ return [list(range(a, b)) for a, b in zip(bounds[:-1], bounds[1:])]
+
+
+def grouped_attention(
+ attentions: tuple[torch.Tensor, ...], groups: list[list[int]]
+) -> torch.Tensor:
+ """Mean over heads and over each layer group, one layer at a time."""
+ batch, _, length, _ = attentions[0].shape
+ result = torch.zeros(
+ len(groups), batch, length, length, device=attentions[0].device
+ )
+ for g, layer_list in enumerate(groups):
+ for layer in layer_list:
+ result[g] += attentions[layer].float().mean(1)
+ result[g] /= len(layer_list)
+ return result
+
+
+@torch.inference_mode()
+def extract_text(args: argparse.Namespace) -> None:
+ records = read_jsonl(Path(args.vg_dir, "text_nodes.jsonl"))
+ if args.limit:
+ records = records[: args.limit]
+ tokenizer = AutoTokenizer.from_pretrained(args.text_model)
+ model = AutoModel.from_pretrained(
+ args.text_model, torch_dtype=torch.bfloat16, attn_implementation="eager"
+ ).to(args.device)
+ model.eval()
+ groups = layer_groups(model.config.num_hidden_layers, 4)
+ separator = tokenizer("\n", add_special_tokens=False)["input_ids"]
+ generator = torch.Generator().manual_seed(args.seed)
+
+ jobs: list[tuple[int, list[int], list[tuple[int, int]]]] = []
+ for index, record in enumerate(records):
+ phrases = record["region_closed"]
+ for _ in range(args.orders):
+ order = torch.randperm(len(phrases), generator=generator)
+ ids: list[int] = []
+ span: list[tuple[int, int]] = [(0, 0)] * len(phrases)
+ for position in order.tolist():
+ tokens = tokenizer(
+ " " + phrases[position].strip(), add_special_tokens=False
+ )["input_ids"]
+ span[position] = (len(ids), len(ids) + len(tokens))
+ ids.extend(tokens + separator)
+ jobs.append((index, ids, span))
+
+ views = len(records[0]["region_closed"])
+ attention_sum = torch.zeros(len(records), len(groups), views, views)
+ state_sum = torch.zeros(len(records), views, model.config.hidden_size)
+
+ for start in tqdm(range(0, len(jobs), args.batch_size), desc="text attention"):
+ batch = jobs[start : start + args.batch_size]
+ longest = max(len(ids) for _, ids, _ in batch)
+ pad_id = tokenizer.pad_token_id or tokenizer.eos_token_id
+ input_ids = torch.full((len(batch), longest), pad_id, dtype=torch.long)
+ attention_mask = torch.zeros_like(input_ids)
+ indicator = torch.zeros(len(batch), views, longest)
+ for row, (_, ids, span) in enumerate(batch):
+ input_ids[row, : len(ids)] = torch.tensor(ids)
+ attention_mask[row, : len(ids)] = 1
+ for view, (a0, a1) in enumerate(span):
+ indicator[row, view, a0:a1] = 1.0 / max(a1 - a0, 1)
+ result = model(
+ input_ids=input_ids.to(args.device),
+ attention_mask=attention_mask.to(args.device),
+ output_attentions=True,
+ return_dict=True,
+ )
+ grouped = grouped_attention(result.attentions, groups) # [G, B, S, S]
+ indicator_device = indicator.to(args.device)
+ pooled = torch.einsum(
+ "bvs,gbst,bwt->gbvw", indicator_device, grouped, indicator_device
+ ).cpu()
+ hidden = result.last_hidden_state.float()
+ states = torch.bmm(indicator_device, hidden).cpu()
+ for row, (index, _, _) in enumerate(batch):
+ attention_sum[index] += pooled[:, row]
+ state_sum[index] += states[row]
+
+ attention_channels = attention_sum / args.orders
+ attention_channels = 0.5 * (
+ attention_channels + attention_channels.transpose(-2, -1)
+ )
+ torch.save(
+ {
+ "side": "text",
+ "model": args.text_model,
+ "node_ids": [record["node_id"] for record in records],
+ "attention_channels": attention_channels,
+ "context_states": state_sum / args.orders,
+ "orders": args.orders,
+ "layer_groups": [len(g) for g in groups],
+ },
+ args.output,
+ )
+ print(f"Wrote {args.output}")
+
+
+@torch.inference_mode()
+def extract_vision(args: argparse.Namespace) -> None:
+ records = read_jsonl(Path(args.vg_dir, "vision_nodes.private.jsonl"))
+ if args.limit:
+ records = records[: args.limit]
+ processor = AutoImageProcessor.from_pretrained(args.vision_model)
+ model = AutoModel.from_pretrained(
+ args.vision_model, torch_dtype=torch.float32, attn_implementation="eager"
+ ).to(args.device)
+ model.eval()
+ groups = layer_groups(model.config.num_hidden_layers, 3)
+ patch = model.config.patch_size
+ grid = args.image_size // patch
+ tokens = grid * grid
+ mean = torch.tensor(processor.image_mean).view(3, 1, 1)
+ std = torch.tensor(processor.image_std).view(3, 1, 1)
+
+ def load_one(record: dict) -> torch.Tensor:
+ path = Path(args.image_cache, f"{record['source_image_id']}.jpg")
+ with Image.open(path) as image:
+ resized = image.convert("RGB").resize(
+ (args.image_size, args.image_size), Image.BILINEAR
+ )
+ pixels = torch.from_numpy(np.asarray(resized).copy()).permute(2, 0, 1)
+ return (pixels.float() / 255.0 - mean) / std
+
+ def region_indicator(record: dict) -> torch.Tensor:
+ width, height = record["width"], record["height"]
+ rows = []
+ centers = torch.arange(grid) + 0.5
+ for region in record["regions"]:
+ x0 = region["x"] / width * grid
+ x1 = (region["x"] + region["width"]) / width * grid
+ y0 = region["y"] / height * grid
+ y1 = (region["y"] + region["height"]) / height * grid
+ in_x = (centers >= x0) & (centers <= x1)
+ in_y = (centers >= y0) & (centers <= y1)
+ mask = (in_y[:, None] & in_x[None, :]).flatten().double()
+ if mask.sum() == 0:
+ cx = min(grid - 1, max(0, int((x0 + x1) / 2)))
+ cy = min(grid - 1, max(0, int((y0 + y1) / 2)))
+ mask[cy * grid + cx] = 1.0
+ rows.append(mask / mask.sum())
+ return torch.stack(rows).float()
+
+ views = len(records[0]["regions"])
+ node_ids: list[str] = []
+ attention_channels: list[torch.Tensor] = []
+ context_states: list[torch.Tensor] = []
+ with ThreadPoolExecutor(max_workers=8) as pool:
+ for indices in tqdm(
+ list(batch_indices(len(records), args.batch_size)),
+ desc="vision attention",
+ ):
+ batch = [records[i] for i in indices]
+ pixels = torch.stack(list(pool.map(load_one, batch)))
+ result = model(
+ pixel_values=pixels.to(args.device),
+ output_attentions=True,
+ return_dict=True,
+ )
+ grouped = grouped_attention(result.attentions, groups)
+ special = grouped.shape[-1] - tokens # CLS and any registers
+ grouped = grouped[:, :, special:, special:]
+ indicator = torch.stack(
+ [region_indicator(record) for record in batch]
+ ).to(args.device)
+ pooled = torch.einsum(
+ "bvs,gbst,bwt->gbvw", indicator, grouped, indicator
+ ).cpu()
+ pooled = 0.5 * (pooled + pooled.transpose(-2, -1))
+ patches = result.last_hidden_state[:, special:].float()
+ states = torch.bmm(indicator, patches).cpu()
+ for row, record in enumerate(batch):
+ node_ids.append(record["node_id"])
+ attention_channels.append(pooled[:, row])
+ context_states.append(states[row])
+ torch.save(
+ {
+ "side": "vision",
+ "model": args.vision_model,
+ "node_ids": node_ids,
+ "attention_channels": torch.stack(attention_channels),
+ "context_states": torch.stack(context_states),
+ "image_size": args.image_size,
+ "layer_groups": [len(g) for g in groups],
+ },
+ args.output,
+ )
+ print(f"Wrote {args.output}")
+
+
+def main() -> None:
+ args = parse_args()
+ if args.side == "text":
+ extract_text(args)
+ else:
+ extract_vision(args)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_attention_probe.py b/worldalign/vg_attention_probe.py
new file mode 100644
index 0000000..3dac7a3
--- /dev/null
+++ b/worldalign/vg_attention_probe.py
@@ -0,0 +1,170 @@
+"""R3 battery evaluation: do model-internal relations align across modalities?
+
+Compares, at the replayed true view correspondence, the cross-modal
+alignment of relation fields read from inside the frozen models
+(cross-phrase attention per layer group, in-context state cosines) against
+the isolated-encoding cosine baseline (matched rho 0.128, z 46.6).
+Spearman is rank-based, so monotone attention transforms are immaterial.
+View truth is evaluation-only.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+
+import torch
+import torch.nn.functional as F
+
+from .common import write_json
+from .vg_view_probe import offdiag, replay_view_permutations, spearman
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf")
+ parser.add_argument("--seed", type=int, default=20260728)
+ parser.add_argument("--views", type=int, default=16)
+ parser.add_argument("--text-attn", default="artifacts/manifold_gate/attn_text_full.pt")
+ parser.add_argument(
+ "--vision-attn", default="artifacts/manifold_gate/attn_vision_full.pt"
+ )
+ parser.add_argument("--nodes", type=int, default=0)
+ parser.add_argument(
+ "--output", default="artifacts/manifold_gate/attention_probe.json"
+ )
+ return parser.parse_args()
+
+
+def channel_fields(state: dict, side: str) -> dict[str, torch.Tensor]:
+ """Named [N, V, V] relation fields for one side."""
+ attention = state["attention_channels"].double()
+ fields = {
+ f"{side}_attn_g{group}": attention[:, group]
+ for group in range(attention.shape[1])
+ }
+ context = F.normalize(state["context_states"].double(), dim=-1)
+ fields[f"{side}_ctx_cos"] = context @ context.transpose(-2, -1)
+ return fields
+
+
+def main() -> None:
+ args = parse_args()
+ mappings = replay_view_permutations(args)
+
+ text_state = torch.load(args.text_attn, map_location="cpu", weights_only=False)
+ vision_state = torch.load(args.vision_attn, map_location="cpu", weights_only=False)
+ text_index = {node: i for i, node in enumerate(text_state["node_ids"])}
+ vision_index = {node: i for i, node in enumerate(vision_state["node_ids"])}
+
+ baseline_vision = torch.load(
+ f"{args.vg_dir}/vision_features.pt", map_location="cpu", weights_only=False
+ )
+ baseline_text = torch.load(
+ f"{args.vg_dir}/text_features.pt", map_location="cpu", weights_only=False
+ )
+ baseline_vision_index = {
+ node: i for i, node in enumerate(baseline_vision["node_ids"])
+ }
+ baseline_text_index = {node: i for i, node in enumerate(baseline_text["node_ids"])}
+
+ text_fields = channel_fields(text_state, "text")
+ vision_fields = channel_fields(vision_state, "vision")
+
+ node_ids = sorted(mappings)
+ if args.nodes:
+ node_ids = node_ids[: args.nodes]
+
+ pairs = [
+ (t_name, v_name) for t_name in text_fields for v_name in vision_fields
+ ]
+ matched: dict[tuple[str, str], list[float]] = {pair: [] for pair in pairs}
+ shuffled: dict[tuple[str, str], list[float]] = {pair: [] for pair in pairs}
+ baseline_matched: list[float] = []
+ baseline_shuffled: list[float] = []
+ generator = torch.Generator().manual_seed(args.seed)
+
+ for node in node_ids:
+ mapping = mappings[node]
+ text_row = text_index[mapping["text_node_id"]]
+ vision_row = vision_index[node]
+ text_to_vision = torch.tensor(mapping["text_to_vision"])
+ vision_to_text = torch.empty_like(text_to_vision)
+ vision_to_text[text_to_vision] = torch.arange(len(text_to_vision))
+ shuffle = torch.randperm(args.views, generator=generator)
+
+ for t_name, v_name in pairs:
+ text_field = text_fields[t_name][text_row]
+ aligned = text_field[vision_to_text][:, vision_to_text]
+ visual_field = vision_fields[v_name][vision_row]
+ matched[(t_name, v_name)].append(
+ spearman(offdiag(visual_field), offdiag(aligned))
+ )
+ scrambled = text_field[shuffle][:, shuffle]
+ shuffled[(t_name, v_name)].append(
+ spearman(offdiag(visual_field), offdiag(scrambled))
+ )
+
+ crops = F.normalize(
+ baseline_vision["region_features"][
+ baseline_vision_index[node]
+ ].double(),
+ dim=-1,
+ )
+ phrases = F.normalize(
+ baseline_text["region_features"][
+ baseline_text_index[mapping["text_node_id"]]
+ ].double(),
+ dim=-1,
+ )
+ crop_field = crops @ crops.T
+ phrase_field = (phrases @ phrases.T)[vision_to_text][:, vision_to_text]
+ baseline_matched.append(spearman(offdiag(crop_field), offdiag(phrase_field)))
+ scrambled = (phrases @ phrases.T)[shuffle][:, shuffle]
+ baseline_shuffled.append(spearman(offdiag(crop_field), offdiag(scrambled)))
+
+ def summarize(matched_values: list[float], shuffled_values: list[float]) -> dict:
+ m = torch.tensor(matched_values)
+ s = torch.tensor(shuffled_values)
+ return {
+ "matched_mean": float(m.mean()),
+ "shuffled_mean": float(s.mean()),
+ "gap_z": float(
+ (m.mean() - s.mean())
+ / (m - s).std().clamp_min(1e-12)
+ * len(m) ** 0.5
+ ),
+ }
+
+ report = {
+ "protocol": (
+ "Replayed view truth, evaluation-only. Each cell is the "
+ "within-node cross-modal Spearman of relation fields at the "
+ "true view correspondence, against a shuffled-view control."
+ ),
+ "nodes": len(node_ids),
+ "baseline_isolated_cosine": summarize(baseline_matched, baseline_shuffled),
+ "channels": {
+ f"{t} x {v}": summarize(matched[(t, v)], shuffled[(t, v)])
+ for t, v in pairs
+ },
+ }
+ write_json(args.output, report)
+ best = max(
+ report["channels"].items(), key=lambda item: item[1]["matched_mean"]
+ )
+ print(
+ json.dumps(
+ {
+ "baseline": report["baseline_isolated_cosine"],
+ "best_channel": {best[0]: best[1]},
+ "nodes": len(node_ids),
+ }
+ )
+ )
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_diagnose.py b/worldalign/vg_diagnose.py
new file mode 100644
index 0000000..4c6d4a8
--- /dev/null
+++ b/worldalign/vg_diagnose.py
@@ -0,0 +1,200 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+from scipy.stats import spearmanr
+import torch
+
+from .common import normalized, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--vision", default="artifacts/vg/vision_features.pt")
+ p.add_argument("--text", default="artifacts/vg/text_features.pt")
+ p.add_argument(
+ "--truth", default="artifacts/vg/ground_truth.private.jsonl"
+ )
+ p.add_argument("--max-nodes", type=int, default=5_000)
+ p.add_argument("--quantiles", type=int, default=33)
+ p.add_argument("--output", default="artifacts/vg/diagnostics.json")
+ return p.parse_args()
+
+
+def read_jsonl(path: str) -> list[dict]:
+ with open(path, encoding="utf-8") as handle:
+ return [json.loads(line) for line in handle if line.strip()]
+
+
+def bundle_signature(
+ features: torch.Tensor, quantiles: int = 33
+) -> torch.Tensor:
+ """Permutation- and rotation-invariant signature of a bag of views."""
+ x = normalized(features.float())
+ gram = x @ x.transpose(1, 2)
+ views = x.shape[1]
+ i, j = torch.triu_indices(views, views, offset=1)
+ pairwise = gram[:, i, j]
+ q = torch.linspace(0, 1, quantiles)
+ distribution = torch.quantile(pairwise, q, dim=1).T
+ eigenvalues = torch.linalg.eigvalsh(gram).flip(-1) / max(views, 1)
+ return torch.cat([distribution, eigenvalues], dim=-1)
+
+
+def rank_normalize_columns(x: torch.Tensor) -> torch.Tensor:
+ order = torch.argsort(x, dim=0)
+ ranks = torch.argsort(order, dim=0).float()
+ ranks = ranks / max(x.shape[0] - 1, 1)
+ return (ranks - 0.5) * 2
+
+
+def aligned_text_indices(
+ vision_ids: list[str], text_ids: list[str], truth_path: str
+) -> torch.Tensor:
+ truth = read_jsonl(truth_path)
+ mapping = {
+ item["vision_node_id"]: item["text_node_id"] for item in truth
+ }
+ text_lookup = {node_id: idx for idx, node_id in enumerate(text_ids)}
+ return torch.tensor([text_lookup[mapping[node_id]] for node_id in vision_ids])
+
+
+def arbitrary_target_retrieval(
+ queries: torch.Tensor, candidates: torch.Tensor, targets: torch.Tensor
+) -> dict:
+ similarity = normalized(queries) @ normalized(candidates).T
+ target_score = similarity[
+ torch.arange(len(queries)), targets.to(similarity.device)
+ ]
+ ranks = (similarity > target_score[:, None]).sum(-1) + 1
+ top_values, top_indices = similarity.topk(
+ min(2, similarity.shape[1]), dim=1
+ )
+ prediction = top_indices[:, 0]
+ correct = prediction == targets.to(prediction.device)
+ if top_values.shape[1] == 2:
+ margin = top_values[:, 0] - top_values[:, 1]
+ else:
+ margin = top_values[:, 0]
+ confidence_order = torch.argsort(margin, descending=True)
+ confidence_precision = {}
+ for count in (10, 50, 100, 500, 1_000):
+ if count <= len(queries):
+ selected = confidence_order[:count]
+ confidence_precision[str(count)] = {
+ "correct": int(correct[selected].sum()),
+ "precision": float(correct[selected].float().mean()),
+ }
+
+ reverse_prediction = similarity.argmax(dim=0)
+ mutual = (
+ reverse_prediction[prediction]
+ == torch.arange(len(queries), device=prediction.device)
+ )
+ mutual_count = int(mutual.sum())
+ result = {
+ "r@1": float((ranks <= 1).float().mean()),
+ "r@5": float((ranks <= 5).float().mean()),
+ "r@10": float((ranks <= 10).float().mean()),
+ "mean_reciprocal_rank": float((1.0 / ranks.float()).mean()),
+ "median_rank": float(ranks.float().median()),
+ "chance_r@1": 1.0 / len(candidates),
+ "chance_r@5": min(5.0 / len(candidates), 1.0),
+ "chance_r@10": min(10.0 / len(candidates), 1.0),
+ "confidence_margin_precision": confidence_precision,
+ "mutual_nearest": {
+ "selected": mutual_count,
+ "correct": int(correct[mutual].sum()),
+ "precision": (
+ float(correct[mutual].float().mean())
+ if mutual_count
+ else 0.0
+ ),
+ },
+ }
+ return result
+
+
+def upper_triangle(x: torch.Tensor) -> np.ndarray:
+ i, j = torch.triu_indices(len(x), len(x), offset=1)
+ return x[i, j].cpu().numpy()
+
+
+def main() -> None:
+ args = parse_args()
+ vision = torch.load(args.vision, map_location="cpu", weights_only=False)
+ text = torch.load(args.text, map_location="cpu", weights_only=False)
+ n = min(args.max_nodes, len(vision["node_ids"]), len(text["node_ids"]))
+ vision_ids = vision["node_ids"][:n]
+ target_full = aligned_text_indices(
+ vision_ids, text["node_ids"], args.truth
+ )
+ candidate_indices = torch.unique(target_full, sorted=False)
+ if len(candidate_indices) != n:
+ raise ValueError("Ground truth is not a one-to-one permutation")
+ candidate_lookup = {
+ int(old): new for new, old in enumerate(candidate_indices.tolist())
+ }
+ targets = torch.tensor(
+ [candidate_lookup[int(old)] for old in target_full.tolist()]
+ )
+
+ v_views = vision["region_features"][:n]
+ t_views = text["region_features"][candidate_indices]
+ v_signature = rank_normalize_columns(
+ bundle_signature(v_views, args.quantiles)
+ )
+ t_signature = rank_normalize_columns(
+ bundle_signature(t_views, args.quantiles)
+ )
+ signature_retrieval = arbitrary_target_retrieval(
+ v_signature, t_signature, targets
+ )
+
+ paired_t_signature = t_signature[targets]
+ signature_cosine = (
+ normalized(v_signature) * normalized(paired_t_signature)
+ ).sum(-1)
+ generator = torch.Generator().manual_seed(20260728)
+ shuffled = paired_t_signature[torch.randperm(n, generator=generator)]
+ shuffled_cosine = (
+ normalized(v_signature) * normalized(shuffled)
+ ).sum(-1)
+
+ v_scene = vision["global_features"][:n]
+ t_scene = text["region_features"][candidate_indices].mean(1)[targets]
+ gv = normalized(v_scene) @ normalized(v_scene).T
+ gt = normalized(t_scene) @ normalized(t_scene).T
+ scene_rho = spearmanr(
+ upper_triangle(gv), upper_triangle(gt)
+ ).statistic
+
+ result = {
+ "nodes": n,
+ "views_per_node": int(v_views.shape[1]),
+ "vision_model": vision["model"],
+ "text_model": text["model"],
+ "text_tier": text["tier"],
+ "bundle_signature_retrieval": signature_retrieval,
+ "paired_bundle_signature_cosine_mean": float(
+ signature_cosine.mean()
+ ),
+ "shuffled_bundle_signature_cosine_mean": float(
+ shuffled_cosine.mean()
+ ),
+ "between_scene_pairwise_cosine_spearman": float(scene_rho),
+ "evaluation_note": (
+ "Ground-truth permutation is used only to score signatures and "
+ "relation geometry, never to fit them."
+ ),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_extract_text.py b/worldalign/vg_extract_text.py
new file mode 100644
index 0000000..affff63
--- /dev/null
+++ b/worldalign/vg_extract_text.py
@@ -0,0 +1,111 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+from tqdm import tqdm
+from transformers import AutoModel, AutoTokenizer
+
+from .common import dtype_for_device
+from .extract_text import mean_pool
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--nodes", default="artifacts/vg/text_nodes.jsonl")
+ p.add_argument("--output", default="artifacts/vg/text_features.pt")
+ p.add_argument(
+ "--tier",
+ choices=["region_closed", "visible_relations", "qa_expanded"],
+ default="region_closed",
+ )
+ p.add_argument("--model", default="Qwen/Qwen2.5-1.5B")
+ p.add_argument("--device", default="cuda:3")
+ p.add_argument("--batch-size", type=int, default=96)
+ p.add_argument("--max-length", type=int, default=48)
+ p.add_argument("--layer", type=int, default=-1)
+ p.add_argument("--views", type=int, default=32)
+ p.add_argument("--limit", type=int)
+ return p.parse_args()
+
+
+def read_jsonl(path: str) -> list[dict]:
+ with open(path, encoding="utf-8") as handle:
+ return [json.loads(line) for line in handle if line.strip()]
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ records = read_jsonl(args.nodes)
+ if args.limit:
+ records = records[: args.limit]
+ if not records:
+ raise ValueError("No text nodes found")
+ for record in records:
+ if args.tier not in record:
+ raise ValueError(
+ f"Tier {args.tier!r} is absent; regenerate bundles with "
+ "--with-extra-tiers if needed"
+ )
+ if len(record[args.tier]) < args.views:
+ raise ValueError(
+ f"Node {record['node_id']} has fewer than {args.views} views"
+ )
+
+ tokenizer = AutoTokenizer.from_pretrained(args.model)
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ tokenizer.padding_side = "right"
+ dtype = dtype_for_device(args.device)
+ model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(
+ args.device
+ )
+ model.eval()
+
+ flat_texts = [
+ text
+ for record in records
+ for text in record[args.tier][: args.views]
+ ]
+ outputs: list[torch.Tensor] = []
+ for start in tqdm(
+ range(0, len(flat_texts), args.batch_size),
+ desc=f"Qwen VG text:{args.tier}",
+ ):
+ tokens = tokenizer(
+ flat_texts[start : start + args.batch_size],
+ padding=True,
+ truncation=True,
+ max_length=args.max_length,
+ return_tensors="pt",
+ )
+ tokens = {key: value.to(args.device) for key, value in tokens.items()}
+ result = model(
+ **tokens, output_hidden_states=True, return_dict=True
+ )
+ outputs.append(
+ mean_pool(
+ result.hidden_states[args.layer], tokens["attention_mask"]
+ )
+ .float()
+ .cpu()
+ )
+ features = torch.cat(outputs).reshape(len(records), args.views, -1)
+ state = {
+ "model": args.model,
+ "layer": args.layer,
+ "tier": args.tier,
+ "node_ids": [record["node_id"] for record in records],
+ "region_features": features,
+ "views_per_node": args.views,
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ print(f"Wrote {args.output}: {tuple(features.shape)}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_extract_vision.py b/worldalign/vg_extract_vision.py
new file mode 100644
index 0000000..504840a
--- /dev/null
+++ b/worldalign/vg_extract_vision.py
@@ -0,0 +1,150 @@
+from __future__ import annotations
+
+import argparse
+from concurrent.futures import ThreadPoolExecutor, as_completed
+import json
+from pathlib import Path
+import time
+
+from PIL import Image
+import requests
+import torch
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel
+
+from .common import dtype_for_device
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument(
+ "--nodes", default="artifacts/vg/vision_nodes.private.jsonl"
+ )
+ p.add_argument("--output", default="artifacts/vg/vision_features.pt")
+ p.add_argument(
+ "--image-cache", default="/tmp/yurenh2-worldalign-vg-images"
+ )
+ p.add_argument("--model", default="facebook/dinov2-small")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--node-batch-size", type=int, default=8)
+ p.add_argument("--download-workers", type=int, default=32)
+ p.add_argument("--limit", type=int)
+ return p.parse_args()
+
+
+def read_jsonl(path: str) -> list[dict]:
+ with open(path, encoding="utf-8") as handle:
+ return [json.loads(line) for line in handle if line.strip()]
+
+
+def download_one(record: dict, image_cache: Path) -> Path:
+ image_cache.mkdir(parents=True, exist_ok=True)
+ output = image_cache / f"{int(record['source_image_id'])}.jpg"
+ if output.exists() and output.stat().st_size > 0:
+ return output
+ temporary = output.with_suffix(".jpg.part")
+ last_error: Exception | None = None
+ for attempt in range(3):
+ try:
+ response = requests.get(record["url"], timeout=30)
+ response.raise_for_status()
+ temporary.write_bytes(response.content)
+ with Image.open(temporary) as image:
+ image.verify()
+ temporary.replace(output)
+ return output
+ except Exception as error:
+ last_error = error
+ if temporary.exists():
+ temporary.unlink()
+ time.sleep(1 + attempt)
+ raise RuntimeError(f"Failed to download {record['url']}") from last_error
+
+
+def crop_views(record: dict, path: Path) -> tuple[Image.Image, list[Image.Image]]:
+ with Image.open(path) as source:
+ image = source.convert("RGB")
+ width, height = image.size
+ crops: list[Image.Image] = []
+ for region in record["regions"]:
+ x0 = max(0, min(int(region["x"]), width - 1))
+ y0 = max(0, min(int(region["y"]), height - 1))
+ x1 = max(x0 + 1, min(x0 + int(region["width"]), width))
+ y1 = max(y0 + 1, min(y0 + int(region["height"]), height))
+ crops.append(image.crop((x0, y0, x1, y1)))
+ return image, crops
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ records = read_jsonl(args.nodes)
+ if args.limit:
+ records = records[: args.limit]
+ if not records:
+ raise ValueError("No vision nodes found")
+ view_count = len(records[0]["regions"])
+ if any(len(record["regions"]) != view_count for record in records):
+ raise ValueError("All nodes must have the same number of region views")
+
+ image_cache = Path(args.image_cache)
+ paths: dict[str, Path] = {}
+ with ThreadPoolExecutor(max_workers=args.download_workers) as executor:
+ futures = {
+ executor.submit(download_one, record, image_cache): record["node_id"]
+ for record in records
+ }
+ for future in tqdm(
+ as_completed(futures), total=len(futures), desc="VG image download"
+ ):
+ paths[futures[future]] = future.result()
+
+ processor = AutoImageProcessor.from_pretrained(args.model)
+ dtype = dtype_for_device(args.device)
+ model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(
+ args.device
+ )
+ model.eval()
+
+ region_outputs: list[torch.Tensor] = []
+ global_outputs: list[torch.Tensor] = []
+ for start in tqdm(
+ range(0, len(records), args.node_batch_size),
+ desc="DINO VG bundles",
+ ):
+ batch = records[start : start + args.node_batch_size]
+ images: list[Image.Image] = []
+ for record in batch:
+ global_image, crops = crop_views(record, paths[record["node_id"]])
+ images.append(global_image)
+ images.extend(crops)
+ pixels = processor(images=images, return_tensors="pt")[
+ "pixel_values"
+ ].to(args.device, dtype=dtype)
+ hidden = model(pixel_values=pixels, return_dict=True).last_hidden_state[
+ :, 0
+ ]
+ hidden = hidden.float().cpu().reshape(
+ len(batch), view_count + 1, -1
+ )
+ global_outputs.append(hidden[:, 0])
+ region_outputs.append(hidden[:, 1:])
+
+ state = {
+ "model": args.model,
+ "node_ids": [record["node_id"] for record in records],
+ "region_features": torch.cat(region_outputs),
+ "global_features": torch.cat(global_outputs),
+ "views_per_node": view_count,
+ "source_metadata_removed": True,
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ print(
+ f"Wrote {args.output}: regions={tuple(state['region_features'].shape)}, "
+ f"global={tuple(state['global_features'].shape)}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_graph_diagnose.py b/worldalign/vg_graph_diagnose.py
new file mode 100644
index 0000000..3a0ad2e
--- /dev/null
+++ b/worldalign/vg_graph_diagnose.py
@@ -0,0 +1,222 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+from scipy import sparse
+from scipy.sparse.linalg import eigsh
+import torch
+
+from .common import normalized, write_json
+from .vg_diagnose import (
+ aligned_text_indices,
+ arbitrary_target_retrieval,
+ bundle_signature,
+ rank_normalize_columns,
+)
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument(
+ "--vision", default="artifacts/vg/vision_features.pt"
+ )
+ parser.add_argument("--text", default="artifacts/vg/text_features.pt")
+ parser.add_argument(
+ "--truth", default="artifacts/vg/ground_truth.private.jsonl"
+ )
+ parser.add_argument("--max-nodes", type=int, default=5_000)
+ parser.add_argument("--knn", type=int, default=32)
+ parser.add_argument("--eigenpairs", type=int, default=96)
+ parser.add_argument("--heat-scales", type=int, default=24)
+ parser.add_argument("--message-passing-steps", type=int, default=3)
+ parser.add_argument("--output", default="artifacts/vg/graph_diagnostics.json")
+ return parser.parse_args()
+
+
+def scene_features(state: dict, indices: torch.Tensor | None = None) -> torch.Tensor:
+ if "global_features" in state:
+ features = state["global_features"]
+ else:
+ features = state["region_features"].mean(1)
+ return features if indices is None else features[indices]
+
+
+def exact_rank_knn_graph(
+ features: torch.Tensor, neighbors: int
+) -> tuple[sparse.csr_matrix, torch.Tensor]:
+ """Build a symmetrized rank-weighted graph and retain density statistics."""
+ x = normalized(features.float())
+ similarity = x @ x.T
+ similarity.fill_diagonal_(-torch.inf)
+ k = min(neighbors, len(x) - 1)
+ values, indices = similarity.topk(k, dim=1)
+
+ # Rank weights avoid assuming that cosine scales agree across modalities.
+ rank_weight = np.exp(
+ -np.arange(k, dtype=np.float64) / max(k / 4.0, 1.0)
+ )
+ rows = np.repeat(np.arange(len(x), dtype=np.int64), k)
+ cols = indices.cpu().numpy().reshape(-1)
+ data = np.tile(rank_weight, len(x))
+ directed = sparse.csr_matrix(
+ (data, (rows, cols)), shape=(len(x), len(x))
+ )
+ adjacency = directed.maximum(directed.T)
+ adjacency.setdiag(0)
+ adjacency.eliminate_zeros()
+
+ density_quantiles = torch.quantile(
+ values,
+ torch.tensor([0.0, 0.25, 0.5, 0.75, 1.0]),
+ dim=1,
+ ).T
+ return adjacency, density_quantiles
+
+
+def heat_kernel_signature(
+ adjacency: sparse.csr_matrix, eigenpairs: int, scales: int
+) -> tuple[torch.Tensor, np.ndarray]:
+ """Return basis-invariant diagonal heat-kernel signatures."""
+ degree = np.asarray(adjacency.sum(axis=1)).reshape(-1)
+ inv_sqrt_degree = 1.0 / np.sqrt(np.maximum(degree, 1e-12))
+ normalizer = sparse.diags(inv_sqrt_degree)
+ normalized_adjacency = normalizer @ adjacency @ normalizer
+ k = min(eigenpairs, adjacency.shape[0] - 2)
+ eigenvalues, eigenvectors = eigsh(
+ normalized_adjacency, k=k, which="LA", tol=1e-4
+ )
+ order = np.argsort(eigenvalues)[::-1]
+ eigenvalues = eigenvalues[order]
+ eigenvectors = eigenvectors[:, order]
+ laplacian_eigenvalues = np.clip(1.0 - eigenvalues, 0.0, None)
+
+ positive = laplacian_eigenvalues[laplacian_eigenvalues > 1e-8]
+ lower = 0.25 / max(float(positive.max(initial=1.0)), 1e-8)
+ upper = 8.0 / max(float(positive.min(initial=1.0)), 1e-8)
+ times = np.geomspace(lower, upper, scales)
+ heat = (eigenvectors**2) @ np.exp(
+ -laplacian_eigenvalues[:, None] * times[None, :]
+ )
+ return torch.from_numpy(heat).float(), degree
+
+
+def graph_signature(
+ features: torch.Tensor,
+ neighbors: int,
+ eigenpairs: int,
+ heat_scales: int,
+ message_passing_steps: int,
+) -> torch.Tensor:
+ adjacency, density = exact_rank_knn_graph(features, neighbors)
+ heat, degree = heat_kernel_signature(
+ adjacency, eigenpairs=eigenpairs, scales=heat_scales
+ )
+ base = rank_normalize_columns(
+ torch.cat(
+ [
+ density,
+ torch.from_numpy(np.log1p(degree)).float()[:, None],
+ heat,
+ ],
+ dim=1,
+ )
+ )
+ transition = sparse.diags(
+ 1.0 / np.maximum(np.asarray(adjacency.sum(axis=1)).reshape(-1), 1e-12)
+ ) @ adjacency
+ current = base
+ for _ in range(message_passing_steps):
+ current_np = current.numpy()
+ neighbor_mean = torch.from_numpy(transition @ current_np).float()
+ neighbor_square_mean = torch.from_numpy(
+ transition @ np.square(current_np)
+ ).float()
+ neighbor_std = torch.sqrt(
+ torch.clamp(neighbor_square_mean - neighbor_mean.square(), min=0)
+ )
+ current = rank_normalize_columns(
+ torch.cat([base, neighbor_mean, neighbor_std], dim=1)
+ )
+ return current
+
+
+def main() -> None:
+ args = parse_args()
+ vision = torch.load(args.vision, map_location="cpu", weights_only=False)
+ text = torch.load(args.text, map_location="cpu", weights_only=False)
+ n = min(args.max_nodes, len(vision["node_ids"]), len(text["node_ids"]))
+ vision_ids = vision["node_ids"][:n]
+ target_full = aligned_text_indices(
+ vision_ids, text["node_ids"], args.truth
+ )
+ candidate_indices = torch.unique(target_full, sorted=False)
+ if len(candidate_indices) != n:
+ raise ValueError("Ground truth is not a one-to-one permutation")
+ candidate_lookup = {
+ int(old): new for new, old in enumerate(candidate_indices.tolist())
+ }
+ targets = torch.tensor(
+ [candidate_lookup[int(old)] for old in target_full.tolist()]
+ )
+
+ vision_scene = scene_features(vision)[:n]
+ text_scene = scene_features(text, candidate_indices)
+ vision_graph = graph_signature(
+ vision_scene,
+ args.knn,
+ args.eigenpairs,
+ args.heat_scales,
+ args.message_passing_steps,
+ )
+ text_graph = graph_signature(
+ text_scene,
+ args.knn,
+ args.eigenpairs,
+ args.heat_scales,
+ args.message_passing_steps,
+ )
+
+ vision_bundle = rank_normalize_columns(
+ bundle_signature(vision["region_features"][:n])
+ )
+ text_bundle = rank_normalize_columns(
+ bundle_signature(text["region_features"][candidate_indices])
+ )
+ combined_vision = torch.cat(
+ [normalized(vision_graph), normalized(vision_bundle)], dim=1
+ )
+ combined_text = torch.cat(
+ [normalized(text_graph), normalized(text_bundle)], dim=1
+ )
+
+ result = {
+ "nodes": n,
+ "knn": min(args.knn, n - 1),
+ "eigenpairs": min(args.eigenpairs, n - 2),
+ "heat_scales": args.heat_scales,
+ "message_passing_steps": args.message_passing_steps,
+ "graph_signature_retrieval": arbitrary_target_retrieval(
+ vision_graph, text_graph, targets
+ ),
+ "bundle_signature_retrieval": arbitrary_target_retrieval(
+ vision_bundle, text_bundle, targets
+ ),
+ "combined_signature_retrieval": arbitrary_target_retrieval(
+ combined_vision, combined_text, targets
+ ),
+ "evaluation_note": (
+ "Each signature is computed independently from one modality. "
+ "Ground truth is loaded only after graph construction to score "
+ "the hidden permutation."
+ ),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_prepare.py b/worldalign/vg_prepare.py
new file mode 100644
index 0000000..7d553a8
--- /dev/null
+++ b/worldalign/vg_prepare.py
@@ -0,0 +1,265 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+import random
+from typing import Any, Iterable
+
+from datasets import Dataset, load_dataset
+
+from .common import write_json
+
+
+DATASET = "ranjaykrishna/visual_genome"
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--output-dir", default="artifacts/vg")
+ p.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf")
+ p.add_argument("--nodes", type=int, default=100_000)
+ p.add_argument("--views", type=int, default=32)
+ p.add_argument("--min-regions", type=int, default=32)
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument(
+ "--with-extra-tiers",
+ action="store_true",
+ help="Also load visible relation/attribute and QA annotations.",
+ )
+ return p.parse_args()
+
+
+def load_config(name: str, cache_dir: str) -> Dataset:
+ return load_dataset(
+ DATASET,
+ name,
+ split="train",
+ trust_remote_code=True,
+ with_image=False,
+ cache_dir=cache_dir,
+ )
+
+
+def by_image_id(dataset: Dataset, field: str) -> dict[int, list[dict[str, Any]]]:
+ return {
+ int(image_id): values
+ for image_id, values in zip(dataset["image_id"], dataset[field])
+ }
+
+
+def first_name(obj: dict[str, Any]) -> str:
+ names = obj.get("names") or []
+ return str(names[0]).strip() if names else ""
+
+
+def visible_relation_texts(
+ image_id: int,
+ relationships: dict[int, list[dict[str, Any]]],
+ attributes: dict[int, list[dict[str, Any]]],
+) -> list[str]:
+ texts: list[str] = []
+ for rel in relationships.get(image_id, []):
+ subject = first_name(rel["subject"])
+ predicate = str(rel.get("predicate") or "").strip()
+ object_ = first_name(rel["object"])
+ if subject and predicate and object_:
+ texts.append(f"{subject} {predicate} {object_}.")
+ for obj in attributes.get(image_id, []):
+ name = first_name(obj)
+ for attribute in obj.get("attributes") or []:
+ attribute = str(attribute).strip()
+ if name and attribute:
+ texts.append(f"{name} is {attribute}.")
+ return texts
+
+
+def qa_texts(
+ image_id: int, questions: dict[int, list[dict[str, Any]]]
+) -> list[str]:
+ texts: list[str] = []
+ for item in questions.get(image_id, []):
+ question = str(item.get("question") or "").strip()
+ answer = str(item.get("answer") or "").strip()
+ if question and answer:
+ texts.append(f"Question: {question} Answer: {answer}")
+ return texts
+
+
+def fixed_mixture(
+ groups: list[list[str]], budget: int, rng: random.Random
+) -> list[str]:
+ """Draw a fixed-size, approximately balanced mixture without count leakage."""
+ groups = [list(dict.fromkeys(group)) for group in groups if group]
+ for group in groups:
+ rng.shuffle(group)
+ selected: list[str] = []
+ while len(selected) < budget and groups:
+ next_groups: list[list[str]] = []
+ for group in groups:
+ if group and len(selected) < budget:
+ selected.append(group.pop())
+ if group:
+ next_groups.append(group)
+ groups = next_groups
+ if not selected:
+ raise ValueError("Cannot construct a text bundle from empty groups")
+ seed_values = selected.copy()
+ while len(selected) < budget:
+ selected.append(rng.choice(seed_values))
+ rng.shuffle(selected)
+ return selected
+
+
+def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> None:
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("w", encoding="utf-8") as handle:
+ for record in records:
+ handle.write(json.dumps(record, ensure_ascii=False) + "\n")
+
+
+def main() -> None:
+ args = parse_args()
+ if args.views < 1:
+ raise ValueError("--views must be positive")
+ output_dir = Path(args.output_dir)
+ output_dir.mkdir(parents=True, exist_ok=True)
+ rng = random.Random(args.seed)
+
+ region_ds = load_config("region_descriptions_v1.2.0", args.cache_dir)
+ relation_lookup: dict[int, list[dict[str, Any]]] = {}
+ attribute_lookup: dict[int, list[dict[str, Any]]] = {}
+ question_lookup: dict[int, list[dict[str, Any]]] = {}
+ if args.with_extra_tiers:
+ relation_lookup = by_image_id(
+ load_config("relationships_v1.2.0", args.cache_dir),
+ "relationships",
+ )
+ attribute_lookup = by_image_id(
+ load_config("attributes_v1.2.0", args.cache_dir),
+ "attributes",
+ )
+ question_lookup = by_image_id(
+ load_config("question_answers_v1.2.0", args.cache_dir),
+ "qas",
+ )
+
+ eligible = [
+ idx
+ for idx, regions in enumerate(region_ds["regions"])
+ if len(regions) >= args.min_regions
+ ]
+ rng.shuffle(eligible)
+ selected = eligible[: min(args.nodes, len(eligible))]
+
+ # The two orders and all opaque identifiers are independent. Source IDs
+ # appear only in the vision preprocessing file and private truth file.
+ vision_order = selected.copy()
+ text_order = selected.copy()
+ rng.shuffle(vision_order)
+ rng.shuffle(text_order)
+ vision_opaque = {idx: f"v{rank:07d}" for rank, idx in enumerate(vision_order)}
+ text_opaque = {idx: f"t{rank:07d}" for rank, idx in enumerate(text_order)}
+
+ vision_records: list[dict[str, Any]] = []
+ text_records: list[dict[str, Any]] = []
+ truth: list[dict[str, Any]] = []
+
+ for idx in selected:
+ row = region_ds[int(idx)]
+ regions = list(row["regions"])
+ # Fix the view count and use a deterministic per-node sample. This
+ # removes region-count and ordering side channels.
+ node_rng = random.Random(args.seed ^ int(row["image_id"]))
+ node_rng.shuffle(regions)
+ regions = regions[: args.views]
+
+ visual_regions = [
+ {
+ "x": int(region["x"]),
+ "y": int(region["y"]),
+ "width": int(region["width"]),
+ "height": int(region["height"]),
+ }
+ for region in regions
+ ]
+ phrases = [str(region["phrase"]).strip() for region in regions]
+ node_rng.shuffle(phrases)
+ image_id = int(row["image_id"])
+ v_id = vision_opaque[idx]
+ t_id = text_opaque[idx]
+ vision_records.append(
+ {
+ "node_id": v_id,
+ "source_image_id": image_id,
+ "url": row["url"],
+ "width": int(row["width"]),
+ "height": int(row["height"]),
+ "regions": visual_regions,
+ }
+ )
+ text_record: dict[str, Any] = {
+ "node_id": t_id,
+ "region_closed": phrases,
+ }
+ if args.with_extra_tiers:
+ relation_text = visible_relation_texts(
+ image_id, relation_lookup, attribute_lookup
+ )
+ qa_text = qa_texts(image_id, question_lookup)
+ text_record["visible_relations"] = fixed_mixture(
+ [phrases, relation_text], args.views, node_rng
+ )
+ text_record["qa_expanded"] = fixed_mixture(
+ [phrases, relation_text, qa_text], args.views, node_rng
+ )
+ text_records.append(text_record)
+ truth.append(
+ {
+ "vision_node_id": v_id,
+ "text_node_id": t_id,
+ "source_image_id": image_id,
+ }
+ )
+
+ rng.shuffle(vision_records)
+ rng.shuffle(text_records)
+ rng.shuffle(truth)
+ write_jsonl(output_dir / "vision_nodes.private.jsonl", vision_records)
+ write_jsonl(output_dir / "text_nodes.jsonl", text_records)
+ write_jsonl(output_dir / "ground_truth.private.jsonl", truth)
+
+ nested_sizes = [
+ size for size in (5_000, 20_000, 50_000, 100_000) if size <= len(selected)
+ ]
+ if len(selected) not in nested_sizes:
+ nested_sizes.append(len(selected))
+ manifest = {
+ "dataset": DATASET,
+ "seed": args.seed,
+ "selected_nodes": len(selected),
+ "eligible_nodes": len(eligible),
+ "views_per_node": args.views,
+ "minimum_available_regions": args.min_regions,
+ "extra_closure_tiers": args.with_extra_tiers,
+ "nested_scaling_sizes": sorted(set(nested_sizes)),
+ "training_files": {
+ "vision": "vision_nodes.private.jsonl",
+ "text": "text_nodes.jsonl",
+ },
+ "evaluation_only": "ground_truth.private.jsonl",
+ "leakage_note": (
+ "The source image ID and URL are required only by the vision "
+ "feature preprocessor. Alignment training must consume extracted "
+ "features with source metadata removed."
+ ),
+ }
+ write_json(output_dir / "manifest.json", manifest)
+ print(
+ f"Wrote {len(selected):,} hidden-pair nodes to {output_dir}; "
+ f"{len(eligible):,} scenes met the region threshold."
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/vg_view_probe.py b/worldalign/vg_view_probe.py
new file mode 100644
index 0000000..322644f
--- /dev/null
+++ b/worldalign/vg_view_probe.py
@@ -0,0 +1,292 @@
+"""View-level shared-structure probe for Visual Genome nodes.
+
+The node-level manifold gate shows that pooled node states expose only a
+low-dimensional coarse shared subspace, which cannot identify individual
+assignments. This probe measures whether additional cross-modal signal
+exists inside nodes, at the level of the sixteen region views, where both
+modalities observe the same sixteen regions of the same scene.
+
+The per-node view permutation is replayed from the preparation seed and the
+private image IDs, verified exactly against the released text bundles, and
+used only for evaluation. No cross-modal map is trained.
+
+Three quantities are reported:
+
+1. within-node relational alignment: correlation between the visual and the
+ truth-aligned text view-relation fields, against a shuffled-view control;
+2. coarse-frame view matching: population-context profiles make single
+ views cross-modally comparable through the node-level frame; assignment
+ accuracy is chance-normalized (16 candidates per node);
+3. a within-node relational-only floor with no population context.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import random
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from scipy.optimize import linear_sum_assignment
+
+from .common import write_json
+from .vg_prepare import load_config
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf")
+ parser.add_argument("--seed", type=int, default=20260728)
+ parser.add_argument("--views", type=int, default=16)
+ parser.add_argument(
+ "--context-nodes",
+ type=int,
+ default=1000,
+ help="Node-level frame size for population-context profiles.",
+ )
+ parser.add_argument("--nodes", type=int, default=0, help="0 keeps all nodes.")
+ parser.add_argument(
+ "--skip-floor",
+ action="store_true",
+ help="Skip the slow within-node relational-only descent floor.",
+ )
+ parser.add_argument(
+ "--output", default="artifacts/manifold_gate/vg_view_probe.json"
+ )
+ return parser.parse_args()
+
+
+def replay_view_permutations(args: argparse.Namespace) -> dict:
+ """Recover text-view -> vision-view mappings from the preparation RNG.
+
+ ``random.shuffle`` consumes the generator identically for equal-length
+ lists, so shuffling index lists replays the region subsample and the
+ phrase shuffle exactly. Every node is verified against the released
+ phrase bundle before use.
+ """
+ truth = [
+ json.loads(line)
+ for line in Path(args.vg_dir, "ground_truth.private.jsonl").read_text(
+ encoding="utf-8"
+ ).splitlines()
+ if line.strip()
+ ]
+ text_nodes = {
+ record["node_id"]: record
+ for line in Path(args.vg_dir, "text_nodes.jsonl").read_text(
+ encoding="utf-8"
+ ).splitlines()
+ if line.strip()
+ for record in [json.loads(line)]
+ }
+ region_ds = load_config("region_descriptions_v1.2.0", args.cache_dir)
+ image_rows = {
+ int(image_id): row for row, image_id in enumerate(region_ds["image_id"])
+ }
+ mappings: dict[str, dict] = {}
+ mismatched = 0
+ for record in truth:
+ image_id = int(record["source_image_id"])
+ row = region_ds[image_rows[image_id]]
+ regions = list(row["regions"])
+ node_rng = random.Random(args.seed ^ image_id)
+ order = list(range(len(regions)))
+ node_rng.shuffle(order)
+ selected = order[: args.views]
+ phrase_perm = list(range(args.views))
+ node_rng.shuffle(phrase_perm)
+ phrases_view_order = [
+ str(regions[index]["phrase"]).strip() for index in selected
+ ]
+ reconstructed = [phrases_view_order[j] for j in phrase_perm]
+ released = text_nodes[record["text_node_id"]]["region_closed"]
+ if reconstructed != released:
+ mismatched += 1
+ continue
+ mappings[record["vision_node_id"]] = {
+ "text_node_id": record["text_node_id"],
+ # text view j describes vision view phrase_perm[j]
+ "text_to_vision": phrase_perm,
+ }
+ if mismatched:
+ print(json.dumps({"replay_mismatched_nodes": mismatched}))
+ return mappings
+
+
+def offdiag(matrix: torch.Tensor) -> torch.Tensor:
+ mask = ~torch.eye(len(matrix), dtype=torch.bool)
+ return matrix[mask]
+
+
+def spearman(x: torch.Tensor, y: torch.Tensor) -> float:
+ ranks = torch.stack(
+ [x.argsort().argsort().double(), y.argsort().argsort().double()]
+ )
+ return float(torch.corrcoef(ranks)[0, 1])
+
+
+def main() -> None:
+ args = parse_args()
+ mappings = replay_view_permutations(args)
+
+ vision = torch.load(
+ Path(args.vg_dir, "vision_features.pt"), map_location="cpu", weights_only=False
+ )
+ text = torch.load(
+ Path(args.vg_dir, "text_features.pt"), map_location="cpu", weights_only=False
+ )
+ vision_index = {node: i for i, node in enumerate(vision["node_ids"])}
+ text_index = {node: i for i, node in enumerate(text["node_ids"])}
+ visual_views = F.normalize(vision["region_features"].double(), dim=-1)
+ text_views = F.normalize(text["region_features"].double(), dim=-1)
+
+ node_ids = sorted(mappings)
+ if args.nodes:
+ node_ids = node_ids[: args.nodes]
+
+ # Node-level coarse frame: view-mean states over a fixed context set,
+ # truth-aligned so both profile axes index the same underlying scenes.
+ # The frame is evaluation-side scaffolding for an upper bound; a blind
+ # system would substitute its recovered coarse alignment here.
+ context_ids = node_ids[: args.context_nodes]
+ visual_frame = F.normalize(
+ torch.stack(
+ [visual_views[vision_index[node]].mean(0) for node in context_ids]
+ ),
+ dim=-1,
+ )
+ text_frame = F.normalize(
+ torch.stack(
+ [
+ text_views[text_index[mappings[node]["text_node_id"]]].mean(0)
+ for node in context_ids
+ ]
+ ),
+ dim=-1,
+ )
+
+ generator = torch.Generator().manual_seed(args.seed)
+ relation_matched: list[float] = []
+ relation_shuffled: list[float] = []
+ profile_hits = 0
+ profile_top3 = 0
+ profile_total = 0
+ floor_hits = 0
+ floor_total = 0
+ matched_profile_corr: list[float] = []
+ mismatched_profile_corr: list[float] = []
+
+ for node in node_ids:
+ mapping = mappings[node]
+ v = visual_views[vision_index[node]]
+ t = text_views[text_index[mapping["text_node_id"]]]
+ text_to_vision = torch.tensor(mapping["text_to_vision"])
+ vision_to_text = torch.empty_like(text_to_vision)
+ vision_to_text[text_to_vision] = torch.arange(len(text_to_vision))
+
+ relation_visual = v @ v.T
+ relation_text = t @ t.T
+ aligned = relation_text[vision_to_text][:, vision_to_text]
+ relation_matched.append(spearman(offdiag(relation_visual), offdiag(aligned)))
+ shuffle = torch.randperm(len(v), generator=generator)
+ shuffled = relation_text[shuffle][:, shuffle]
+ relation_shuffled.append(spearman(offdiag(relation_visual), offdiag(shuffled)))
+
+ # Population-context profiles through the coarse frame.
+ profile_visual = v @ visual_frame.T
+ profile_text = t @ text_frame.T
+ rank_visual = profile_visual.argsort(-1).argsort(-1).double()
+ rank_text = profile_text.argsort(-1).argsort(-1).double()
+ rank_visual = rank_visual - rank_visual.mean(-1, keepdim=True)
+ rank_text = rank_text - rank_text.mean(-1, keepdim=True)
+ rank_visual = F.normalize(rank_visual, dim=-1)
+ rank_text = F.normalize(rank_text, dim=-1)
+ similarity = rank_visual @ rank_text.T # [vision view, text view]
+ rows, cols = linear_sum_assignment(-similarity.numpy())
+ truth_cols = vision_to_text.numpy()
+ profile_hits += int((cols == truth_cols).sum())
+ ranks = (
+ similarity
+ >= similarity.gather(1, vision_to_text[:, None])
+ ).sum(-1)
+ profile_top3 += int((ranks <= 3).sum())
+ profile_total += len(rows)
+ matched_profile_corr.extend(
+ similarity.gather(1, vision_to_text[:, None]).squeeze(1).tolist()
+ )
+ mismatched_profile_corr.extend(
+ similarity[~torch.eye(len(v), dtype=torch.bool)][:64].tolist()
+ )
+
+ # Relational-only floor: 2-swap descent on within-node fields from a
+ # random start, no population context.
+ if args.skip_floor:
+ continue
+ state = torch.randperm(len(v), generator=generator)
+ for _ in range(200):
+ improved = False
+ current = relation_text[state][:, state]
+ base = float(((relation_visual - current) ** 2).mean())
+ for p in range(len(v)):
+ for q in range(p + 1, len(v)):
+ trial = state.clone()
+ trial[[p, q]] = trial[[q, p]]
+ candidate = relation_text[trial][:, trial]
+ if float(((relation_visual - candidate) ** 2).mean()) < base - 1e-12:
+ state = trial
+ improved = True
+ break
+ if improved:
+ break
+ if not improved:
+ break
+ floor_hits += int((state == vision_to_text).sum())
+ floor_total += len(v)
+
+ matched = torch.tensor(relation_matched)
+ shuffled = torch.tensor(relation_shuffled)
+ matched_corr = torch.tensor(matched_profile_corr)
+ mismatched_corr = torch.tensor(mismatched_profile_corr)
+ report = {
+ "protocol": (
+ "View permutations are replayed from the preparation seed and "
+ "verified against released phrase bundles; they are used only "
+ "for evaluation. The coarse frame uses truth-aligned node "
+ "states, so profile matching is an upper bound on what a "
+ "recovered node alignment would enable."
+ ),
+ "nodes": len(node_ids),
+ "verified_fraction": len(mappings) / 5000.0,
+ "within_node_relation": {
+ "matched_spearman_mean": float(matched.mean()),
+ "matched_spearman_std": float(matched.std()),
+ "shuffled_spearman_mean": float(shuffled.mean()),
+ "shuffled_spearman_std": float(shuffled.std()),
+ "gap_z_across_nodes": float(
+ (matched.mean() - shuffled.mean())
+ / (matched - shuffled).std().clamp_min(1e-12)
+ * len(matched) ** 0.5
+ ),
+ },
+ "coarse_frame_view_matching": {
+ "context_nodes": len(context_ids),
+ "accuracy": profile_hits / max(profile_total, 1),
+ "top3_accuracy": profile_top3 / max(profile_total, 1),
+ "chance": 1.0 / args.views,
+ "matched_profile_corr_mean": float(matched_corr.mean()),
+ "mismatched_profile_corr_mean": float(mismatched_corr.mean()),
+ },
+ "within_node_relational_floor": {
+ "accuracy": floor_hits / max(floor_total, 1),
+ "chance": 1.0 / args.views,
+ },
+ }
+ write_json(args.output, report)
+ print(json.dumps(report, indent=2))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/view_anchor_affinity.py b/worldalign/view_anchor_affinity.py
new file mode 100644
index 0000000..3c921bb
--- /dev/null
+++ b/worldalign/view_anchor_affinity.py
@@ -0,0 +1,271 @@
+"""C5 cascade step: anchored node affinity from view-level matching.
+
+For every node pair the sixteen-view anchor similarity matrix is built
+from per-view features and frozen world maps, solved by the Hungarian
+method, and the optimal matching value becomes the node affinity. The
+anchor-optimal view assignment can also score the within-node relational
+agreement at that fixed assignment -- fixed-sigma evaluation avoids the
+free-permutation overfit that killed the min-fit affinity.
+
+Hidden pairs score retrieval and seed precision only; the affinity uses
+per-node features and frozen maps, nothing cross-modal that is learned.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from scipy.optimize import linear_sum_assignment
+
+from .common import write_json
+
+CHANNELS = ("color", "size", "light", "horizontal", "vertical")
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument(
+ "--features", default="artifacts/manifold_gate/view_anchor_features.pt"
+ )
+ parser.add_argument(
+ "--structure", default="artifacts/manifold_gate/attn_text_full.pt"
+ )
+ parser.add_argument(
+ "--vision-structure", default="artifacts/manifold_gate/attn_vision_full.pt"
+ )
+ parser.add_argument("--samples", type=int, default=512)
+ parser.add_argument("--subset-seeds", default="0,1,2")
+ parser.add_argument("--relational-weight", type=float, default=0.5)
+ parser.add_argument("--chunk", type=int, default=64)
+ parser.add_argument(
+ "--output", default="artifacts/manifold_gate/view_anchor_affinity.json"
+ )
+ return parser.parse_args()
+
+
+def rank_z_rows(values: torch.Tensor) -> torch.Tensor:
+ order = values.argsort(-1).argsort(-1).double()
+ order = order - order.mean(-1, keepdim=True)
+ std = order.std(-1, keepdim=True).clamp_min(1e-9)
+ return order / std
+
+
+def standardized_fields(views: torch.Tensor) -> torch.Tensor:
+ views = F.normalize(views.double(), dim=-1)
+ relation = views @ views.transpose(-2, -1)
+ size = relation.shape[-1]
+ mask = ~torch.eye(size, dtype=torch.bool)
+ values = relation[..., mask]
+ mean = values.mean(-1, keepdim=True)
+ std = values.std(-1, keepdim=True).clamp_min(1e-9)
+ standardized = (relation - mean[..., None]) / std[..., None]
+ return standardized.masked_fill(~mask, 0.0)
+
+
+def evaluate_subset(
+ data: dict, subset: torch.Tensor, args: argparse.Namespace
+) -> dict:
+ n = len(subset)
+ views = data["vision_color"].shape[1]
+ vision_color = data["vision_color"][subset]
+ text_color = data["text_color"][subset]
+ text_informative = data["text_informative"][subset] # [n, C, V] bool
+ vision_scalar = data["vision_scalar"][subset] # [n, 3, V] rank-z
+ text_scalar = data["text_scalar"][subset] # [n, 3, V] rank-z
+ text_scalar_mask = data["text_scalar_mask"][subset] # [n, 3, V]
+ text_fields = data["text_fields"][subset]
+ vision_fields = data["vision_fields"][subset]
+
+ anchor_cost = torch.zeros(n, n)
+ combined_cost = torch.zeros(n, n)
+ for start in range(0, n, args.chunk):
+ stop = min(start + args.chunk, n)
+ block = slice(start, stop)
+ color = torch.einsum(
+ "avc,btc->abvt", vision_color[block].double(), text_color.double()
+ )
+ color = color * text_informative[None, :, 0, None, :]
+ channels = [color / max(float(color[color != 0].std()), 1e-9) if (color != 0).any() else color]
+ for c in range(4):
+ vision_values = vision_scalar[block, c] # [a, V]
+ text_values = text_scalar[:, c] # [n, V]
+ difference = -(
+ vision_values[:, None, :, None] - text_values[None, :, None, :]
+ ).abs()
+ difference = difference * text_scalar_mask[:, c][None, :, None, :]
+ scale = float(difference[difference != 0].std()) if (difference != 0).any() else 1.0
+ channels.append(difference / max(scale, 1e-9))
+ stacked = torch.stack(channels) # [C, a, b, V, V]
+ present = torch.stack(
+ [
+ text_informative[:, 0].any(-1),
+ text_scalar_mask[:, 0].any(-1),
+ text_scalar_mask[:, 1].any(-1),
+ text_scalar_mask[:, 2].any(-1),
+ text_scalar_mask[:, 3].any(-1),
+ ]
+ ).double() # [C, b]
+ weight = present / present.sum(0, keepdim=True).clamp_min(1.0)
+ anchor_matrix = torch.einsum("cabvt,cb->abvt", stacked, weight)
+ for a in range(stop - start):
+ for b in range(n):
+ matrix = anchor_matrix[a, b].numpy()
+ rows, cols = linear_sum_assignment(-matrix)
+ value = float(matrix[rows, cols].mean())
+ anchor_cost[start + a, b] = -value
+ if args.relational_weight:
+ aligned = text_fields[b][cols][:, cols]
+ relational = float(
+ ((aligned - vision_fields[start + a]) ** 2)[
+ ~torch.eye(views, dtype=torch.bool)
+ ].mean()
+ )
+ combined_cost[start + a, b] = (
+ -value + args.relational_weight * relational
+ )
+ truth = torch.arange(n)
+
+ def metrics(cost: torch.Tensor) -> dict:
+ centered = cost - cost.mean(0, keepdim=True)
+ ranks = (centered <= centered.gather(1, truth[:, None])).sum(-1)
+ rows, cols = linear_sum_assignment(cost.numpy())
+ forward = centered.argmin(-1)
+ backward = centered.argmin(0)
+ mutual = backward[forward] == truth
+ sorted_cost = centered.sort(-1).values
+ margin = sorted_cost[:, 1] - sorted_cost[:, 0]
+ top = margin.argsort(descending=True)[: max(1, n // 10)]
+ true_costs = cost.diagonal()
+ permuted = []
+ rng = np.random.default_rng(0)
+ for _ in range(200):
+ permutation = torch.from_numpy(rng.permutation(n))
+ permuted.append(float(cost[truth, permutation].mean()))
+ permuted = torch.tensor(permuted)
+ return {
+ "r@1": float((ranks <= 1).double().mean()),
+ "r@5": float((ranks <= 5).double().mean()),
+ "r@10": float((ranks <= 10).double().mean()),
+ "median_rank": float(ranks.double().median()),
+ "true_z": float(
+ (permuted.mean() - true_costs.mean()) / permuted.std().clamp_min(1e-12)
+ ),
+ "hungarian_accuracy": float(
+ (torch.from_numpy(cols) == truth).double().mean()
+ ),
+ "mutual_nn_count": int(mutual.sum()),
+ "mutual_nn_precision": float(
+ (forward[mutual] == truth[mutual]).double().mean()
+ )
+ if mutual.any()
+ else None,
+ "top_margin_decile_precision": float(
+ (forward[top] == truth[top]).double().mean()
+ ),
+ }
+
+ result = {"anchors_only": metrics(anchor_cost)}
+ if args.relational_weight:
+ result["anchors_plus_relational"] = metrics(combined_cost)
+ return result
+
+
+def main() -> None:
+ args = parse_args()
+ state = torch.load(args.features, map_location="cpu", weights_only=False)
+ truth_pairs = [
+ json.loads(line)
+ for line in open(f"{args.vg_dir}/ground_truth.private.jsonl", encoding="utf-8")
+ if line.strip()
+ ]
+ text_structure = torch.load(args.structure, map_location="cpu", weights_only=False)
+ vision_structure = torch.load(
+ args.vision_structure, map_location="cpu", weights_only=False
+ )
+ text_index = {node: i for i, node in enumerate(text_structure["node_ids"])}
+ vision_index = {node: i for i, node in enumerate(vision_structure["node_ids"])}
+
+ vision_color, text_color, text_informative = [], [], []
+ vision_scalar, text_scalar, text_scalar_mask = [], [], []
+ text_fields, vision_fields = [], []
+ for pair in truth_pairs:
+ vision_features = state["vision_view_features"][pair["vision_node_id"]]
+ text_features = state["text_view_features"][pair["text_node_id"]]
+ hue = torch.from_numpy(vision_features["hue"])
+ color = torch.from_numpy(text_features["color"])
+ vision_color.append(F.normalize(hue, dim=-1))
+ text_color.append(F.normalize(color, dim=-1))
+ text_informative.append((color.sum(-1) > 0)[None, :])
+ vision_scalar.append(
+ torch.stack(
+ [
+ rank_z_rows(torch.from_numpy(vision_features["area"])[None])[0],
+ rank_z_rows(torch.from_numpy(vision_features["luminance"])[None])[0],
+ rank_z_rows(torch.from_numpy(vision_features["x"])[None])[0],
+ rank_z_rows(torch.from_numpy(vision_features["y"])[None])[0],
+ ]
+ )
+ )
+ scalars, masks = [], []
+ for key in ("size", "light", "horizontal", "vertical"):
+ values = torch.from_numpy(text_features[key])
+ scalars.append(rank_z_rows(values[None])[0])
+ masks.append(values != 0)
+ text_scalar.append(torch.stack(scalars))
+ text_scalar_mask.append(torch.stack(masks))
+ text_fields.append(
+ standardized_fields(
+ torch.as_tensor(
+ text_structure["context_states"][
+ text_index[pair["text_node_id"]]
+ ]
+ )[None]
+ )[0]
+ )
+ vision_fields.append(
+ standardized_fields(
+ torch.as_tensor(
+ vision_structure["context_states"][
+ vision_index[pair["vision_node_id"]]
+ ]
+ )[None]
+ )[0]
+ )
+ data = {
+ "vision_color": torch.stack(vision_color),
+ "text_color": torch.stack(text_color),
+ "text_informative": torch.stack(text_informative),
+ "vision_scalar": torch.stack(vision_scalar),
+ "text_scalar": torch.stack(text_scalar),
+ "text_scalar_mask": torch.stack(text_scalar_mask),
+ "text_fields": torch.stack(text_fields),
+ "vision_fields": torch.stack(vision_fields),
+ }
+ report: dict = {
+ "protocol": (
+ "Node affinity is the Hungarian value of the per-pair view "
+ "anchor matrix (plus optional fixed-assignment relational "
+ "agreement). Hidden pairs score retrieval and seeds only."
+ ),
+ "samples": args.samples,
+ "subsets": [],
+ }
+ for seed in (int(s) for s in args.subset_seeds.split(",")):
+ generator = torch.Generator().manual_seed(seed)
+ subset = torch.randperm(len(truth_pairs), generator=generator)[
+ : args.samples
+ ]
+ result = {"subset_seed": seed, **evaluate_subset(data, subset, args)}
+ report["subsets"].append(result)
+ print(json.dumps(result))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/view_anchor_battery.py b/worldalign/view_anchor_battery.py
new file mode 100644
index 0000000..f039c63
--- /dev/null
+++ b/worldalign/view_anchor_battery.py
@@ -0,0 +1,365 @@
+"""C5 second wave: fine-grained anchors at view level.
+
+The node-level anchor wave failed because scene-level attributes are
+coarse-redundant with the shared subspace. Region-level attributes are
+not: a phrase names the color, relative size, lighting, or position of
+one region, and the region's own pixels and box realize them. With
+sixteen candidates per node, weak per-view anchors carry usable bits.
+
+Per-view anchors, computed independently per modality:
+
+- text, per phrase: color-lexicon histogram, size score, light score,
+ horizontal/vertical position scores, numeral content;
+- vision, per region: crop hue-band histogram, crop luminance, box area
+ fraction, box center coordinates.
+
+Scalar channels are rank-standardized within the node (sixteen values),
+so they compare relative attributes inside one scene; lexicon-to-physics
+maps are frozen world knowledge. Channels contribute only where the
+phrase carries the attribute. Replayed view truth scores matching;
+nothing here uses node pairs or the coarse frame.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import re
+from concurrent.futures import ThreadPoolExecutor
+from pathlib import Path
+
+import numpy as np
+import torch
+from PIL import Image
+from scipy.optimize import linear_sum_assignment
+
+from .anchor_battery import (
+ COLOR_BANDS,
+ LIGHT_WORDS,
+ NUMBER_WORDS,
+ SIZE_WORDS,
+ read_jsonl,
+)
+from .common import write_json
+from .vg_view_probe import replay_view_permutations
+
+HORIZONTAL_WORDS = {"left": -1.0, "right": 1.0}
+VERTICAL_WORDS = {"top": -1.0, "upper": -1.0, "above": -1.0, "bottom": 1.0, "lower": 1.0, "below": 1.0}
+# Twelve color classes: eight hue bands plus achromatic and brown classes
+# with pixel rules on value and saturation. Frozen world knowledge.
+COLOR_CLASSES = list(COLOR_BANDS) + ["white", "black", "gray", "brown"]
+COLOR_SYNONYMS = {"grey": "gray", "tan": "brown", "beige": "brown", "golden": "yellow", "gold": "yellow", "silver": "gray"}
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf")
+ parser.add_argument("--seed", type=int, default=20260728)
+ parser.add_argument("--views", type=int, default=16)
+ parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images")
+ parser.add_argument("--workers", type=int, default=16)
+ parser.add_argument("--structure", default="artifacts/manifold_gate/attn_text_full.pt")
+ parser.add_argument(
+ "--vision-structure", default="artifacts/manifold_gate/attn_vision_full.pt"
+ )
+ parser.add_argument("--nodes", type=int, default=0)
+ parser.add_argument(
+ "--output", default="artifacts/manifold_gate/view_anchor_battery.json"
+ )
+ parser.add_argument(
+ "--anchors-output", default="artifacts/manifold_gate/view_anchor_maps.pt"
+ )
+ return parser.parse_args()
+
+
+def phrase_features(phrase: str) -> dict:
+ tokens = [
+ COLOR_SYNONYMS.get(token, token)
+ for token in re.findall(r"[a-z]+|\d+", phrase.lower())
+ ]
+ color = np.zeros(len(COLOR_CLASSES))
+ for index, name in enumerate(COLOR_CLASSES):
+ color[index] = sum(token == name for token in tokens)
+ numeral = 0.0
+ for token in tokens:
+ if token.isdigit() and int(token) <= 50:
+ numeral += int(token)
+ elif token in NUMBER_WORDS:
+ numeral += NUMBER_WORDS[token]
+ return {
+ "color": color,
+ "size": float(sum(SIZE_WORDS.get(token, 0.0) for token in tokens)),
+ "light": float(sum(LIGHT_WORDS.get(token, 0.0) for token in tokens)),
+ "horizontal": float(sum(HORIZONTAL_WORDS.get(token, 0.0) for token in tokens)),
+ "vertical": float(sum(VERTICAL_WORDS.get(token, 0.0) for token in tokens)),
+ "numeral": numeral,
+ }
+
+
+def region_features(record: dict, args: argparse.Namespace) -> dict:
+ path = Path(args.image_cache, f"{record['source_image_id']}.jpg")
+ with Image.open(path) as image:
+ hsv = np.asarray(image.convert("HSV"), dtype=np.float64)
+ height, width = hsv.shape[:2]
+ hues, lums, areas, xs, ys = [], [], [], [], []
+ for region in record["regions"]:
+ x0 = max(0, min(int(region["x"]), width - 1))
+ y0 = max(0, min(int(region["y"]), height - 1))
+ x1 = max(x0 + 1, min(x0 + int(region["width"]), width))
+ y1 = max(y0 + 1, min(y0 + int(region["height"]), height))
+ crop = hsv[y0:y1, x0:x1]
+ hue = crop[..., 0] * (360.0 / 255.0)
+ saturation = crop[..., 1] / 255.0
+ value = crop[..., 2] / 255.0
+ white = (value > 0.75) & (saturation < 0.25)
+ black = value < 0.25
+ gray = (~white) & (~black) & (saturation < 0.25)
+ brown_hue = np.minimum(np.abs(hue - 30.0), 360.0 - np.abs(hue - 30.0)) < 25.0
+ brown = brown_hue & (saturation >= 0.25) & (value >= 0.2) & (value < 0.55)
+ chromatic = (saturation >= 0.25) & (value >= 0.2) & (~brown)
+ histogram = np.zeros(len(COLOR_CLASSES))
+ for index, center in enumerate(COLOR_BANDS.values()):
+ distance = np.minimum(np.abs(hue - center), 360.0 - np.abs(hue - center))
+ histogram[index] = float(((distance < 25.0) & chromatic).mean())
+ histogram[len(COLOR_BANDS) + 0] = float(white.mean())
+ histogram[len(COLOR_BANDS) + 1] = float(black.mean())
+ histogram[len(COLOR_BANDS) + 2] = float(gray.mean())
+ histogram[len(COLOR_BANDS) + 3] = float(brown.mean())
+ hues.append(histogram)
+ lums.append(float(value.mean()))
+ areas.append((x1 - x0) * (y1 - y0) / (width * height))
+ xs.append((x0 + x1) / 2.0 / width)
+ ys.append((y0 + y1) / 2.0 / height)
+ return {
+ "hue": np.stack(hues),
+ "luminance": np.array(lums),
+ "area": np.array(areas),
+ "x": np.array(xs),
+ "y": np.array(ys),
+ }
+
+
+def rank_z(values: np.ndarray) -> np.ndarray:
+ order = values.argsort().argsort().astype(np.float64)
+ std = order.std()
+ return (order - order.mean()) / (std if std > 1e-9 else 1.0)
+
+
+def node_anchor_matrix(
+ text: list[dict], vision: dict
+) -> tuple[np.ndarray, np.ndarray]:
+ """[V, V] anchor similarity (vision rows, text columns) and coverage."""
+ views = len(text)
+ channels = []
+ coverage = np.zeros(5)
+
+ text_color = np.stack([p["color"] for p in text])
+ informative = text_color.sum(1) > 0
+ if informative.any():
+ tn = text_color / np.linalg.norm(text_color, axis=1, keepdims=True).clip(min=1e-9)
+ vn = vision["hue"] / np.linalg.norm(vision["hue"], axis=1, keepdims=True).clip(min=1e-9)
+ color = vn @ tn.T
+ color[:, ~informative] = 0.0
+ flat = color[:, informative]
+ scale = flat.std()
+ channels.append(color / (scale if scale > 1e-9 else 1.0))
+ coverage[0] = informative.mean()
+
+ for slot, (text_key, vision_values, sign) in enumerate(
+ (
+ ("size", rank_z(vision["area"]), 1.0),
+ ("light", rank_z(vision["luminance"]), 1.0),
+ ("horizontal", rank_z(vision["x"]), 1.0),
+ ("vertical", rank_z(vision["y"]), 1.0),
+ ),
+ start=1,
+ ):
+ scores = np.array([p[text_key] for p in text])
+ informative = scores != 0
+ if informative.sum() < 2:
+ continue
+ text_rank = rank_z(scores)
+ similarity = -np.abs(vision_values[:, None] - sign * text_rank[None, :])
+ similarity[:, ~informative] = 0.0
+ flat = similarity[:, informative]
+ scale = flat.std()
+ channels.append(similarity / (scale if scale > 1e-9 else 1.0))
+ coverage[slot] = informative.mean()
+ if not channels:
+ return np.zeros((views, views)), coverage
+ return np.mean(channels, axis=0), coverage
+
+
+def structure_signature_matrix(
+ text_views: torch.Tensor, visual_views: torch.Tensor
+) -> np.ndarray:
+ """Frame-free structural channel: sorted within-node relation rows."""
+ def signatures(views: torch.Tensor) -> torch.Tensor:
+ views = torch.nn.functional.normalize(views.double(), dim=-1)
+ relation = views @ views.T
+ size = len(relation)
+ mask = ~torch.eye(size, dtype=torch.bool)
+ rows = relation.masked_select(mask).reshape(size, size - 1)
+ rows = (rows - rows.mean()) / rows.std().clamp_min(1e-9)
+ return rows.sort(-1).values
+
+ a = signatures(visual_views)
+ b = signatures(text_views)
+ cost = ((a[:, None, :] - b[None, :, :]) ** 2).sum(-1)
+ similarity = -cost
+ return ((similarity - similarity.mean()) / similarity.std().clamp_min(1e-9)).numpy()
+
+
+def main() -> None:
+ args = parse_args()
+ mappings = replay_view_permutations(args)
+ text_nodes = {
+ record["node_id"]: record
+ for record in read_jsonl(Path(args.vg_dir, "text_nodes.jsonl"))
+ }
+ vision_nodes = {
+ record["node_id"]: record
+ for record in read_jsonl(Path(args.vg_dir, "vision_nodes.private.jsonl"))
+ }
+ text_structure = torch.load(args.structure, map_location="cpu", weights_only=False)
+ vision_structure = torch.load(
+ args.vision_structure, map_location="cpu", weights_only=False
+ )
+ text_index = {node: i for i, node in enumerate(text_structure["node_ids"])}
+ vision_index = {node: i for i, node in enumerate(vision_structure["node_ids"])}
+
+ node_ids = sorted(mappings)
+ if args.nodes:
+ node_ids = node_ids[: args.nodes]
+
+ vision_records = [vision_nodes[node] for node in node_ids]
+ with ThreadPoolExecutor(max_workers=args.workers) as pool:
+ vision_results = list(
+ pool.map(lambda record: region_features(record, args), vision_records)
+ )
+
+ generator = np.random.default_rng(args.seed)
+ accuracies = {"anchors": [], "structure": [], "anchors_plus_structure": []}
+ matched_scores, shuffled_scores = [], []
+ coverage_totals = np.zeros(5)
+ anchor_maps = {}
+ informative_view_hits: list[bool] = []
+ seed_pool: list[tuple[float, bool]] = []
+ text_view_features: dict[str, dict] = {}
+ vision_view_features: dict[str, dict] = {}
+
+ for node, vision_features_one in zip(node_ids, vision_results):
+ mapping = mappings[node]
+ phrases = text_nodes[mapping["text_node_id"]]["region_closed"]
+ text_features_list = [phrase_features(p) for p in phrases]
+ anchors, coverage = node_anchor_matrix(text_features_list, vision_features_one)
+ coverage_totals += coverage
+ text_to_vision = np.array(mapping["text_to_vision"])
+ vision_to_text = np.empty_like(text_to_vision)
+ vision_to_text[text_to_vision] = np.arange(len(text_to_vision))
+
+ structure = structure_signature_matrix(
+ torch.as_tensor(
+ text_structure["context_states"][text_index[mapping["text_node_id"]]]
+ ),
+ torch.as_tensor(vision_structure["context_states"][vision_index[node]]),
+ )
+ combined = anchors + structure
+
+ informative_columns = np.abs(anchors).sum(0) > 1e-9
+ for name, matrix in (
+ ("anchors", anchors),
+ ("structure", structure),
+ ("anchors_plus_structure", combined),
+ ):
+ rows, cols = linear_sum_assignment(-matrix)
+ accuracies[name].append(float((cols == vision_to_text).mean()))
+ if name == "anchors" and informative_columns.any():
+ assigned_text = cols # vision row i -> text column
+ correct = assigned_text == vision_to_text
+ text_informative_hit = informative_columns[assigned_text]
+ informative_view_hits.extend(
+ correct[text_informative_hit].tolist()
+ )
+ sorted_scores = np.sort(matrix, axis=1)[:, ::-1]
+ margins = sorted_scores[:, 0] - sorted_scores[:, 1]
+ for row in range(len(matrix)):
+ if informative_columns[assigned_text[row]]:
+ seed_pool.append(
+ (float(margins[row]), bool(correct[row]))
+ )
+ matched = anchors[np.arange(len(anchors)), vision_to_text]
+ shuffle = generator.permutation(len(anchors))
+ matched_scores.append(float(matched.mean()))
+ shuffled_scores.append(
+ float(anchors[np.arange(len(anchors)), shuffle].mean())
+ )
+ anchor_maps[node] = torch.from_numpy(anchors).float()
+ text_view_features[mapping["text_node_id"]] = {
+ "color": np.stack([p["color"] for p in text_features_list]),
+ "size": np.array([p["size"] for p in text_features_list]),
+ "light": np.array([p["light"] for p in text_features_list]),
+ "horizontal": np.array([p["horizontal"] for p in text_features_list]),
+ "vertical": np.array([p["vertical"] for p in text_features_list]),
+ "numeral": np.array([p["numeral"] for p in text_features_list]),
+ }
+ vision_view_features[node] = vision_features_one
+
+ matched_arr = np.array(matched_scores)
+ shuffled_arr = np.array(shuffled_scores)
+ report = {
+ "protocol": (
+ "Per-view anchors and frame-free structural signatures; "
+ "replayed view truth scores matching only. No node pairs, no "
+ "coarse frame."
+ ),
+ "nodes": len(node_ids),
+ "chance": 1.0 / args.views,
+ "true_frame_profile_baseline": 0.1448,
+ "channel_coverage_mean": {
+ "color": float(coverage_totals[0] / len(node_ids)),
+ "size": float(coverage_totals[1] / len(node_ids)),
+ "light": float(coverage_totals[2] / len(node_ids)),
+ "horizontal": float(coverage_totals[3] / len(node_ids)),
+ "vertical": float(coverage_totals[4] / len(node_ids)),
+ },
+ "anchor_matched_vs_shuffled_z": float(
+ (matched_arr - shuffled_arr).mean()
+ / (matched_arr - shuffled_arr).std().clip(min=1e-12)
+ * np.sqrt(len(matched_arr))
+ ),
+ "view_matching_accuracy": {
+ name: float(np.mean(values)) for name, values in accuracies.items()
+ },
+ "informative_view_accuracy": float(np.mean(informative_view_hits))
+ if informative_view_hits
+ else None,
+ "informative_view_count": len(informative_view_hits),
+ }
+ if seed_pool:
+ seed_pool.sort(key=lambda item: -item[0])
+ calibration = {}
+ for fraction in (0.01, 0.05, 0.10, 0.25):
+ k = max(1, int(len(seed_pool) * fraction))
+ calibration[f"top_{fraction:.0%}_margin"] = {
+ "count": k,
+ "precision": float(np.mean([hit for _, hit in seed_pool[:k]])),
+ }
+ report["seed_calibration"] = calibration
+ torch.save(
+ {
+ "anchor_maps": anchor_maps,
+ "text_view_features": text_view_features,
+ "vision_view_features": vision_view_features,
+ "color_classes": COLOR_CLASSES,
+ "order_note": "vision rows, text columns, released view order",
+ },
+ args.anchors_output,
+ )
+ write_json(args.output, report)
+ print(json.dumps(report, indent=2))
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/view_gate.py b/worldalign/view_gate.py
new file mode 100644
index 0000000..a09fdca
--- /dev/null
+++ b/worldalign/view_gate.py
@@ -0,0 +1,436 @@
+"""View-level assignment gate on structured Visual Genome node states.
+
+A node state here is the set of sixteen region views with its internal
+standardized relation field, not a pooled vector. Two questions are gated,
+per the current experimental gate in LAB_NOTES:
+
+A. Ordering at view level, given the true coarse frame (evaluation upper
+ bound): is the true view assignment a local optimum of (1) the
+ within-node relational energy, (2) the population-context profile
+ energy, (3) their combination? All 120 transpositions per node are
+ enumerated exactly.
+
+B. Blind structured-state node affinity, frame-free: the matching cost
+ between two nodes' internal relation fields (signature-initialized
+ Hungarian plus exact batched 2-swap descent) is used as a node-level
+ linear assignment energy. Retrieval, global Hungarian recovery, and
+ margin-calibrated precision are compared against the pooled-state and
+ graph-signature baselines.
+
+Replayed view truth and node pairs are evaluation-only throughout.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+
+import torch
+import torch.nn.functional as F
+from scipy.optimize import linear_sum_assignment
+
+from .common import seed_everything, write_json
+from .vg_view_probe import replay_view_permutations
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--cache-dir", default="/tmp/yurenh2-worldalign-vg-hf")
+ parser.add_argument("--seed", type=int, default=20260728)
+ parser.add_argument("--views", type=int, default=16)
+ parser.add_argument("--part-a-nodes", type=int, default=5000)
+ parser.add_argument("--context-nodes", type=int, default=1000)
+ parser.add_argument("--joint-weights", default="0.5,2.0")
+ parser.add_argument("--part-b-nodes", type=int, default=512)
+ parser.add_argument("--part-b-seeds", default="0,1,2")
+ parser.add_argument("--descent-sweeps", type=int, default=40)
+ parser.add_argument("--chunk", type=int, default=48)
+ parser.add_argument("--output", default="artifacts/manifold_gate/view_gate.json")
+ parser.add_argument(
+ "--matrices-output",
+ default="",
+ help="Optional .pt path stem for saving raw/null cost matrices.",
+ )
+ return parser.parse_args()
+
+
+def standardized_field(views: torch.Tensor) -> torch.Tensor:
+ """Within-node standardized relation field with zeroed diagonal."""
+ relation = views @ views.transpose(-2, -1)
+ size = relation.shape[-1]
+ mask = ~torch.eye(size, dtype=torch.bool)
+ values = relation[..., mask]
+ mean = values.mean(-1, keepdim=True)
+ std = values.std(-1, keepdim=True).clamp_min(1e-9)
+ standardized = (relation - mean[..., None]) / std[..., None]
+ return standardized.masked_fill(~mask, 0.0)
+
+
+def batched_transposition_delta(
+ text_fields: torch.Tensor, visual_fields: torch.Tensor
+) -> torch.Tensor:
+ """Exact MSE delta of every transposition for a batch of field pairs.
+
+ Same cancellation as manifold_gate.all_transposition_delta_mse, with a
+ batch dimension: delta[b] applies to (text_fields[b], visual_fields[b]).
+ """
+ size = text_fields.shape[-1]
+ count = size * (size - 1)
+ cross = torch.bmm(text_fields, visual_fields)
+ self_terms = (text_fields * visual_fields).sum(-1)
+ corrections = 2.0 * text_fields * visual_fields
+ total = (
+ self_terms[:, :, None]
+ + self_terms[:, None, :]
+ - cross
+ - cross.transpose(-2, -1)
+ - corrections
+ )
+ delta = (4.0 / count) * total
+ diagonal = torch.eye(size, dtype=torch.bool)
+ return delta.masked_fill(diagonal, 0.0)
+
+
+def field_mse(text_fields: torch.Tensor, visual_fields: torch.Tensor) -> torch.Tensor:
+ size = text_fields.shape[-1]
+ mask = ~torch.eye(size, dtype=torch.bool)
+ return ((text_fields - visual_fields) ** 2)[..., mask].mean(-1)
+
+
+def batched_descent(
+ text_fields: torch.Tensor,
+ visual_fields: torch.Tensor,
+ permutations: torch.Tensor,
+ sweeps: int,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Steepest 2-swap descent run in parallel over a batch of field pairs."""
+ batch, size = permutations.shape
+ index = torch.arange(batch)
+ for _ in range(sweeps):
+ permuted = text_fields[
+ index[:, None, None],
+ permutations[:, :, None],
+ permutations[:, None, :],
+ ]
+ delta = batched_transposition_delta(permuted, visual_fields)
+ upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1)
+ masked = delta.masked_fill(~upper, float("inf"))
+ flat = masked.reshape(batch, -1)
+ best = flat.argmin(-1)
+ best_value = flat.gather(-1, best[:, None]).squeeze(-1)
+ improvable = best_value < -1e-12
+ if not improvable.any():
+ break
+ p = best // size
+ q = best % size
+ rows = index[improvable]
+ permutations[rows] = permutations[rows].scatter(
+ 1,
+ torch.stack([p[improvable], q[improvable]], -1),
+ torch.stack(
+ [
+ permutations[rows, q[improvable]],
+ permutations[rows, p[improvable]],
+ ],
+ -1,
+ ),
+ )
+ final = text_fields[
+ index[:, None, None], permutations[:, :, None], permutations[:, None, :]
+ ]
+ return field_mse(final, visual_fields), permutations
+
+
+def load_aligned_views(args: argparse.Namespace) -> dict:
+ mappings = replay_view_permutations(args)
+ vision = torch.load(
+ f"{args.vg_dir}/vision_features.pt", map_location="cpu", weights_only=False
+ )
+ text = torch.load(
+ f"{args.vg_dir}/text_features.pt", map_location="cpu", weights_only=False
+ )
+ vision_index = {node: i for i, node in enumerate(vision["node_ids"])}
+ text_index = {node: i for i, node in enumerate(text["node_ids"])}
+ node_ids = sorted(mappings)
+ visual_views = []
+ text_views_aligned = []
+ text_views_raw = []
+ for node in node_ids:
+ mapping = mappings[node]
+ v = F.normalize(
+ vision["region_features"][vision_index[node]].double(), dim=-1
+ )
+ t = F.normalize(
+ text["region_features"][text_index[mapping["text_node_id"]]].double(),
+ dim=-1,
+ )
+ text_to_vision = torch.tensor(mapping["text_to_vision"])
+ vision_to_text = torch.empty_like(text_to_vision)
+ vision_to_text[text_to_vision] = torch.arange(len(text_to_vision))
+ visual_views.append(v)
+ text_views_aligned.append(t[vision_to_text])
+ text_views_raw.append(t)
+ return {
+ "node_ids": node_ids,
+ "visual_views": torch.stack(visual_views),
+ "text_views_aligned": torch.stack(text_views_aligned),
+ "text_views_raw": torch.stack(text_views_raw),
+ }
+
+
+def part_a(data: dict, args: argparse.Namespace) -> dict:
+ """Ordering gates at the true view assignment, true coarse frame."""
+ nodes = min(args.part_a_nodes, len(data["node_ids"]))
+ visual = data["visual_views"][:nodes]
+ text = data["text_views_aligned"][:nodes]
+ visual_fields = standardized_field(visual)
+ text_fields = standardized_field(text)
+ size = visual.shape[1]
+ upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1)
+
+ delta_e1 = batched_transposition_delta(text_fields, visual_fields)
+ e1_improving = (delta_e1[:, upper] < 0).double().mean(-1)
+
+ context = min(args.context_nodes, nodes)
+ visual_frame = F.normalize(visual[:context].mean(1), dim=-1)
+ text_frame = F.normalize(text[:context].mean(1), dim=-1)
+ profile_visual = visual @ visual_frame.T
+ profile_text = text @ text_frame.T
+
+ def centered_ranks(profiles: torch.Tensor) -> torch.Tensor:
+ ranks = profiles.argsort(-1).argsort(-1).double()
+ ranks = ranks - ranks.mean(-1, keepdim=True)
+ return F.normalize(ranks, dim=-1)
+
+ similarity = torch.bmm(
+ centered_ranks(profile_visual), centered_ranks(profile_text).transpose(-2, -1)
+ )
+ matched = similarity.diagonal(dim1=-2, dim2=-1)
+ delta_e2 = (
+ matched[:, :, None] + matched[:, None, :]
+ - similarity - similarity.transpose(-2, -1)
+ )
+ e2_improving = (delta_e2[:, upper] < 0).double().mean(-1)
+
+ report = {
+ "nodes": nodes,
+ "context_nodes": context,
+ "e1_within_node_relational": {
+ "improving_fraction_mean": float(e1_improving.mean()),
+ "strict_local_min_nodes": float((e1_improving == 0).double().mean()),
+ },
+ "e2_coarse_frame_profile": {
+ "improving_fraction_mean": float(e2_improving.mean()),
+ "strict_local_min_nodes": float((e2_improving == 0).double().mean()),
+ },
+ "joint": {},
+ }
+ for weight in (float(w) for w in args.joint_weights.split(",")):
+ delta_joint = delta_e1 + weight * delta_e2
+ joint_improving = (delta_joint[:, upper] < 0).double().mean(-1)
+ report["joint"][f"lambda_{weight}"] = {
+ "improving_fraction_mean": float(joint_improving.mean()),
+ "strict_local_min_nodes": float((joint_improving == 0).double().mean()),
+ }
+ return report
+
+
+def part_b_one_seed(
+ data: dict, args: argparse.Namespace, subset_seed: int
+) -> dict:
+ """Frame-free structured node affinity on one node subset."""
+ generator = torch.Generator().manual_seed(subset_seed)
+ subset = torch.randperm(len(data["node_ids"]), generator=generator)[
+ : args.part_b_nodes
+ ]
+ visual_fields = standardized_field(data["visual_views"][subset])
+ text_fields = standardized_field(data["text_views_raw"][subset])
+ n, size = visual_fields.shape[0], visual_fields.shape[-1]
+
+ columns_without_diag = torch.stack(
+ [
+ torch.tensor([j for j in range(size) if j != i], dtype=torch.long)
+ for i in range(size)
+ ]
+ )
+ row_index = torch.arange(size)[:, None]
+ signature_visual = visual_fields[
+ :, row_index, columns_without_diag
+ ].sort(-1).values
+ signature_text = text_fields[:, row_index, columns_without_diag].sort(-1).values
+ visual_norms = (signature_visual**2).sum(-1)
+ text_norms = (signature_text**2).sum(-1)
+
+ # Structure-scrambled null per text node: off-diagonal values are
+ # randomly reassigned to positions (marginals preserved, metric
+ # consistency destroyed). The excess of structured over null matching
+ # cost removes the free-permutation overfitting capacity that dominates
+ # raw min-cost at sixteen views.
+ null_generator = torch.Generator().manual_seed(subset_seed + 1)
+ upper = torch.triu(torch.ones(size, size, dtype=torch.bool), diagonal=1)
+ null_text_fields = torch.empty_like(text_fields)
+ for node in range(n):
+ values = text_fields[node][upper]
+ scrambled = values[
+ torch.randperm(len(values), generator=null_generator)
+ ]
+ field = torch.zeros(size, size, dtype=text_fields.dtype)
+ field[upper] = scrambled
+ null_text_fields[node] = field + field.T
+
+ def matched_costs(
+ target_fields: torch.Tensor,
+ ) -> tuple[torch.Tensor, torch.Tensor]:
+ """Descended min-fit cost and zero-capacity signature-Hungarian cost."""
+ target_signature = target_fields[:, row_index, columns_without_diag].sort(
+ -1
+ ).values
+ target_norms = (target_signature**2).sum(-1)
+ result = torch.zeros(n, n)
+ signature_result = torch.zeros(n, n)
+ for start in range(0, n, args.chunk):
+ stop = min(start + args.chunk, n)
+ block = signature_visual[start:stop]
+ init_cost = (
+ visual_norms[start:stop, None, :, None]
+ + target_norms[None, :, None, :]
+ - 2.0 * torch.einsum("cvs,nws->cnvw", block, target_signature)
+ )
+ chunk_pairs = init_cost.shape[0] * init_cost.shape[1]
+ permutations = torch.empty(chunk_pairs, size, dtype=torch.long)
+ flat_cost = init_cost.reshape(chunk_pairs, size, size)
+ for pair in range(chunk_pairs):
+ pair_rows, columns = linear_sum_assignment(flat_cost[pair].numpy())
+ permutations[pair] = torch.from_numpy(columns)
+ signature_result.view(-1)[
+ (start * n) + pair
+ ] = float(flat_cost[pair][pair_rows, columns].sum())
+ pair_text = (
+ target_fields[None, :, :, :]
+ .expand(stop - start, n, size, size)
+ .reshape(chunk_pairs, size, size)
+ )
+ pair_visual = (
+ visual_fields[start:stop, None, :, :]
+ .expand(stop - start, n, size, size)
+ .reshape(chunk_pairs, size, size)
+ )
+ final_cost, _ = batched_descent(
+ pair_text, pair_visual, permutations, args.descent_sweeps
+ )
+ result[start:stop] = final_cost.reshape(stop - start, n)
+ return result, signature_result
+
+ raw_cost, signature_cost = matched_costs(text_fields)
+ null_cost, _ = matched_costs(null_text_fields)
+ cost = raw_cost - null_cost
+
+ truth = torch.arange(n)
+ centered = cost - cost.mean(0, keepdim=True)
+ ranks = (centered <= centered.gather(1, truth[:, None])).sum(-1)
+ raw_ranks = (raw_cost <= raw_cost.gather(1, truth[:, None])).sum(-1)
+ rows, cols = linear_sum_assignment(cost.numpy())
+ true_cost = cost.diagonal()
+ random_costs = []
+ for _ in range(300):
+ permutation = torch.argsort(torch.rand(n, generator=generator))
+ random_costs.append(float(cost[truth, permutation].mean()))
+ random_costs = torch.tensor(random_costs)
+ sorted_cost = centered.sort(-1).values
+ margin = sorted_cost[:, 1] - sorted_cost[:, 0]
+ order = margin.argsort(descending=True)
+ top = order[: max(1, n // 10)]
+ mutual = centered.argmin(-1)[centered.argmin(0)] == torch.arange(n)
+ forward = centered.argmin(-1)
+ if args.matrices_output:
+ torch.save(
+ {
+ "subset_seed": subset_seed,
+ "raw_cost": raw_cost,
+ "null_cost": null_cost,
+ },
+ args.matrices_output.replace(".pt", f"_seed{subset_seed}.pt"),
+ )
+ return {
+ "subset_seed": subset_seed,
+ "nodes": n,
+ "retrieval_null_corrected_centered": {
+ "r@1": float((ranks <= 1).double().mean()),
+ "r@5": float((ranks <= 5).double().mean()),
+ "r@10": float((ranks <= 10).double().mean()),
+ "median_rank": float(ranks.double().median()),
+ "chance_r@1": 1.0 / n,
+ },
+ "retrieval_raw": {
+ "r@10": float((raw_ranks <= 10).double().mean()),
+ "median_rank": float(raw_ranks.double().median()),
+ },
+ "retrieval_signature_only": {
+ "r@10": float(
+ (
+ (
+ signature_cost
+ <= signature_cost.gather(1, truth[:, None])
+ ).sum(-1)
+ <= 10
+ )
+ .double()
+ .mean()
+ ),
+ "true_z": float(
+ (
+ signature_cost.mean()
+ - signature_cost.diagonal().mean()
+ )
+ / signature_cost.std().clamp_min(1e-12)
+ ),
+ },
+ "linear_assignment_energy_null_corrected": {
+ "true_mean_cost": float(true_cost.mean()),
+ "random_mean_cost": float(random_costs.mean()),
+ "random_std": float(random_costs.std()),
+ "true_z": float(
+ (random_costs.mean() - true_cost.mean())
+ / random_costs.std().clamp_min(1e-12)
+ ),
+ },
+ "hungarian_recovery_accuracy": float(
+ (torch.from_numpy(cols) == truth).double().mean()
+ ),
+ "top_margin_decile_precision": float((forward[top] == top).double().mean()),
+ "mutual_nn_count": int(mutual.sum()),
+ "mutual_nn_precision": float(
+ (forward[mutual] == torch.arange(n)[mutual]).double().mean()
+ )
+ if mutual.any()
+ else None,
+ }
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ data = load_aligned_views(args)
+ report: dict = {
+ "protocol": (
+ "Node states are sixteen-view sets with internal standardized "
+ "relation fields. Part A gates orderings at the true view "
+ "assignment given the true coarse frame (upper bound); part B "
+ "matches internal fields blind, with no frame and no pooled "
+ "vectors, and uses hidden pairs only for scoring."
+ ),
+ "part_a_view_ordering": part_a(data, args),
+ }
+ print(json.dumps({"part_a": report["part_a_view_ordering"]}))
+ report["part_b_structured_affinity"] = []
+ for subset_seed in (int(s) for s in args.part_b_seeds.split(",")):
+ result = part_b_one_seed(data, args, subset_seed)
+ report["part_b_structured_affinity"].append(result)
+ print(json.dumps({"part_b": result}))
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/worldalign/world_match.py b/worldalign/world_match.py
new file mode 100644
index 0000000..e719a9f
--- /dev/null
+++ b/worldalign/world_match.py
@@ -0,0 +1,138 @@
+"""World matching: spectral initialisation plus energy refinement.
+
+The two solvers are complementary. A spectral solver reads the coarse
+correspondence out of the eigenstructure of two relation fields in
+polynomial time and without any search, but its rounding is noisy. Exact
+steepest descent on the closed-form energy repairs the rounding but
+cannot find the basin from a random start. Composed, they recover the
+hidden assignment.
+
+Neither stage sees a pair. The relation fields are built independently
+per modality; the hidden permutation is applied to the text field and
+read back only to score.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+
+from .common import seed_everything, write_json
+from .spectral_match import grampa, spectral_profile, umeyama
+from .synth_fast_gate import ClosedFormEnergy, all_swaps, steepest_descent
+from .synth_triangle_gate import standardized
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--fields", required=True)
+ parser.add_argument("--eta", type=float, default=1.0)
+ parser.add_argument("--pair-weight", type=float, default=1.0)
+ parser.add_argument("--triangle-weight", type=float, default=1.0)
+ parser.add_argument("--refine-steps", type=int, default=4000)
+ parser.add_argument("--chunk", type=int, default=256)
+ parser.add_argument("--restarts", type=int, default=1)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="")
+ return parser.parse_args()
+
+
+def normalise(matrix: np.ndarray) -> np.ndarray:
+ size = len(matrix)
+ mask = ~np.eye(size, dtype=bool)
+ values = matrix[mask]
+ out = (matrix - values.mean()) / values.std()
+ np.fill_diagonal(out, 0.0)
+ return out
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ state = torch.load(args.fields, map_location="cpu", weights_only=False)
+ visual = normalise(state["visual_field"].double().numpy().copy())
+ text = normalise(state["text_field"].double().numpy().copy())
+ size = len(visual)
+
+ generator = np.random.default_rng(args.seed)
+ hidden = generator.permutation(size)
+ shuffled = text[np.ix_(hidden, hidden)]
+ truth = np.arange(size)
+
+ def accuracy(assignment: np.ndarray) -> float:
+ return float((hidden[assignment] == truth).mean())
+
+ device = torch.device(args.device)
+ text_tensor = standardized(torch.from_numpy(shuffled).to(device)).float()
+ visual_tensor = standardized(torch.from_numpy(visual).to(device)).float()
+ energy = ClosedFormEnergy(
+ text_tensor,
+ visual_tensor,
+ args.pair_weight,
+ args.triangle_weight,
+ args.chunk,
+ )
+ swaps = all_swaps(size, device)
+ truth_permutation = torch.from_numpy(np.argsort(hidden)).to(device)
+ true_energy = float(energy.energy(truth_permutation[None])[0])
+
+ report = {
+ "protocol": (
+ "Fields are built per modality without pairs; the hidden "
+ "permutation is applied to the text field and used only to "
+ "score. Spectral initialisation is followed by exact "
+ "steepest descent on the closed-form energy."
+ ),
+ "samples": size,
+ "field_correlation_at_truth": float(
+ np.corrcoef(
+ visual[~np.eye(size, dtype=bool)],
+ text[~np.eye(size, dtype=bool)],
+ )[0, 1]
+ ),
+ "true_energy": true_energy,
+ "chance": 1.0 / size,
+ "spectra": {
+ "visual": spectral_profile(visual),
+ "text": spectral_profile(shuffled),
+ },
+ "stages": {},
+ }
+
+ spectral = grampa(visual, shuffled, args.eta)
+ report["stages"]["spectral"] = {"accuracy": accuracy(spectral)}
+ initial = torch.from_numpy(spectral.copy()).to(device)
+ refined, value = steepest_descent(energy, initial, swaps, args.refine_steps)
+ report["stages"]["spectral_plus_refinement"] = {
+ "accuracy": accuracy(refined.cpu().numpy()),
+ "energy": value,
+ "energy_over_true": value / abs(true_energy) - true_energy / abs(true_energy),
+ }
+
+ quenches = []
+ for restart in range(args.restarts):
+ start = torch.from_numpy(generator.permutation(size)).to(device)
+ final, final_value = steepest_descent(
+ energy, start, swaps, args.refine_steps
+ )
+ quenches.append(
+ {
+ "accuracy": accuracy(final.cpu().numpy()),
+ "energy": final_value,
+ }
+ )
+ report["stages"]["random_start_refinement"] = quenches
+
+ print(json.dumps({key: value for key, value in report["stages"].items()}))
+ if args.output:
+ write_json(args.output, report)
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()