summaryrefslogtreecommitdiff
path: root/worldalign/view_gate.py
diff options
context:
space:
mode:
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()