diff options
Diffstat (limited to 'worldalign/train_bridge.py')
| -rw-r--r-- | worldalign/train_bridge.py | 255 |
1 files changed, 255 insertions, 0 deletions
diff --git a/worldalign/train_bridge.py b/worldalign/train_bridge.py new file mode 100644 index 0000000..da29194 --- /dev/null +++ b/worldalign/train_bridge.py @@ -0,0 +1,255 @@ +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() |
