summaryrefslogtreecommitdiff
path: root/worldalign/vg_diagnose.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/vg_diagnose.py')
-rw-r--r--worldalign/vg_diagnose.py200
1 files changed, 200 insertions, 0 deletions
diff --git a/worldalign/vg_diagnose.py b/worldalign/vg_diagnose.py
new file mode 100644
index 0000000..4c6d4a8
--- /dev/null
+++ b/worldalign/vg_diagnose.py
@@ -0,0 +1,200 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+from scipy.stats import spearmanr
+import torch
+
+from .common import normalized, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--vision", default="artifacts/vg/vision_features.pt")
+ p.add_argument("--text", default="artifacts/vg/text_features.pt")
+ p.add_argument(
+ "--truth", default="artifacts/vg/ground_truth.private.jsonl"
+ )
+ p.add_argument("--max-nodes", type=int, default=5_000)
+ p.add_argument("--quantiles", type=int, default=33)
+ p.add_argument("--output", default="artifacts/vg/diagnostics.json")
+ return p.parse_args()
+
+
+def read_jsonl(path: str) -> list[dict]:
+ with open(path, encoding="utf-8") as handle:
+ return [json.loads(line) for line in handle if line.strip()]
+
+
+def bundle_signature(
+ features: torch.Tensor, quantiles: int = 33
+) -> torch.Tensor:
+ """Permutation- and rotation-invariant signature of a bag of views."""
+ x = normalized(features.float())
+ gram = x @ x.transpose(1, 2)
+ views = x.shape[1]
+ i, j = torch.triu_indices(views, views, offset=1)
+ pairwise = gram[:, i, j]
+ q = torch.linspace(0, 1, quantiles)
+ distribution = torch.quantile(pairwise, q, dim=1).T
+ eigenvalues = torch.linalg.eigvalsh(gram).flip(-1) / max(views, 1)
+ return torch.cat([distribution, eigenvalues], dim=-1)
+
+
+def rank_normalize_columns(x: torch.Tensor) -> torch.Tensor:
+ order = torch.argsort(x, dim=0)
+ ranks = torch.argsort(order, dim=0).float()
+ ranks = ranks / max(x.shape[0] - 1, 1)
+ return (ranks - 0.5) * 2
+
+
+def aligned_text_indices(
+ vision_ids: list[str], text_ids: list[str], truth_path: str
+) -> torch.Tensor:
+ truth = read_jsonl(truth_path)
+ mapping = {
+ item["vision_node_id"]: item["text_node_id"] for item in truth
+ }
+ text_lookup = {node_id: idx for idx, node_id in enumerate(text_ids)}
+ return torch.tensor([text_lookup[mapping[node_id]] for node_id in vision_ids])
+
+
+def arbitrary_target_retrieval(
+ queries: torch.Tensor, candidates: torch.Tensor, targets: torch.Tensor
+) -> dict:
+ similarity = normalized(queries) @ normalized(candidates).T
+ target_score = similarity[
+ torch.arange(len(queries)), targets.to(similarity.device)
+ ]
+ ranks = (similarity > target_score[:, None]).sum(-1) + 1
+ top_values, top_indices = similarity.topk(
+ min(2, similarity.shape[1]), dim=1
+ )
+ prediction = top_indices[:, 0]
+ correct = prediction == targets.to(prediction.device)
+ if top_values.shape[1] == 2:
+ margin = top_values[:, 0] - top_values[:, 1]
+ else:
+ margin = top_values[:, 0]
+ confidence_order = torch.argsort(margin, descending=True)
+ confidence_precision = {}
+ for count in (10, 50, 100, 500, 1_000):
+ if count <= len(queries):
+ selected = confidence_order[:count]
+ confidence_precision[str(count)] = {
+ "correct": int(correct[selected].sum()),
+ "precision": float(correct[selected].float().mean()),
+ }
+
+ reverse_prediction = similarity.argmax(dim=0)
+ mutual = (
+ reverse_prediction[prediction]
+ == torch.arange(len(queries), device=prediction.device)
+ )
+ mutual_count = int(mutual.sum())
+ result = {
+ "r@1": float((ranks <= 1).float().mean()),
+ "r@5": float((ranks <= 5).float().mean()),
+ "r@10": float((ranks <= 10).float().mean()),
+ "mean_reciprocal_rank": float((1.0 / ranks.float()).mean()),
+ "median_rank": float(ranks.float().median()),
+ "chance_r@1": 1.0 / len(candidates),
+ "chance_r@5": min(5.0 / len(candidates), 1.0),
+ "chance_r@10": min(10.0 / len(candidates), 1.0),
+ "confidence_margin_precision": confidence_precision,
+ "mutual_nearest": {
+ "selected": mutual_count,
+ "correct": int(correct[mutual].sum()),
+ "precision": (
+ float(correct[mutual].float().mean())
+ if mutual_count
+ else 0.0
+ ),
+ },
+ }
+ return result
+
+
+def upper_triangle(x: torch.Tensor) -> np.ndarray:
+ i, j = torch.triu_indices(len(x), len(x), offset=1)
+ return x[i, j].cpu().numpy()
+
+
+def main() -> None:
+ args = parse_args()
+ vision = torch.load(args.vision, map_location="cpu", weights_only=False)
+ text = torch.load(args.text, map_location="cpu", weights_only=False)
+ n = min(args.max_nodes, len(vision["node_ids"]), len(text["node_ids"]))
+ vision_ids = vision["node_ids"][:n]
+ target_full = aligned_text_indices(
+ vision_ids, text["node_ids"], args.truth
+ )
+ candidate_indices = torch.unique(target_full, sorted=False)
+ if len(candidate_indices) != n:
+ raise ValueError("Ground truth is not a one-to-one permutation")
+ candidate_lookup = {
+ int(old): new for new, old in enumerate(candidate_indices.tolist())
+ }
+ targets = torch.tensor(
+ [candidate_lookup[int(old)] for old in target_full.tolist()]
+ )
+
+ v_views = vision["region_features"][:n]
+ t_views = text["region_features"][candidate_indices]
+ v_signature = rank_normalize_columns(
+ bundle_signature(v_views, args.quantiles)
+ )
+ t_signature = rank_normalize_columns(
+ bundle_signature(t_views, args.quantiles)
+ )
+ signature_retrieval = arbitrary_target_retrieval(
+ v_signature, t_signature, targets
+ )
+
+ paired_t_signature = t_signature[targets]
+ signature_cosine = (
+ normalized(v_signature) * normalized(paired_t_signature)
+ ).sum(-1)
+ generator = torch.Generator().manual_seed(20260728)
+ shuffled = paired_t_signature[torch.randperm(n, generator=generator)]
+ shuffled_cosine = (
+ normalized(v_signature) * normalized(shuffled)
+ ).sum(-1)
+
+ v_scene = vision["global_features"][:n]
+ t_scene = text["region_features"][candidate_indices].mean(1)[targets]
+ gv = normalized(v_scene) @ normalized(v_scene).T
+ gt = normalized(t_scene) @ normalized(t_scene).T
+ scene_rho = spearmanr(
+ upper_triangle(gv), upper_triangle(gt)
+ ).statistic
+
+ result = {
+ "nodes": n,
+ "views_per_node": int(v_views.shape[1]),
+ "vision_model": vision["model"],
+ "text_model": text["model"],
+ "text_tier": text["tier"],
+ "bundle_signature_retrieval": signature_retrieval,
+ "paired_bundle_signature_cosine_mean": float(
+ signature_cosine.mean()
+ ),
+ "shuffled_bundle_signature_cosine_mean": float(
+ shuffled_cosine.mean()
+ ),
+ "between_scene_pairwise_cosine_spearman": float(scene_rho),
+ "evaluation_note": (
+ "Ground-truth permutation is used only to score signatures and "
+ "relation geometry, never to fit them."
+ ),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2))
+
+
+if __name__ == "__main__":
+ main()