summaryrefslogtreecommitdiff
path: root/worldalign/view_anchor_affinity.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/view_anchor_affinity.py')
-rw-r--r--worldalign/view_anchor_affinity.py271
1 files changed, 271 insertions, 0 deletions
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()