summaryrefslogtreecommitdiff
path: root/worldalign/precompute_gw.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/precompute_gw.py')
-rw-r--r--worldalign/precompute_gw.py51
1 files changed, 51 insertions, 0 deletions
diff --git a/worldalign/precompute_gw.py b/worldalign/precompute_gw.py
new file mode 100644
index 0000000..8eac056
--- /dev/null
+++ b/worldalign/precompute_gw.py
@@ -0,0 +1,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()