"""R6 battery: cross-view predictable (content) subspace projection. Views of the same scene share content and differ in style. Directions that maximize between-scene over within-scene variance are estimated from within-modality orbit structure alone (no pairs anywhere), then states are projected onto the top content directions. Hidden pairs are used only to score the cross-modal effect. Note the contrast with population whitening, which was destructive: the generalized eigenproblem whitens the within-scene (view-noise) covariance, not the total covariance, so directions where redescriptions of the same scene agree are amplified rather than equalized away. """ from __future__ import annotations import argparse import json import numpy as np import torch import torch.nn.functional as F from scipy.linalg import eigh from .common import read_json, write_json from .io import load_feature_pair, select_rows from .manifold_gate import all_transposition_delta_mse, standardize_relation def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() 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("--samples", type=int, default=512) parser.add_argument("--dims", default="8,16,32,64,128,256") parser.add_argument("--shrinkage", type=float, default=0.05) parser.add_argument( "--output", default="artifacts/manifold_gate/content_projection.json" ) return parser.parse_args() def content_directions( views: torch.Tensor, shrinkage: float ) -> tuple[np.ndarray, np.ndarray]: """Generalized eigenvectors of between-scene vs within-scene covariance. views: [scenes, views_per_scene, dim], any consistent preprocessing. Returns (mean, directions) with directions sorted by decreasing ratio. """ flat = views.reshape(-1, views.shape[-1]).double().numpy() mean = flat.mean(0) scene_means = views.double().mean(1).numpy() within = views.double().numpy() - scene_means[:, None, :] within = within.reshape(-1, views.shape[-1]) sigma_within = within.T @ within / max(len(within) - 1, 1) centered_means = scene_means - mean sigma_between = ( centered_means.T @ centered_means / max(len(centered_means) - 1, 1) ) trace_scale = np.trace(sigma_within) / len(sigma_within) regularized = sigma_within + shrinkage * trace_scale * np.eye(len(sigma_within)) values, vectors = eigh(sigma_between, regularized) order = np.argsort(values)[::-1] return mean, vectors[:, order] def project( states: torch.Tensor, mean: np.ndarray, directions: np.ndarray, dims: int ) -> torch.Tensor: basis = torch.from_numpy(directions[:, :dims]).double() centered = states.double() - torch.from_numpy(mean).double() return F.normalize(centered @ basis, dim=-1) def relation_spearman(text_states: torch.Tensor, visual_states: torch.Tensor) -> float: text_relation = text_states @ text_states.T visual_relation = visual_states @ visual_states.T mask = ~torch.eye(len(text_relation), dtype=torch.bool) t, v = text_relation[mask], visual_relation[mask] ranks = torch.stack( [t.argsort().argsort().double(), v.argsort().argsort().double()] ) return float(torch.corrcoef(ranks)[0, 1]) def improving_fraction( text_states: torch.Tensor, visual_states: torch.Tensor ) -> float: text_field, _, _ = standardize_relation(text_states @ text_states.T) visual_field, _, _ = standardize_relation(visual_states @ visual_states.T) delta = all_transposition_delta_mse(text_field, visual_field) upper = torch.triu(torch.ones_like(delta, dtype=torch.bool), diagonal=1) return float((delta[upper] < 0).double().mean()) def flickr_battery(args: argparse.Namespace, dims: list[int]) -> dict: 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_rows = [int(r) for r in manifest["text_only_train"]] test_rows = manifest["test"][: args.samples] train_views = features[[lookup[r] for r in train_rows]] mean, directions = content_directions(train_views, args.shrinkage) visual = select_rows(vision["features"], vision_lookup, test_rows).double() visual = F.normalize(visual, dim=-1) test_views = features[[lookup[int(r)] for r in test_rows]] orbit_mean = F.normalize(test_views.mean(1), dim=-1) single = test_views[:, 0] report = { "baseline_orbit_mean": { "spearman": relation_spearman(orbit_mean, visual), "improving_fraction": improving_fraction(orbit_mean, visual), }, "baseline_single": { "spearman": relation_spearman(F.normalize(single, dim=-1), visual), "improving_fraction": improving_fraction( F.normalize(single, dim=-1), visual ), }, "projected": {}, } for k in dims: projected_mean = project(test_views.mean(1), mean, directions, k) projected_single = project(single, mean, directions, k) report["projected"][k] = { "orbit_mean_spearman": relation_spearman(projected_mean, visual), "orbit_mean_improving_fraction": improving_fraction( projected_mean, visual ), "single_spearman": relation_spearman(projected_single, visual), } return report def vg_battery(args: argparse.Namespace, dims: list[int]) -> dict: vision = torch.load(args.vg_vision, map_location="cpu", weights_only=False) text = torch.load(args.vg_text, map_location="cpu", weights_only=False) 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["region_features"].double(), dim=-1)[ vision_order ] text_views = F.normalize(text["region_features"].double(), dim=-1)[text_order] visual_mean, visual_directions = content_directions(visual_views, args.shrinkage) text_mean, text_directions = content_directions(text_views, args.shrinkage) generator = torch.Generator().manual_seed(0) subset = torch.randperm(len(visual_views), generator=generator)[: args.samples] visual_node = F.normalize(visual_views[subset].mean(1), dim=-1) text_node = F.normalize(text_views[subset].mean(1), dim=-1) report = { "baseline": { "spearman": relation_spearman(text_node, visual_node), "improving_fraction": improving_fraction(text_node, visual_node), }, "projected": {}, } for k in dims: projected_text = project( text_views[subset].mean(1), text_mean, text_directions, k ) projected_visual = project( visual_views[subset].mean(1), visual_mean, visual_directions, k ) report["projected"][k] = { "both_sides_spearman": relation_spearman(projected_text, projected_visual), "both_sides_improving_fraction": improving_fraction( projected_text, projected_visual ), "text_only_spearman": relation_spearman(projected_text, visual_node), } return report def main() -> None: args = parse_args() dims = [int(d) for d in args.dims.split(",")] report = { "protocol": ( "Content directions maximize between-scene over within-scene " "variance of view states, fitted per modality on unpaired " "orbit structure only. Hidden pairs score the effect." ), "flickr": flickr_battery(args, dims), "vg_region_closed": vg_battery(args, dims), } write_json(args.output, report) print(json.dumps(report, indent=2)) if __name__ == "__main__": main()