summaryrefslogtreecommitdiff
path: root/worldalign/encoder_matrix.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/encoder_matrix.py')
-rw-r--r--worldalign/encoder_matrix.py223
1 files changed, 223 insertions, 0 deletions
diff --git a/worldalign/encoder_matrix.py b/worldalign/encoder_matrix.py
new file mode 100644
index 0000000..7a65fdb
--- /dev/null
+++ b/worldalign/encoder_matrix.py
@@ -0,0 +1,223 @@
+"""Do stronger unimodal encoders covary more? Re-priced with the right instrument.
+
+An earlier battery concluded that general-purpose frozen encoders underperform
+corpus-fitted statistics -- 0.19 against 0.656 -- and the project has proceeded
+on PPMI vectors ever since. That comparison was made in field correlation, and
+field correlation was shown today not to govern recovery: the bag-of-words text
+field scores 0.470 where the PPMI field scores 0.731, and yet its anchor bound
+is *higher*, 0.330 against 0.291. Any conclusion resting on the retired
+statistic has to be re-taken before it can be used to rule an option out.
+
+The option it currently rules out is the expensive one -- putting larger models
+behind each modality. Visual Genome's cross-modal ceiling with this encoder pair
+is 0.36 even when part correspondence is handed over, so the question is whether
+that is a property of the corpus or of DINOv2 crossed with PPMI. If a stronger
+text encoder moves the anchor bound, encoder scale is a live lever and worth
+real compute; if every one of them lands near 0.36, it is not, and no amount of
+scale will change it.
+
+Only unimodally trained encoders are admissible. A vision tower from a
+vision-language model is contrastively trained on image-text pairs and would
+import the very correspondence the project claims to search for.
+"""
+
+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.optimize import linear_sum_assignment
+
+from .common import seed_everything, write_json
+from .natural_families import load_phrases
+from .natural_pipeline import content_directions, word_vectors
+from .synth_cc_battery import moment_field
+
+# name -> (huggingface id, pooling). All trained on text alone.
+TEXT_ENCODERS = {
+ "MiniLM-L6 (22M)": ("sentence-transformers/all-MiniLM-L6-v2", "mean"),
+ "mpnet-base (110M)": ("sentence-transformers/all-mpnet-base-v2", "mean"),
+ "bert-base (110M)": ("bert-base-uncased", "mean"),
+ "BGE-large (335M)": ("BAAI/bge-large-en-v1.5", "cls"),
+ "Qwen2.5-0.5B": ("Qwen/Qwen2.5-0.5B", "mean"),
+ "Qwen2.5-1.5B": ("Qwen/Qwen2.5-1.5B", "mean"),
+}
+
+
+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/nobj_base_gpu.pt")
+ parser.add_argument("--samples", type=int, default=256)
+ parser.add_argument("--fit-scenes", type=int, default=900)
+ parser.add_argument("--word-vectors", type=int, default=128)
+ 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("--keep", type=int, default=64)
+ parser.add_argument("--shrinkage", type=float, default=0.05)
+ parser.add_argument("--weight-power", type=float, default=0.5)
+ parser.add_argument("--batch-size", type=int, default=64)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--only", nargs="*", default=None)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="artifacts/vg_5k/encoder_matrix.json")
+ return parser.parse_args()
+
+
+def standardise(matrix: np.ndarray) -> np.ndarray:
+ mask = ~np.eye(len(matrix), dtype=bool)
+ values = matrix[mask]
+ out = (matrix - values.mean()) / values.std()
+ np.fill_diagonal(out, 0.0)
+ return out
+
+
+def anchor_bound(visual: np.ndarray, text: np.ndarray, repeats: int = 5) -> float:
+ size = len(visual)
+ scores = []
+ for repeat in range(repeats):
+ generator = np.random.default_rng(repeat)
+ shuffle = generator.permutation(size)
+ half = size // 2
+ anchors, probe = shuffle[:half], shuffle[half:]
+ order = generator.permutation(len(probe))
+
+ def profile(field, rows):
+ block = field[np.ix_(rows, anchors)]
+ block = block - block.mean(1, keepdims=True)
+ return block / block.std(1, keepdims=True).clip(1e-9)
+
+ similarity = profile(visual, probe) @ profile(text, probe[order]).T / half
+ _, columns = linear_sum_assignment(-similarity)
+ scores.append(float((order[columns] == np.arange(len(probe))).mean()))
+ return float(np.mean(scores))
+
+
+@torch.inference_mode()
+def embed_phrases(phrases: list[str], model_id: str, pooling: str,
+ args: argparse.Namespace) -> np.ndarray:
+ from transformers import AutoModel, AutoTokenizer
+
+ tokenizer = AutoTokenizer.from_pretrained(model_id)
+ if tokenizer.pad_token is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ model = AutoModel.from_pretrained(model_id, torch_dtype=torch.float32).to(args.device)
+ model.eval()
+ out = []
+ for start in range(0, len(phrases), args.batch_size):
+ batch = phrases[start : start + args.batch_size]
+ tokens = tokenizer(batch, padding=True, truncation=True, max_length=32,
+ return_tensors="pt").to(args.device)
+ hidden = model(**tokens, return_dict=True).last_hidden_state
+ if pooling == "cls":
+ pooled = hidden[:, 0]
+ else:
+ mask = tokens["attention_mask"].unsqueeze(-1).float()
+ pooled = (hidden * mask).sum(1) / mask.sum(1).clamp_min(1e-9)
+ out.append(pooled.float().cpu().numpy())
+ del model
+ torch.cuda.empty_cache()
+ return np.concatenate(out).astype(np.float64)
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ vg_dir = Path(args.vg_dir)
+ truth = [json.loads(l) for l in
+ (vg_dir / "ground_truth.private.jsonl").read_text().splitlines() if l.strip()]
+ records = {json.loads(l)["node_id"]: json.loads(l) for l in
+ (vg_dir / "text_nodes.jsonl").read_text().splitlines() if l.strip()}
+ state = torch.load(args.objects, map_location="cpu", weights_only=False)
+ index = {n: i for i, n in enumerate(state["node_ids"])}
+ pairs = [p for p in truth
+ if p["vision_node_id"] in index and p["text_node_id"] in records]
+ needed = pairs[: args.samples + args.fit_scenes]
+
+ # vision side is held fixed throughout
+ vision_sets = {i: np.stack([s["feature"] for s in
+ state["segments"][index[p["vision_node_id"]]]]).astype(np.float64)
+ for i, p in enumerate(needed)
+ if len(state["segments"][index[p["vision_node_id"]]]) >= 2}
+ usable = sorted(vision_sets)
+ evaluated = [i for i in usable if i < args.samples][: args.samples]
+ fit = [i for i in usable if i >= args.samples]
+
+ def build(parts: dict[int, np.ndarray]) -> np.ndarray:
+ fitted = content_directions([parts[i] for i in fit if i in parts], args.shrinkage)
+ centre, basis, values = fitted
+ width = min(args.keep, basis.shape[1])
+ scale = values[:width] ** args.weight_power
+ sets = [F.normalize(torch.tensor((parts[i] - centre) @ basis[:, :width] * scale,
+ dtype=torch.float32), dim=-1)
+ for i in evaluated]
+ return standardise(moment_field(sets).double().numpy())
+
+ visual_field = build(vision_sets)
+ mask = ~np.eye(len(evaluated), dtype=bool)
+
+ # flatten every phrase once, remember which scene it belongs to
+ flat, owner = [], []
+ for i in usable:
+ for phrase in records[needed[i]["text_node_id"]]["region_closed"]:
+ flat.append(phrase)
+ owner.append(i)
+ owner = np.array(owner)
+ print(f"{len(evaluated)} evaluated scenes, {len(flat)} phrases", flush=True)
+
+ rows = []
+
+ def score(name: str, per_phrase: np.ndarray) -> None:
+ parts = {i: per_phrase[owner == i] for i in usable}
+ parts = {i: v for i, v in parts.items() if len(v) >= 2}
+ if not all(i in parts for i in evaluated):
+ print(f" {name:<22} skipped (missing scenes)", flush=True)
+ return
+ text_field = build(parts)
+ row = {
+ "encoder": name,
+ "dimension": int(per_phrase.shape[1]),
+ "correlation": float(np.corrcoef(visual_field[mask], text_field[mask])[0, 1]),
+ "anchor_bound": anchor_bound(visual_field, text_field),
+ }
+ rows.append(row)
+ print(f" {name:<22} dim={row['dimension']:<5} rho={row['correlation']:.3f} "
+ f"bound={row['anchor_bound']:.4f}", flush=True)
+
+ # the incumbent: corpus-fitted PPMI vectors averaged per phrase
+ vectors = word_vectors(load_phrases(vg_dir / "text_nodes.jsonl", "region_closed"), args)
+ ppmi = []
+ for phrase in flat:
+ hits = [vectors[t] for t in re.findall(r"[a-z]+", phrase.lower()) if t in vectors]
+ ppmi.append(np.mean(hits, axis=0) if hits else np.zeros(args.word_vectors))
+ score("PPMI (incumbent)", np.stack(ppmi))
+
+ chosen = args.only or list(TEXT_ENCODERS)
+ for name in chosen:
+ model_id, pooling = TEXT_ENCODERS[name]
+ try:
+ score(name, embed_phrases(flat, model_id, pooling, args))
+ except Exception as error:
+ print(f" {name:<22} FAILED: {str(error)[:90]}", flush=True)
+
+ write_json(args.output, {
+ "protocol": (
+ "Vision side held fixed (DINOv2-base segments). Only the text "
+ "encoder varies. All encoders are trained on text alone; no "
+ "vision-language model is admissible. Priced by the anchor bound, "
+ "since field correlation was shown not to govern recovery."
+ ),
+ "vg_part_oracle_ceiling": 0.359,
+ "rows": rows,
+ })
+ print(json.dumps({"done": True}))
+
+
+if __name__ == "__main__":
+ main()