from __future__ import annotations import argparse from pathlib import Path import numpy as np import torch from torch.optim import AdamW from tqdm import tqdm from scipy.optimize import linear_sum_assignment from .common import ( cosine_isometry_loss, cosine_loss, cosine_schedule, parameter_count, read_json, retrieval_metrics, seed_everything, sliced_wasserstein, ) from .gw import gw_pseudo_targets from .io import load_feature_pair, select_rows from .models import Bridge def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser() p.add_argument("--mode", choices=["paired", "unpaired_swd", "unpaired_gw"]) 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", required=True) p.add_argument("--device", default="cuda:1") p.add_argument("--steps", type=int, default=4_000) p.add_argument("--batch-size", type=int, default=256) p.add_argument("--hidden-dim", type=int, default=1536) p.add_argument("--linear", action="store_true") p.add_argument("--lr", type=float, default=3e-4) p.add_argument("--warmup", type=int, default=200) p.add_argument("--clusters", type=int, default=128) p.add_argument( "--gw-cache", help="Optional path for reusable GW prototype coupling and assignments.", ) p.add_argument("--swd-weight", type=float, default=10.0) p.add_argument("--isometry-weight", type=float, default=1.0) p.add_argument("--gw-weight", type=float, default=1.0) p.add_argument("--seed", type=int, default=20260728) p.add_argument("--eval-every", type=int, default=200) return p.parse_args() def evaluate( bridge: Bridge, x: torch.Tensor, y: torch.Tensor, device: str, ) -> dict[str, float]: bridge.eval() mapped = [] with torch.inference_mode(): for chunk in x.split(512): mapped.append(bridge(chunk.to(device)).cpu()) bridge.train() return retrieval_metrics(torch.cat(mapped), y) def evaluate_gw_cluster_mapping( gw: dict, val_x: torch.Tensor, val_y: torch.Tensor, ) -> dict[str, float]: vx = torch.nn.functional.normalize(gw["vision_centers"], dim=-1) ty = torch.nn.functional.normalize(gw["text_centers"], dim=-1) x_labels = (torch.nn.functional.normalize(val_x, dim=-1) @ vx.T).argmax(1) y_labels = (torch.nn.functional.normalize(val_y, dim=-1) @ ty.T).argmax(1) predicted_map = gw["coupling"].argmax(1) predicted = predicted_map[x_labels] accuracy = (predicted == y_labels).float().mean().item() k = vx.shape[0] contingency = torch.zeros(k, k, dtype=torch.float64) for i, j in zip(x_labels.tolist(), y_labels.tolist()): contingency[i, j] += 1 row, col = linear_sum_assignment(-contingency.numpy()) oracle_correct = contingency[row, col].sum().item() return { "paired_val_cluster_accuracy": float(accuracy), "paired_val_cluster_chance": 1.0 / k, "paired_val_cluster_oracle_permutation": float( oracle_correct / max(len(val_x), 1) ), } 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) if args.mode == "paired": rows = manifest["paired_train"] train_x = select_rows(vision["features"], vlookup, rows) train_y = select_rows(text["features"], tlookup, rows) target_by_cluster = None assignments = None gw_meta = {} gw_result = None else: train_x = select_rows( vision["features"], vlookup, manifest["vision_only_train"] ) train_y = select_rows( text["features"], tlookup, manifest["text_only_train"] ) target_by_cluster = None assignments = None gw_meta = {} gw_result = None if args.mode == "unpaired_gw": if args.gw_cache and Path(args.gw_cache).exists(): print(f"Loading GW coupling from {args.gw_cache}") gw = torch.load( args.gw_cache, map_location="cpu", weights_only=False ) else: print("Computing unpaired GW prototype coupling...") gw = gw_pseudo_targets( train_x, train_y, args.clusters, args.seed ) if args.gw_cache: Path(args.gw_cache).parent.mkdir(parents=True, exist_ok=True) torch.save(gw, args.gw_cache) print(f"Wrote GW coupling to {args.gw_cache}") gw_result = gw target_by_cluster = gw["target_centers"] assignments = gw["vision_assignments"] gw_meta = { "gw_distance": gw["gw_distance"], "coupling_row_entropy": gw["coupling_row_entropy"], "clusters": gw["clusters"], } print(f"GW diagnostics: {gw_meta}") val_rows = manifest["val"] val_x = select_rows(vision["features"], vlookup, val_rows) val_y = select_rows(text["features"], tlookup, val_rows) if gw_result is not None: gw_meta.update(evaluate_gw_cluster_mapping(gw_result, val_x, val_y)) print(f"GW paired-eval diagnostics (not used for training): {gw_meta}") bridge = Bridge( train_x.shape[-1], train_y.shape[-1], hidden_dim=args.hidden_dim, linear=args.linear, ).to(args.device) print(f"Bridge parameters: {parameter_count(bridge):,}") optimizer = AdamW(bridge.parameters(), lr=args.lr, weight_decay=1e-4) generator = torch.Generator().manual_seed(args.seed) history = [] best_score = -1.0 best_state = None progress = tqdm(range(args.steps), desc=f"bridge:{args.mode}") for step in progress: ix = torch.randint( len(train_x), (args.batch_size,), generator=generator ) if args.mode == "paired": iy = ix else: iy = torch.randint( len(train_y), (args.batch_size,), generator=generator ) x = train_x[ix].to(args.device) y = train_y[iy].to(args.device) mapped = bridge(x) if args.mode == "paired": alignment = cosine_loss(mapped, y) swd = mapped.new_zeros(()) gw_loss = mapped.new_zeros(()) else: alignment = mapped.new_zeros(()) swd = sliced_wasserstein(mapped, y, num_projections=64) if target_by_cluster is not None and assignments is not None: pseudo = target_by_cluster[assignments[ix]].to(args.device) gw_loss = cosine_loss(mapped, pseudo) else: gw_loss = mapped.new_zeros(()) isometry = cosine_isometry_loss(x, mapped) loss = ( alignment + args.swd_weight * swd + args.gw_weight * gw_loss + args.isometry_weight * isometry ) optimizer.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(bridge.parameters(), 1.0) optimizer.step() scale = cosine_schedule(step, args.steps, args.warmup) for group in optimizer.param_groups: group["lr"] = args.lr * scale if step % 20 == 0: progress.set_postfix( loss=f"{loss.item():.3f}", align=f"{alignment.item():.3f}", swd=f"{swd.item():.3g}", gw=f"{gw_loss.item():.3f}", iso=f"{isometry.item():.3f}", ) if step % args.eval_every == 0 or step == args.steps - 1: metrics = evaluate(bridge, val_x, val_y, args.device) record = {"step": step, "loss": float(loss.item()), **metrics} history.append(record) # Paired validation is recorded for scientific evaluation, not used to # select the unsupervised checkpoint. Save final for unpaired modes. if args.mode == "paired" and metrics["i2t_r@1"] > best_score: best_score = metrics["i2t_r@1"] best_state = { k: v.detach().cpu().clone() for k, v in bridge.state_dict().items() } print(record) if args.mode == "paired" and best_state is not None: bridge.load_state_dict(best_state) checkpoint_selection = "best paired validation R@1 (upper bound only)" else: checkpoint_selection = "final step; no paired validation selection" state = { "config": bridge.config(), "state_dict": bridge.state_dict(), "mode": args.mode, "args": vars(args), "vision_model": vision["model"], "text_model": text["model"], "history": history, "gw": gw_meta, "checkpoint_selection": checkpoint_selection, } Path(args.output).parent.mkdir(parents=True, exist_ok=True) torch.save(state, args.output) print(f"Wrote {args.output}") if __name__ == "__main__": main()