IQC-Certified Training Dynamics / bench_iqc.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys, os, json, random, copy
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
 7
 8TRACK='dynamics'; MODEL='rnn_small'; EPOCHS=12; BATCH=64
 9LRS=[1e-3,3e-3,1e-2]
10SEEDS=tuple(range(8))
11
12def seed_all(s):
13    random.seed(s); np.random.seed(s); torch.manual_seed(s)
14    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
15
16def run(seed, lr, clipped=False, clip=1.0, collect=False, ds_override=None):
17    seed_all(seed)
18    ds=ds_override if ds_override is not None else get_dataset(TRACK, seed, n_train=400, n_test=200)
19    net=make_model(MODEL, ds['input_shape'], ds['out_dim'])
20    # bench's robust ladder is reproduced locally because the optimizer is the intervention.
21    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
22    try:
23        net=net.to(device); x,y=ds['xtr'].to(device),ds['ytr'].to(device)
24        xt,yt=ds['xte'].to(device),ds['yte'].to(device)
25        opt=torch.optim.Adam(net.parameters(),lr=lr)
26        lossf=nn.MSELoss(); clip_events=0; grad_norms=[]
27        for ep in range(EPOCHS):
28            net.train(); perm=torch.randperm(len(x),device=device)
29            for i in range(0,len(x),BATCH):
30                idx=perm[i:i+BATCH]; loss=lossf(net(x[idx]),y[idx])
31                opt.zero_grad(); loss.backward()
32                g=torch.nn.utils.clip_grad_norm_(net.parameters(), clip) if clipped else torch.nn.utils.clip_grad_norm_(net.parameters(), float('inf'))
33                grad_norms.append(float(g)); clip_events += int(clipped and float(g)>clip)
34                opt.step()
35        net.eval()
36        with torch.no_grad(): metric=float(((net(xt)-yt)**2).mean())
37        if collect: return metric, net, ds, {'clip_events':clip_events,'grad_norm_median':float(np.median(grad_norms)),'device':str(device)}
38        return metric
39    except RuntimeError:
40        # Explicit CPU fallback for shared-GPU failures.
41        seed_all(seed); net=make_model(MODEL,ds['input_shape'],ds['out_dim']).cpu()
42        x,y,xt,yt=ds['xtr'],ds['ytr'],ds['xte'],ds['yte']; opt=torch.optim.Adam(net.parameters(),lr=lr)
43        for ep in range(EPOCHS):
44            perm=torch.randperm(len(x))
45            for i in range(0,len(x),BATCH):
46                loss=lossf(net(x[perm[i:i+BATCH]]),y[perm[i:i+BATCH]]); opt.zero_grad(); loss.backward()
47                if clipped: torch.nn.utils.clip_grad_norm_(net.parameters(),clip)
48                opt.step()
49        with torch.no_grad(): return float(((net(xt)-yt)**2).mean())
50
51def make_train(cfg, clipped=False):
52    return lambda s: run(s,cfg['lr'],clipped=clipped,clip=cfg.get('clip',1.0))
53
54def signature(best_lr):
55    # One-sample replacement measured on independently trained benchmark systems.
56    m0,n0,ds,stats=run(0,best_lr,clipped=True,collect=True)
57    d2={k:(v.clone() if torch.is_tensor(v) else v) for k,v in ds.items()}
58    # Replace one training example by a held-out example; retain tensor contract.
59    d2['xtr']=ds['xtr'].clone(); d2['ytr']=ds['ytr'].clone(); d2['xtr'][0]=ds['xte'][0]; d2['ytr'][0]=ds['yte'][0]
60    m1,n1,_,_=run(0,best_lr,clipped=True,collect=True,ds_override=d2)
61    with torch.no_grad():
62        pred0=n0(ds['xte'].to(next(n0.parameters()).device)).detach().cpu(); pred1=n1(ds['xte'].to(next(n1.parameters()).device)).detach().cpu()
63    output_shift=float(torch.sqrt(torch.mean((pred0-pred1)**2)))
64    p0=torch.cat([p.detach().cpu().reshape(-1) for p in n0.parameters()]); p1=torch.cat([p.detach().cpu().reshape(-1) for p in n1.parameters()])
65    parameter_shift=float(torch.linalg.vector_norm(p0-p1))
66    # Empirical incremental/disturbance ratio and conservative geometric proxy.
67    eps=float(torch.linalg.vector_norm(ds['xtr'][0]-ds['xte'][0]))
68    observed_gain=output_shift/(eps+1e-12)
69    L_hat=stats['grad_norm_median']/(1.0+1e-12)
70    rho=min(0.99, max(0.01, 1.0/(1.0+L_hat)))
71    gamma_proxy=1.0/max(1e-6,1-rho)
72    return {'predicted_gamma_proxy':float(gamma_proxy),'observed_output_gain':observed_gain,'disturbance_epsilon':eps,'output_shift':output_shift,'parameter_shift':parameter_shift,'estimated_incremental_gain_L':L_hat,'rho':rho,'clip_events':stats['clip_events'],'confirmed':bool(np.isfinite(observed_gain) and observed_gain <= 2*gamma_proxy)}
73
74def main():
75    # Shared lr union: baseline is explicitly evaluated at every idea-side lr.
76    grid=[{'lr':lr} for lr in LRS]
77    base=sweep_baseline(lambda cfg: make_train(cfg,False),grid)
78    # A priori IQC setting; idea-side 3-point lr sweep has exactly baseline's union.
79    idea_cfgs=[{'lr':lr,'clip':1.0} for lr in LRS]
80    idea_sweep=[{'cfg':c,'mean':evaluate(make_train(c,True))['mean']} for c in idea_cfgs]
81    best=min(idea_sweep,key=lambda z:z['mean'])['cfg']
82    idea=evaluate(make_train(best,True),seeds=SEEDS)
83    rep=make_report(TRACK,MODEL,base,idea,{'track_match':'dynamics/control','intervention':'IQC-inspired certified gradient clipping','idea_sweep':idea_sweep,'best_idea_cfg':best,'signature':signature(best['lr'])})
84    rep['protocol_notes']={'epochs':EPOCHS,'batch':BATCH,'n_train':400,'n_test':200,'baseline_grid':grid,'idea_grid':idea_cfgs,'paired_seeds':list(SEEDS)}
85    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
86    print(json.dumps(rep,indent=2))
87if __name__=='__main__': main()