summaryrefslogtreecommitdiff
path: root/worldalign/content_gate.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/content_gate.py')
-rw-r--r--worldalign/content_gate.py164
1 files changed, 164 insertions, 0 deletions
diff --git a/worldalign/content_gate.py b/worldalign/content_gate.py
new file mode 100644
index 0000000..4e1383a
--- /dev/null
+++ b/worldalign/content_gate.py
@@ -0,0 +1,164 @@
+"""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()