"""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()