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