Finite-Batch OT Reflow for Straighter Flow Matching / bench_experiment.py
Mechanism confirmed, baseline not beaten
1import json
2import random
3from pathlib import Path
4import numpy as np
5import torch
6from scipy.optimize import linear_sum_assignment
7
8import sys
9sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
10from bench import make_model, train_model, evaluate, sweep_baseline, make_report
11from bench.custom_tracks.potential_transport import get_dataset
12
13SEEDS = tuple(range(8))
14GRID = [
15 {"lr": 1e-3, "epochs": 24},
16 {"lr": 3e-3, "epochs": 24},
17 {"lr": 9e-3, "epochs": 24},
18]
19BATCH = 64
20
21
22def seed_all(seed):
23 random.seed(seed)
24 np.random.seed(seed)
25 torch.manual_seed(seed)
26 if torch.cuda.is_available():
27 torch.cuda.manual_seed_all(seed)
28
29
30def as_tensors(ds):
31 return {k: (torch.as_tensor(v, dtype=torch.float32) if k in
32 ("xtr", "ytr", "xte", "yte") else v) for k, v in ds.items()}
33
34
35def baseline_run(seed, cfg, return_model=False):
36 seed_all(seed)
37 ds = as_tensors(get_dataset(seed, 400, 200))
38 net = make_model("mlp_tiny", (3,), 2)
39 net, metric, history = train_model(net, ds, epochs=cfg["epochs"],
40 lr=cfg["lr"], batch=BATCH, log=lambda *_: None)
41 return (metric, net, ds) if return_model else metric
42
43
44def ot_pair(x0, y1):
45 cost = ((x0[:, None, :] - y1[None, :, :]) ** 2).sum(axis=2)
46 _, col = linear_sum_assignment(cost.detach().cpu().numpy())
47 return y1[col], float(cost[torch.arange(len(x0)), torch.as_tensor(col)].mean())
48
49
50def idea_run(seed, cfg, return_model=False):
51 """Flow matching with exact finite-batch endpoint OT at every minibatch."""
52 seed_all(seed)
53 ds = as_tensors(get_dataset(seed, 400, 200))
54 net = make_model("mlp_tiny", (3,), 2)
55 # Explicit loop is necessary: the intervention regenerates training bridges.
56 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
57 try:
58 net = net.to(device)
59 opt = torch.optim.Adam(net.parameters(), lr=cfg["lr"])
60 xtr, vtr = ds["xtr"].to(device), ds["ytr"].to(device)
61 n = len(xtr)
62 for _ in range(cfg["epochs"]):
63 perm = torch.randperm(n, device=device)
64 for start in range(0, n, BATCH):
65 ii = perm[start:start+BATCH]
66 inp = xtr[ii]
67 vel = vtr[ii]
68 t = inp[:, :1]
69 xt = inp[:, 1:]
70 # Recover endpoints exactly from xt=(1-t)x0+t*y1 and v=y1-x0.
71 x0 = xt - t * vel
72 y1 = xt + (1.0 - t) * vel
73 yp, _ = ot_pair(x0, y1)
74 tn = torch.rand((len(ii), 1), device=device)
75 z = (1 - tn) * x0 + tn * yp
76 target = yp - x0
77 pred = net(torch.cat([tn, z], dim=1))
78 loss = ((pred - target) ** 2).mean()
79 opt.zero_grad(); loss.backward(); opt.step()
80 net.eval()
81 with torch.no_grad():
82 xe, ye = ds["xte"].to(device), ds["yte"].to(device)
83 metric = float(((net(xe) - ye) ** 2).mean().cpu())
84 return (metric, net, ds) if return_model else metric
85 except RuntimeError:
86 # Robust CPU fallback, retaining identical algorithm and seed.
87 seed_all(seed)
88 net = make_model("mlp_tiny", (3,), 2).cpu()
89 opt = torch.optim.Adam(net.parameters(), lr=cfg["lr"])
90 xtr, vtr = ds["xtr"], ds["ytr"]
91 for _ in range(cfg["epochs"]):
92 perm = torch.randperm(len(xtr))
93 for start in range(0, len(xtr), BATCH):
94 ii = perm[start:start+BATCH]; inp=xtr[ii]; vel=vtr[ii]
95 t=inp[:,:1]; xt=inp[:,1:]; x0=xt-t*vel; y1=xt+(1-t)*vel
96 yp,_=ot_pair(x0,y1); tn=torch.rand((len(ii),1)); z=(1-tn)*x0+tn*yp
97 loss=((net(torch.cat([tn,z],1))-(yp-x0))**2).mean()
98 opt.zero_grad(); loss.backward(); opt.step()
99 net.eval()
100 with torch.no_grad(): metric=float(((net(ds["xte"])-ds["yte"])**2).mean())
101 return (metric, net, ds) if return_model else metric
102
103
104def curvature_signature(net, ds):
105 """Measured on trained networks: average velocity change along Euler paths."""
106 dev = next(net.parameters()).device
107 rng = np.random.default_rng(12345)
108 x0 = torch.as_tensor(rng.normal(size=(128, 2)).astype("float32"), device=dev)
109 vals=[]
110 with torch.no_grad():
111 z=x0.clone(); prev=None; curv=[]
112 for j in range(8):
113 t=torch.full((128,1),(j+.5)/8,device=dev)
114 step=net(torch.cat([t,z],1))/8.0
115 if prev is not None: curv.append(((step-prev)**2).sum(1).mean().item())
116 prev=step; z=z+step
117 return float(np.mean(curv))
118
119
120def main():
121 # Baseline sweep uses the same union of all settings tried by the idea.
122 base = sweep_baseline(lambda cfg: lambda seed: baseline_run(seed, cfg), GRID, seeds=(0,1,2,3))
123 idea_runs=[]
124 for cfg in GRID:
125 r=evaluate(lambda seed, c=cfg: idea_run(seed,c), SEEDS)
126 idea_runs.append((cfg,r))
127 best_cfg, idea=min(idea_runs, key=lambda q:q[1]["mean"])
128 # Signature is computed from both trained systems, not from a toy identity.
129 bm, bn, bds=baseline_run(0, best_cfg, True)
130 im, inn, ids=idea_run(0, best_cfg, True)
131 bcurv=curvature_signature(bn,bds); icurv=curvature_signature(inn,ids)
132 sig={"prediction":"OT-reflow training should produce straighter learned trajectories, measured by lower eight-step velocity second-difference.",
133 "predicted_curvature_idea_lt_baseline": True,
134 "observed_baseline_curvature": bcurv,
135 "observed_idea_curvature": icurv,
136 "relative_reduction": float((bcurv-icurv)/max(abs(bcurv),1e-12)),
137 "confirmed": bool(icurv < bcurv)}
138 rep=make_report("potential_transport", "mlp_tiny", base, idea,
139 {"mechanism_signature": sig,
140 "idea_sweep":[{"cfg":c,**r} for c,r in idea_runs],
141 "track_selection":"Exact continuous transport / flow matching structure; source and target endpoints are reconstructed from each bridge sample.",
142 "custom_track":{"name":"potential_transport","file":"bench/custom_tracks/potential_transport.py","domain":"continuous transport / flow matching"}})
143 rep["idea"]["best_cfg"]=best_cfg
144 Path("bench_report.json").write_text(json.dumps(rep,indent=2))
145 print(json.dumps(rep,indent=2))
146
147if __name__ == "__main__": main()