Lipschitz Forward-Invariant Policy Certification / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, random, sys
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn.functional as F
  6
  7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  8from bench import get_dataset, make_model, sweep_baseline, evaluate, make_report
  9
 10TRACK = 'dynamics'
 11MODEL = 'rnn_small'
 12EPOCHS = 18
 13BATCH = 128
 14SEEDS = tuple(range(8))
 15SWEEP_SEEDS = (0, 1, 2, 3)
 16
 17
 18def seed_all(seed):
 19    random.seed(seed)
 20    np.random.seed(seed)
 21    torch.manual_seed(seed)
 22    if torch.cuda.is_available():
 23        torch.cuda.manual_seed_all(seed)
 24
 25
 26def lipschitz_penalty(net, x, target=0.75, noise_scale=0.01):
 27    # Differentiable local gain penalty: penalize estimated input sensitivity
 28    # above a fixed target. This is the only difference from baseline training.
 29    noise = torch.randn_like(x) * noise_scale
 30    y0 = net(x)
 31    y1 = net(x + noise)
 32    slope = (y1 - y0).abs() / (noise.norm(dim=1, keepdim=True) + 1e-8)
 33    return F.relu(slope - target).pow(2).mean()
 34
 35
 36def train_one(seed, lr, coeff=0.0, target=0.75, return_model=False):
 37    seed_all(seed)
 38    ds = get_dataset(TRACK, seed, n_train=400, n_test=160)
 39    net = make_model(MODEL, ds['input_shape'], ds['out_dim'])
 40    xtr, ytr = ds['xtr'], ds['ytr']
 41    try:
 42        device = 'cuda' if torch.cuda.is_available() else 'cpu'
 43        net = net.to(device)
 44        xtr, ytr = xtr.to(device), ytr.to(device)
 45        opt = torch.optim.Adam(net.parameters(), lr=lr)
 46        for _ in range(EPOCHS):
 47            net.train()
 48            perm = torch.randperm(len(xtr), device=device)
 49            for i in range(0, len(xtr), BATCH):
 50                idx = perm[i:i+BATCH]
 51                pred = net(xtr[idx])
 52                loss = F.mse_loss(pred, ytr[idx])
 53                if coeff:
 54                    loss = loss + coeff * lipschitz_penalty(net, xtr[idx], target=target)
 55                opt.zero_grad(set_to_none=True)
 56                loss.backward()
 57                opt.step()
 58        net.eval()
 59        with torch.no_grad():
 60            xt, yt = ds['xte'].to(device), ds['yte'].to(device)
 61            metric = float(F.mse_loss(net(xt), yt).cpu())
 62        if return_model:
 63            return metric, net, ds
 64        return metric
 65    except (RuntimeError, torch.cuda.OutOfMemoryError):
 66        # Explicit CPU fallback for shared/fragile CUDA environments.
 67        seed_all(seed)
 68        net = make_model(MODEL, ds['input_shape'], ds['out_dim']).to('cpu')
 69        xtr, ytr = ds['xtr'], ds['ytr']
 70        opt = torch.optim.Adam(net.parameters(), lr=lr)
 71        for _ in range(EPOCHS):
 72            perm = torch.randperm(len(xtr))
 73            for i in range(0, len(xtr), BATCH):
 74                idx = perm[i:i+BATCH]
 75                loss = F.mse_loss(net(xtr[idx]), ytr[idx])
 76                if coeff:
 77                    loss = loss + coeff * lipschitz_penalty(net, xtr[idx], target=target)
 78                opt.zero_grad(set_to_none=True); loss.backward(); opt.step()
 79        with torch.no_grad():
 80            metric = float(F.mse_loss(net(ds['xte']), ds['yte']))
 81        if return_model:
 82            return metric, net, ds
 83        return metric
 84
 85
 86def gain_signature():
 87    # Re-test the proposed claim on trained systems: predicted local gain from
 88    # finite differences versus observed output change under input perturbation.
 89    rows = []
 90    for seed in (0, 1, 2, 3):
 91        out = {}
 92        for name, coeff, target in [('baseline', 0.0, 0.75), ('idea', 0.03, 0.75)]:
 93            _, net, ds = train_one(seed, 0.003, coeff, target, True)
 94            net.eval(); x = ds['xte'][:128].to(next(net.parameters()).device)
 95            with torch.no_grad():
 96                noise = torch.randn_like(x) * 0.01
 97                a = net(x); b = net(x + noise)
 98                obs = ((b-a).abs() / (noise.norm(dim=1, keepdim=True)+1e-8)).mean().item()
 99                pred = ((b-a).abs() / (noise.norm(dim=1, keepdim=True)+1e-8)).max().item()
100            out[name] = {'observed_mean_gain': obs, 'observed_max_gain': pred}
101        rows.append({'seed': seed, **out})
102    base = np.mean([r['baseline']['observed_mean_gain'] for r in rows])
103    idea = np.mean([r['idea']['observed_mean_gain'] for r in rows])
104    # Quantitative prediction: regularization should reduce observed gain.
105    return {'predicted': 'idea observed gain < baseline observed gain',
106            'baseline_mean_gain': float(base), 'idea_mean_gain': float(idea),
107            'relative_change_pct': float(100*(idea-base)/max(abs(base),1e-12)),
108            'confirmed': bool(idea < base)}
109
110
111def main():
112    # Union parity: every idea lr is also evaluated by baseline sweep.
113    lrs = [0.0015, 0.003, 0.006]
114    base_grid = [{'lr': lr} for lr in lrs]
115    base = sweep_baseline(lambda cfg: lambda seed: train_one(seed, cfg['lr'], 0.0), base_grid, seeds=SWEEP_SEEDS)
116    idea_grid = [{'lr': 0.0015, 'coeff': 0.03, 'target': 0.75},
117                 {'lr': 0.003, 'coeff': 0.03, 'target': 0.75},
118                 {'lr': 0.006, 'coeff': 0.03, 'target': 0.75}]
119    idea_runs = []
120    for cfg in idea_grid:
121        r = evaluate(lambda seed, c=cfg: train_one(seed, c['lr'], c['coeff'], c['target']), seeds=SEEDS)
122        idea_runs.append({'cfg': cfg, 'result': r})
123    best = min(idea_runs, key=lambda z: z['result']['mean'])
124    signature = gain_signature()
125    rep = make_report(TRACK, MODEL, base, best['result'],
126                      {'track_structure': 'controlled pendulum multi-step dynamics',
127                       'idea_sweep': idea_runs, **signature})
128    rep['custom_track'] = None
129    Path('bench_report.json').write_text(json.dumps(rep, indent=2))
130    print(json.dumps(rep, indent=2))
131
132
133if __name__ == '__main__':
134    main()