summaryrefslogtreecommitdiff
path: root/worldalign/view_gate.py
diff options
context:
space:
mode:
authorYuren Hao <blackhao0426@gmail.com>2026-08-01 14:10:03 -0500
committerYuren Hao <blackhao0426@gmail.com>2026-08-01 14:10:03 -0500
commita62cf4d2a99b4a7985c61b2a7feb92a82a8218b7 (patch)
treeee2248078db7edf3812a07f195afa3d9bd6f10c6 /worldalign/view_gate.py
World Alignment: unpaired cross-modal correspondence by relational identifiability
Method: scene states are sets of part states; relation fields are built within each modality and are invariant to how each side labels its own features; the cross-modal bridge is a coupling searched under an energy that is a closed-form functional of one matrix; solving is spectral initialisation followed by exact local refinement. Evidence: in a procedurally generated closed world, blind recovery of a hidden image-caption correspondence reaches 95.3% at 256 scenes against 0.39% chance, and the recovered pairs transfer to 200 held-out scenes at 93.0% exact retrieval with random-pair and shuffled-image controls at or near chance. Cross-modal value correspondence is derived from disjoint corpora rather than declared. On Visual Genome the field correlation reaches 0.656 against the 0.9 that polynomial recovery needs, with the deficit attributed away from segmentation and discretisation. Protocol: no image-text pair enters any objective, optimiser, initialisation, or model selection; hidden pairs score orderings only. Co-Authored-By: Claude <noreply@anthropic.com>
Diffstat (limited to 'worldalign/view_gate.py')
-rw-r--r--worldalign/view_gate.py436
1 files changed, 436 insertions, 0 deletions
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()