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()
|