summaryrefslogtreecommitdiff
path: root/worldalign/train_bridge.py
diff options
context:
space:
mode:
authorYuren Hao <blackhao0426@gmail.com>2026-08-01 14:10:03 -0500
committerYuren Hao <blackhao0426@gmail.com>2026-08-01 14:10:03 -0500
commita62cf4d2a99b4a7985c61b2a7feb92a82a8218b7 (patch)
treeee2248078db7edf3812a07f195afa3d9bd6f10c6 /worldalign/train_bridge.py
World Alignment: unpaired cross-modal correspondence by relational identifiability
Method: scene states are sets of part states; relation fields are built within each modality and are invariant to how each side labels its own features; the cross-modal bridge is a coupling searched under an energy that is a closed-form functional of one matrix; solving is spectral initialisation followed by exact local refinement. Evidence: in a procedurally generated closed world, blind recovery of a hidden image-caption correspondence reaches 95.3% at 256 scenes against 0.39% chance, and the recovered pairs transfer to 200 held-out scenes at 93.0% exact retrieval with random-pair and shuffled-image controls at or near chance. Cross-modal value correspondence is derived from disjoint corpora rather than declared. On Visual Genome the field correlation reaches 0.656 against the 0.9 that polynomial recovery needs, with the deficit attributed away from segmentation and discretisation. Protocol: no image-text pair enters any objective, optimiser, initialisation, or model selection; hidden pairs score orderings only. Co-Authored-By: Claude <noreply@anthropic.com>
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()