"""Sweep the content projection, the largest measured lever on natural data. Content projection took the field correlation from 0.489 to 0.656, further than any other single choice, and it was adopted at one setting without a sweep. This prices its free parameters -- how many directions to keep, how hard to shrink the within-scene scatter -- and two variants of the projection itself. The whitened variant rescales each kept direction by its own between-scene spread, so a direction that separates scenes weakly is not drowned by one that separates them strongly. The kernel variant fits the same discriminant in a random Fourier feature space, testing whether the directions that separate scenes are linear in the encoder's coordinates at all. Word vectors and object states are computed once and reused, so the whole sweep costs one pipeline run. """ from __future__ import annotations import argparse import json import re from pathlib import Path import numpy as np import torch import torch.nn.functional as F from scipy.linalg import eigh from .common import seed_everything, write_json from .natural_families import load_phrases from .natural_pipeline import word_vectors from .synth_cc_battery import moment_field def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--vg-dir", default="artifacts/vg_5k") parser.add_argument("--objects", default="artifacts/vg_5k/natural_objects.pt") parser.add_argument("--samples", type=int, default=256) parser.add_argument("--fit-scenes", type=int, default=1200) parser.add_argument("--word-vectors", type=int, default=48) parser.add_argument("--min-count", type=int, default=60) parser.add_argument("--max-vocabulary", type=int, default=600) parser.add_argument("--context-window", type=int, default=2) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--rff", type=int, default=256, help="Random Fourier features for the kernel variant.") parser.add_argument("--output", default="artifacts/vg_5k/projection_sweep.json") return parser.parse_args() def discriminant( sets: list[np.ndarray], shrinkage: float ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Directions separating scenes more than parts, with their eigenvalues.""" means = np.stack([item.mean(0) for item in sets]) centre = means.mean(0) within = np.concatenate([item - item.mean(0, keepdims=True) for item in sets]) scatter_within = within.T @ within / max(len(within) - 1, 1) centred = means - centre scatter_between = centred.T @ centred / max(len(centred) - 1, 1) trace = np.trace(scatter_within) / len(scatter_within) values, vectors = eigh( scatter_between, scatter_within + shrinkage * trace * np.eye(len(scatter_within)) ) order = np.argsort(values)[::-1] return centre, vectors[:, order], np.maximum(values[order], 0.0) def lift(raw: np.ndarray, weights: np.ndarray, offsets: np.ndarray) -> np.ndarray: """Random Fourier features: an explicit map into an RBF feature space.""" return np.cos(raw @ weights + offsets) * np.sqrt(2.0 / weights.shape[1]) def correlation(vision_sets, text_sets) -> float: visual_field = moment_field(vision_sets) text_field = moment_field(text_sets) mask = ~np.eye(len(vision_sets), dtype=bool) return float( np.corrcoef( visual_field.double().numpy()[mask], text_field.double().numpy()[mask] )[0, 1] ) def main() -> None: args = parse_args() seed_everything(args.seed) vg_dir = Path(args.vg_dir) truth = [ json.loads(line) for line in (vg_dir / "ground_truth.private.jsonl") .read_text(encoding="utf-8") .splitlines() if line.strip() ] records = { json.loads(line)["node_id"]: json.loads(line) for line in (vg_dir / "text_nodes.jsonl").read_text(encoding="utf-8").splitlines() if line.strip() } state = torch.load(args.objects, map_location="cpu", weights_only=False) index = {node: position for position, node in enumerate(state["node_ids"])} vectors = word_vectors(load_phrases(vg_dir / "text_nodes.jsonl", "region_closed"), args) def language_state(node: str) -> np.ndarray | None: rows = [] for phrase in records[node]["region_closed"]: hits = [vectors[token] for token in re.findall(r"[a-z]+", phrase.lower()) if token in vectors] if hits: rows.append(np.mean(hits, axis=0)) return np.stack(rows) if rows else None def vision_state(node: str) -> np.ndarray | None: segments = state["segments"][index[node]] return ( np.stack([item["feature"] for item in segments]).astype(np.float64) if segments else None ) pairs = [ pair for pair in truth if pair["vision_node_id"] in index and pair["text_node_id"] in records ] fit = pairs[args.samples : args.samples + args.fit_scenes] language_fit = [language_state(p["text_node_id"]) for p in fit] language_fit = [i for i in language_fit if i is not None and len(i) > 1] vision_fit = [vision_state(p["vision_node_id"]) for p in fit] vision_fit = [i for i in vision_fit if i is not None and len(i) > 1] evaluated = [ pair for pair in pairs[: args.samples] if language_state(pair["text_node_id"]) is not None and vision_state(pair["vision_node_id"]) is not None ] language_eval = [language_state(p["text_node_id"]) for p in evaluated] vision_eval = [vision_state(p["vision_node_id"]) for p in evaluated] print(f"fit {len(vision_fit)} vision / {len(language_fit)} text, " f"evaluated {len(evaluated)}", flush=True) rows = [] def record(variant: str, keep, shrinkage, value: float) -> None: rows.append({"variant": variant, "keep": keep, "shrinkage": shrinkage, "correlation": value}) print(f"{variant:<10} keep={str(keep):<5} shrink={shrinkage:<6} rho={value:.4f}", flush=True) # --- linear discriminant: sweep kept width and shrinkage --- for shrinkage in (0.01, 0.05, 0.2, 1.0): lang = discriminant(language_fit, shrinkage) vis = discriminant(vision_fit, shrinkage) for keep in (4, 8, 16, 32, 64): def project(raw, fitted): centre, basis, _ = fitted width = min(keep, basis.shape[1]) return F.normalize( torch.tensor((raw - centre) @ basis[:, :width], dtype=torch.float32), dim=-1 ) value = correlation( [project(v, vis) for v in vision_eval], [project(t, lang) for t in language_eval], ) record("linear", keep, shrinkage, value) # --- eigenvalue-weighted: scale each direction by its own separation --- for shrinkage in (0.05, 0.2): lang = discriminant(language_fit, shrinkage) vis = discriminant(vision_fit, shrinkage) for keep in (8, 16, 32, 64): for power in (0.5, 1.0): def project(raw, fitted): centre, basis, values = fitted width = min(keep, basis.shape[1]) scale = values[:width] ** power return F.normalize( torch.tensor( (raw - centre) @ basis[:, :width] * scale, dtype=torch.float32 ), dim=-1 ) value = correlation( [project(v, vis) for v in vision_eval], [project(t, lang) for t in language_eval], ) record(f"weighted^{power}", keep, shrinkage, value) # --- kernel discriminant in a random Fourier feature space --- generator = np.random.default_rng(args.seed) def kernel_space(fit_sets, eval_sets, gamma_scale): stacked = np.concatenate(fit_sets) spread = np.median(np.linalg.norm(stacked - stacked.mean(0), axis=1)) gamma = gamma_scale / max(spread ** 2, 1e-9) weights = generator.normal( 0.0, np.sqrt(2 * gamma), size=(stacked.shape[1], args.rff) ) offsets = generator.uniform(0, 2 * np.pi, size=args.rff) return ( [lift(item, weights, offsets) for item in fit_sets], [lift(item, weights, offsets) for item in eval_sets], ) for gamma_scale in (0.25, 1.0): lang_fit_k, lang_eval_k = kernel_space(language_fit, language_eval, gamma_scale) vis_fit_k, vis_eval_k = kernel_space(vision_fit, vision_eval, gamma_scale) lang = discriminant(lang_fit_k, 0.05) vis = discriminant(vis_fit_k, 0.05) for keep in (8, 16, 32): def project(raw, fitted): centre, basis, _ = fitted width = min(keep, basis.shape[1]) return F.normalize( torch.tensor((raw - centre) @ basis[:, :width], dtype=torch.float32), dim=-1 ) value = correlation( [project(v, vis) for v in vis_eval_k], [project(t, lang) for t in lang_eval_k], ) record(f"kernel g={gamma_scale}", keep, 0.05, value) best = max(rows, key=lambda row: row["correlation"]) summary = { "protocol": ( "Projection fitted per modality on scenes outside the evaluated " "set; hidden pairs read only for the correlation. Baseline is " "linear, keep=8, shrinkage=0.05." ), "evaluated": len(evaluated), "rows": rows, "best": best, "recovery_threshold": 0.9, } print(json.dumps(best)) write_json(args.output, summary) if __name__ == "__main__": main()