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