summaryrefslogtreecommitdiff
path: root/worldalign/diagnose.py
blob: 2d05fd9f21079974d3bd114b9983d2ee2fbbb2c7 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
from __future__ import annotations

import argparse
from pathlib import Path

import numpy as np
import torch
from scipy.stats import spearmanr

from .common import linear_cka, normalized, read_json, write_json
from .io import load_feature_pair, select_rows


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser()
    p.add_argument("--manifest", default="artifacts/manifest.json")
    p.add_argument("--vision", default="artifacts/vision.pt")
    p.add_argument("--text", default="artifacts/text.pt")
    p.add_argument("--split", choices=["val", "test"], default="val")
    p.add_argument("--max-samples", type=int, default=1_000)
    p.add_argument("--permutations", type=int, default=20)
    p.add_argument("--seed", type=int, default=20260728)
    p.add_argument("--output", default="artifacts/diagnostics.json")
    return p.parse_args()


def upper_triangle(x: torch.Tensor) -> np.ndarray:
    n = x.shape[0]
    i, j = torch.triu_indices(n, n, offset=1)
    return x[i, j].cpu().numpy()


def main() -> None:
    args = parse_args()
    manifest = read_json(args.manifest)
    vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
    rows = manifest[args.split][: args.max_samples]
    x = select_rows(vision["features"], vlookup, rows)
    y = select_rows(text["features"], tlookup, rows)

    gx = normalized(x) @ normalized(x).T
    gy = normalized(y) @ normalized(y).T
    gx_upper = upper_triangle(gx)
    gy_upper = upper_triangle(gy)
    rho = spearmanr(gx_upper, gy_upper).statistic
    generator = torch.Generator().manual_seed(args.seed)
    shuffled_rhos = []
    for _ in range(args.permutations):
        permutation = torch.randperm(len(y), generator=generator)
        shuffled_gy = gy[permutation][:, permutation]
        shuffled_rhos.append(
            float(spearmanr(gx_upper, upper_triangle(shuffled_gy)).statistic)
        )
    shuffled_mean = float(np.mean(shuffled_rhos))
    shuffled_std = float(np.std(shuffled_rhos))
    result = {
        "split": args.split,
        "samples": len(rows),
        "linear_cka": linear_cka(x, y),
        "pairwise_cosine_spearman": float(rho),
        "shuffled_spearman_mean": shuffled_mean,
        "shuffled_spearman_std": shuffled_std,
        "spearman_shuffle_z": float(
            (rho - shuffled_mean) / max(shuffled_std, 1e-12)
        ),
        "permutations": args.permutations,
        "vision_dim": x.shape[-1],
        "text_dim": y.shape[-1],
    }
    Path(args.output).parent.mkdir(parents=True, exist_ok=True)
    write_json(args.output, result)
    print(result)


if __name__ == "__main__":
    main()