Latent-Component Schrödinger Bridge / bench_runner.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
 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()