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