summaryrefslogtreecommitdiff
path: root/worldalign/train_bridge.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/train_bridge.py')
-rw-r--r--worldalign/train_bridge.py255
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()