diff options
Diffstat (limited to 'worldalign/encoder_matrix.py')
| -rw-r--r-- | worldalign/encoder_matrix.py | 223 |
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() |
