Delay-aware event-triggered optimizer / stage2_bench.py
Unverified
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, sweep_baseline, make_report, count_params
8
9TRACK='dynamics'; MODEL='rnn_small'; SEEDS=tuple(range(8))
10# Same union of lr/delay settings is used for both methods.
11GRID=[{'lr':1e-3,'delay':0,'epsilon':0.05},
12 {'lr':3e-3,'delay':1,'epsilon':0.18},
13 {'lr':1e-2,'delay':3,'epsilon':0.50}]
14EPOCHS=12; NTR=400; NTE=200; BATCH=128; WD=0.0
15
16def seed_all(s):
17 random.seed(s); np.random.seed(s); torch.manual_seed(s)
18 if torch.cuda.is_available():
19 try: torch.cuda.manual_seed_all(s)
20 except Exception: pass
21
22def train_one(kind,cfg,seed,collect=False):
23 seed_all(seed); ds=get_dataset(TRACK,seed,n_train=NTR,n_test=NTE)
24 net=make_model(MODEL,ds['input_shape'],ds['out_dim'])
25 lossf=nn.MSELoss(); queued={}; theta_hat=None; event_steps=[]; drift_norms=[]
26 for devname in (['cuda','cpu'] if torch.cuda.is_available() else ['cpu']):
27 try:
28 device=torch.device(devname); net=net.to(device)
29 x,y=ds['xtr'].to(device),ds['ytr'].to(device)
30 opt=torch.optim.SGD(net.parameters(),lr=cfg['lr'],weight_decay=WD)
31 theta_hat=[p.detach().clone() for p in net.parameters()]
32 step=0
33 for ep in range(EPOCHS):
34 gen=torch.Generator(device='cpu').manual_seed(seed+ep)
35 order=torch.randperm(len(x),generator=gen).tolist()
36 for start in range(0,len(x),BATCH):
37 # Apply stale queued corrections before computing the next local update.
38 for corr in queued.pop(step,[]):
39 with torch.no_grad():
40 for p,u in zip(net.parameters(),corr): p.add_(u)
41 idx=order[start:start+BATCH]
42 opt.zero_grad(set_to_none=True); out=net(x[idx]); loss=lossf(out,y[idx]); loss.backward()
43 upd=[(-cfg['lr']*p.grad.detach()).clone() if p.grad is not None else torch.zeros_like(p) for p in net.parameters()]
44 with torch.no_grad():
45 for p,u in zip(net.parameters(),upd): p.add_(u)
46 drift2=sum(((p.detach()-h)**2).sum() for p,h in zip(net.parameters(),theta_hat))
47 grad2=sum((p.grad.detach()**2).sum() for p in net.parameters() if p.grad is not None)
48 v=grad2+1e-3*drift2
49 if kind=='baseline' or drift2 <= cfg['epsilon']*torch.clamp(v,min=1e-30):
50 if kind=='baseline':
51 queued.setdefault(step+cfg['delay'],[]).append(upd)
52 # Idea refreshes reference only on an event; baseline reference is irrelevant.
53 else:
54 queued.setdefault(step+cfg['delay'],[]).append(upd)
55 theta_hat=[p.detach().clone() for p in net.parameters()]
56 event_steps.append(step)
57 drift_norms.append(float(torch.sqrt(drift2).detach().cpu())); step+=1
58 # Flush updates whose actuation time lies within the budget, as a real delayed system does.
59 with torch.no_grad():
60 for due in sorted(queued):
61 for corr in queued[due]:
62 for p,u in zip(net.parameters(),corr): p.add_(u)
63 net.eval(); metric=float(((net(ds['xte'].to(device))-ds['yte'].to(device))**2).mean().detach().cpu())
64 sig=None
65 if collect:
66 gaps=np.diff(event_steps) if len(event_steps)>1 else np.array([])
67 sig={'trained_event_count':len(event_steps),'total_steps':len(drift_norms),
68 'communication_fraction':float(len(event_steps)/max(1,len(drift_norms))),
69 'min_inter_event_gap':int(gaps.min()) if len(gaps) else None,
70 'predicted_gap_proxy':'positive for nonzero threshold',
71 'observed_positive_gap':bool(len(gaps)==0 or gaps.min()>0),
72 'mean_drift_norm':float(np.mean(drift_norms)),
73 'confirmed':bool(len(event_steps)>0 and (len(gaps)==0 or gaps.min()>0))}
74 return metric,sig
75 except RuntimeError:
76 if devname=='cuda':
77 try: torch.cuda.empty_cache()
78 except Exception: pass
79 continue
80 return float('nan'),None
81
82def main():
83 def base_factory(cfg):
84 return lambda s: train_one('baseline',cfg,s)[0]
85 base=sweep_baseline(base_factory,GRID,seeds=(0,1,2,3))
86 runs=[]
87 for cfg in GRID:
88 vals=[]
89 sig=None
90 for s in SEEDS:
91 v,sg=train_one('idea',cfg,s,collect=(s==0)); vals.append(v)
92 if sg is not None: sig=sg
93 runs.append({'cfg':cfg,'result':{'mean':float(np.nanmean(vals)),'std':float(np.nanstd(vals)),
94 'per_seed':[float(v) for v in vals],'n':len(vals)},'signature':sig})
95 best=min(runs,key=lambda r:r['result']['mean'])
96 report=make_report(TRACK,MODEL,base,best['result'],extra={
97 'prediction':'triggered updates reduce transmissions while retaining a positive inter-event gap under delayed actuation',
98 'trained_model_signature':best['signature'],
99 'idea_sweep':[{'cfg':r['cfg'],'mean':r['result']['mean']} for r in runs],
100 'parameter_counts':{'baseline':count_params(make_model(MODEL,(24,),1)),'idea':count_params(make_model(MODEL,(24,),1))},
101 'protocol_note':'baseline and idea use identical rnn_small architecture, SGD lr/delay union, epochs, batches, and paired seeds'})
102 report['protocol']={'seeds':list(SEEDS),'epochs':EPOCHS,'n_train':NTR,'n_test':NTE,'batch':BATCH,'grid':GRID}
103 Path('bench_report.json').write_text(json.dumps(report,indent=2))
104 print(json.dumps(report,indent=2))
105if __name__=='__main__': main()