"""Full assignment gate on content-projected states. The R6 battery collapsed the improving-swap fraction by one to two orders of magnitude. This runs the complete gate -- global ranking, exact transposition enumeration, descent from the truth, and counterfeit search from random starts -- on the projected states, which the fraction alone cannot decide. Projection hygiene: Flickr directions are fitted on the text-only training orbits and applied to held-out test states. VG directions are fitted only on nodes outside the evaluated subset. No pairs anywhere in the fit; hidden pairs score orderings only. """ from __future__ import annotations import argparse import json from pathlib import Path import torch import torch.nn.functional as F from .common import read_json, seed_everything, write_json from .content_projection import content_directions, project from .io import load_feature_pair, select_rows from .manifold_gate import standardize_relation from .ricci_control import run_gates def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--dataset", choices=["flickr", "vg"], default="flickr") parser.add_argument("--manifest", default="artifacts/manifest.json") parser.add_argument("--vision", default="artifacts/vision.pt") parser.add_argument("--text", default="artifacts/text.pt") parser.add_argument("--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt") parser.add_argument("--vg-vision", default="artifacts/vg_5k/vision_features.pt") parser.add_argument("--vg-text", default="artifacts/vg_5k/text_features.pt") parser.add_argument( "--vg-ground-truth", default="artifacts/vg_5k/ground_truth.private.jsonl" ) parser.add_argument("--split", choices=["val", "test"], default="test") parser.add_argument("--samples", type=int, default=512) parser.add_argument("--subset-seed", type=int, default=0) parser.add_argument("--dims", type=int, default=32) parser.add_argument("--shrinkage", type=float, default=0.05) parser.add_argument("--random-perms", type=int, default=1000) parser.add_argument("--descent-restarts", type=int, default=3) parser.add_argument("--descent-max-steps", type=int, default=200000) parser.add_argument( "--descent-objective", default="mse", choices=["mse", "m30_total"] ) parser.add_argument("--descent-verify-top", type=int, default=64) parser.add_argument("--seed", type=int, default=20260729) parser.add_argument("--output", required=True) return parser.parse_args() def flickr_states(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]: manifest = read_json(args.manifest) vision, _, vision_lookup, _ = load_feature_pair(args.vision, args.text) orbits = torch.load(args.text_orbits, map_location="cpu", weights_only=False) lookup = {int(row): i for i, row in enumerate(orbits["rows"])} features = F.normalize(orbits["features"].double(), dim=-1) train_views = features[[lookup[int(r)] for r in manifest["text_only_train"]]] mean, directions = content_directions(train_views, args.shrinkage) rows = manifest[args.split][: args.samples] text_states = project( features[[lookup[int(r)] for r in rows]].mean(1), mean, directions, args.dims ) visual_states = F.normalize( select_rows(vision["features"], vision_lookup, rows).double(), dim=-1 ) return text_states, visual_states def vg_states(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]: vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False) text = torch.load(args.vg_text, map_location="cpu", weights_only=False) vision_key = "context_states" if "context_states" in vision else "region_features" text_key = "context_states" if "context_states" in text else "region_features" pairs = [ json.loads(line) for line in open(args.vg_ground_truth, encoding="utf-8") if line.strip() ] vision_index = {node: i for i, node in enumerate(vision["node_ids"])} text_index = {node: i for i, node in enumerate(text["node_ids"])} vision_order = [vision_index[p["vision_node_id"]] for p in pairs] text_order = [text_index[p["text_node_id"]] for p in pairs] visual_views = F.normalize(vision[vision_key].double(), dim=-1)[vision_order] text_views = F.normalize(text[text_key].double(), dim=-1)[text_order] generator = torch.Generator().manual_seed(args.subset_seed) order = torch.randperm(len(visual_views), generator=generator) subset = order[: args.samples] holdout = order[args.samples :] visual_mean, visual_directions = content_directions( visual_views[holdout], args.shrinkage ) text_mean, text_directions = content_directions( text_views[holdout], args.shrinkage ) text_states = project( text_views[subset].mean(1), text_mean, text_directions, args.dims ) visual_states = project( visual_views[subset].mean(1), visual_mean, visual_directions, args.dims ) return text_states, visual_states def main() -> None: args = parse_args() seed_everything(args.seed) if args.dataset == "flickr": text_states, visual_states = flickr_states(args) else: text_states, visual_states = vg_states(args) text_channels = standardize_relation(text_states @ text_states.T)[0][None] visual_channels = standardize_relation(visual_states @ visual_states.T)[0][None] generator = torch.Generator().manual_seed(args.seed) report = { "protocol": ( "Content-projected states, directions fitted without pairs on " "held-out scenes; the complete assignment gate is scored with " "hidden pairs." ), "dataset": args.dataset, "dims": args.dims, "samples": args.samples, **run_gates( text_channels, visual_channels, args, generator ), } verdict = { "true_z_mse": report["gate_a"]["random"]["mse"]["true_z"], "improving_fraction": report["gate_b"]["improving_fraction"], "identity_strict_local_min": report["gate_b"]["identity_is_local_min_mse"], "descent_keeps": report["descent_from_true"]["final_accuracy"], "true_mse": report["gate_a"]["true"]["mse"], "best_random_descent": min( (r["final_objective"] for r in report["descent_from_random"]), default=None, ), "best_random_accuracy": max( (r["final_accuracy"] for r in report["descent_from_random"]), default=None, ), } verdict["counterfeit_found"] = bool( verdict["best_random_descent"] is not None and verdict["best_random_descent"] < verdict["true_mse"] and verdict["best_random_accuracy"] < 0.5 ) report["verdict"] = verdict print(json.dumps({"verdict": verdict})) Path(args.output).parent.mkdir(parents=True, exist_ok=True) write_json(args.output, report) print(f"Wrote {args.output}") if __name__ == "__main__": main()