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