Delay-aware event-triggered optimizer / stage2_bench.py

Unverified

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