summaryrefslogtreecommitdiff
path: root/worldalign/precompute_gw.py
blob: 8eac0563dd436b17253451dea38c4f1d813a01b5 (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
from __future__ import annotations

import argparse
from pathlib import Path

import torch

from .common import read_json, seed_everything
from .gw import gw_pseudo_targets
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("--output", default="artifacts/gw.pt")
    p.add_argument("--clusters", type=int, default=128)
    p.add_argument("--seed", type=int, default=20260728)
    p.add_argument("--max-iter", type=int, default=100)
    return p.parse_args()


def main() -> None:
    args = parse_args()
    seed_everything(args.seed)
    manifest = read_json(args.manifest)
    vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
    x = select_rows(
        vision["features"], vlookup, manifest["vision_only_train"]
    )
    y = select_rows(text["features"], tlookup, manifest["text_only_train"])
    result = gw_pseudo_targets(
        x,
        y,
        clusters=args.clusters,
        seed=args.seed,
        max_iter=args.max_iter,
    )
    Path(args.output).parent.mkdir(parents=True, exist_ok=True)
    torch.save(result, args.output)
    print(
        f"Wrote {args.output}: K={result['clusters']}, "
        f"GW={result['gw_distance']:.6f}, "
        f"row_entropy={result['coupling_row_entropy']:.6f}"
    )


if __name__ == "__main__":
    main()