diff options
Diffstat (limited to 'worldalign/content_projection.py')
| -rw-r--r-- | worldalign/content_projection.py | 214 |
1 files changed, 214 insertions, 0 deletions
diff --git a/worldalign/content_projection.py b/worldalign/content_projection.py new file mode 100644 index 0000000..b3e7422 --- /dev/null +++ b/worldalign/content_projection.py @@ -0,0 +1,214 @@ +"""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() |
