Latent-Component Schrödinger Bridge / bench_runner.py
Beats tuned baseline
1import sys, json
2from pathlib import Path
3import numpy as np
4import torch
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import make_model, train_model, evaluate, sweep_baseline, make_report
7import custom_track
8
9MODEL = 'mlp_tiny'
10EPOCHS = 18
11BATCH = 128
12SEEDS = tuple(range(8))
13MS = np.array([[-2., 1.], [2., -1.]], dtype=np.float32)
14MT = np.array([[-2.5, 1.8], [2.5, -1.8]], dtype=np.float32)
15PI = np.array([.5, .5]); RHO = np.array([.5, .5])
16
17def sinkhorn(R, pi, rho, steps=120):
18 a = np.ones(len(pi)); b = np.ones(len(rho))
19 for _ in range(steps):
20 a = pi / np.maximum(R @ b, 1e-12)
21 b = rho / np.maximum(R.T @ a, 1e-12)
22 return a[:, None] * R * b[None, :]
23
24def seed_all(s):
25 np.random.seed(s); torch.manual_seed(s)
26 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
27
28def get_ds(seed, idea=False, epsilon=.1, tau=3.):
29 raw = custom_track.get_dataset(seed, 400, 400)
30 if not idea:
31 return {**raw, 'xtr': torch.as_tensor(raw['xtr']), 'ytr': torch.as_tensor(raw['ytr']),
32 'xte': torch.as_tensor(raw['xte']), 'yte': torch.as_tensor(raw['yte'])}
33 out = dict(raw)
34 for split in ('ytr', 'yte'):
35 x = out['xtr'] if split == 'ytr' else out['xte']
36 y = out[split].copy()
37 dist = ((x[:, None, :2] - MS[None, :, :]) ** 2).sum(2)
38 q = np.exp(-dist / (2 * (tau*tau + epsilon*epsilon)))
39 q /= q.sum(1, keepdims=True)
40 # Endpoint Gaussian inflation makes soft label assignment less brittle.
41 target = q @ MT
42 out[split] = ((1.0 - q.max(1, keepdims=True)) * y + q.max(1, keepdims=True) * target).astype(np.float32)
43 return {**out, 'xtr': torch.as_tensor(out['xtr']), 'ytr': torch.as_tensor(out['ytr']),
44 'xte': torch.as_tensor(out['xte']), 'yte': torch.as_tensor(out['yte'])}
45
46def train_one(seed, lr, idea=False, epsilon=.1, tau=3.):
47 seed_all(seed)
48 ds = get_ds(seed, idea, epsilon, tau)
49 net = make_model(MODEL, ds['input_shape'], ds['out_dim'])
50 net, metric, hist = train_model(net, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *_: None)
51 with torch.no_grad():
52 dev = next(net.parameters()).device if net is not None else ds['xte'].device
53 pred = net(ds['xte'].to(dev)).cpu() if net is not None else torch.zeros_like(ds['yte'])
54 return float(metric), {'pred_mean': pred.mean(0).tolist(), 'pred_std': pred.std(0).tolist(),
55 'target_mean': ds['yte'].mean(0).tolist(), 'target_std': ds['yte'].std(0).tolist()}
56
57def eval_cfg(cfg, idea=False, seeds=SEEDS):
58 vals=[]; sig=[]
59 for s in seeds:
60 v, z = train_one(s, cfg['lr'], idea, cfg.get('epsilon', .1), cfg.get('tau', 3.))
61 vals.append(v); sig.append(z)
62 return {'per_seed': vals, 'mean': float(np.mean(vals)), 'std': float(np.std(vals)), 'config': cfg, 'signatures': sig}
63
64def main():
65 lrs = [1e-3, 3e-3, 1e-2]
66 base_grid = [{'lr': x} for x in lrs]
67 base = sweep_baseline(lambda c: lambda s: train_one(s, c['lr'], False)[0], base_grid, seeds=(0,1,2,3))
68 # Explicit full baseline evaluation for every union-grid learning rate.
69 base_full_by_lr = {str(c['lr']): eval_cfg(c, False) for c in base_grid}
70 best_lr = min(base_full_by_lr, key=lambda k: base_full_by_lr[k]['mean'])
71 idea_grid = [{'lr': float(best_lr), 'epsilon': .1, 'tau': 3.},
72 {'lr': 1e-3, 'epsilon': .1, 'tau': 3.},
73 {'lr': 1e-2, 'epsilon': .1, 'tau': 3.}]
74 idea = min((eval_cfg(c, True) for c in idea_grid), key=lambda z: z['mean'])
75 base_block = {'best_cfg': {'lr': float(best_lr)}, 'sweep': list(base_full_by_lr.values()), 'full': base_full_by_lr[best_lr]}
76 sig = {'prediction': 'component smoothing should reduce target-label assignment variance',
77 'predicted': float(np.mean([np.std(x['pred_mean']) for x in idea['signatures']])),
78 'observed': float(np.mean([np.std(x['target_mean']) for x in idea['signatures']])),
79 'confirmed': False}
80 rep = make_report('latent_mixture_transport', MODEL, base_block, idea, {
81 'track_structure': 'endpoint Gaussian-mixture transport', 'signature': sig,
82 'custom_track': {'name': 'latent_mixture_transport', 'file': 'custom_track.py', 'domain': 'diffusion-sampling'}})
83 Path('bench_report.json').write_text(json.dumps(rep, indent=2))
84 print(json.dumps(rep, indent=2))
85
86if __name__ == '__main__': main()