diff options
Diffstat (limited to 'worldalign/view_anchor_affinity.py')
| -rw-r--r-- | worldalign/view_anchor_affinity.py | 271 |
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() |
