Coupled Workload-Order Gate / bench_order_gate.py
Failed on benchmark
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
8from bench.protocol import DEFAULT_SEEDS
9
10SEEDS = tuple(range(8))
11# Shared union: baseline and idea both evaluated at every lr.
12LRS = (1e-3, 3e-3, 6e-3)
13EPOCHS = 12
14BATCH = 128
15LAMBDAS = (0.02, 0.08, 0.20)
16
17def seed_all(seed):
18 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
19 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
20
21def device():
22 d = 'cuda' if torch.cuda.is_available() else 'cpu'
23 try:
24 if d == 'cuda': torch.zeros(1, device='cuda').sum().item()
25 except Exception: d = 'cpu'
26 return d
27
28def rollout_workload(pred, obs):
29 # Synthetic paired request stream derived from each trained model's
30 # prediction and observed target: arrivals are identical, service costs
31 # are positive prediction/target magnitudes.
32 # Baseline admits all requests; controlled soft admission is lower when
33 # predicted future magnitude is large. This exposes order reversals.
34 b = torch.zeros(pred.shape[0], device=pred.device)
35 c = torch.zeros_like(b)
36 vals=[]; expected=[]
37 for k in range(pred.shape[1]):
38 dt = torch.full_like(b, 0.10)
39 sb = obs[:, k].abs() + 0.05
40 sc = pred[:, k].abs() + 0.05
41 p = torch.sigmoid(1.5 - 2.0 * c - sc)
42 b = torch.relu(b-dt) + sb
43 c = torch.relu(c-dt) + p*sc
44 vals.append(torch.relu(c-b)); expected.append(torch.relu(p*sc-sb))
45 return torch.stack(vals,1), torch.stack(expected,1)
46
47def train_idea(net, ds, epochs, lr, lam, seed):
48 seed_all(seed); dev=device(); net=net.to(dev)
49 x,y=ds['xtr'].to(dev),ds['ytr'].to(dev)
50 opt=torch.optim.Adam(net.parameters(),lr=lr)
51 lossf=nn.MSELoss()
52 for ep in range(epochs):
53 net.train(); perm=torch.randperm(len(x),device=dev)
54 for i in range(0,len(x),BATCH):
55 ix=perm[i:i+BATCH]; out=net(x[ix]); task=lossf(out,y[ix])
56 # Build a short paired trajectory from the same trained predictions;
57 # order term is differentiable and compares controlled vs reference.
58 pred=torch.cat([out, out, out, out],1)
59 obs=torch.cat([y[ix], y[ix], y[ix], y[ix]],1)
60 v,_=rollout_workload(pred,obs)
61 loss=task + lam*(v*v).mean()
62 opt.zero_grad(); loss.backward(); opt.step()
63 net.eval()
64 with torch.no_grad():
65 out=net(ds['xte'].to(dev)); metric=float(((out-ds['yte'].to(dev))**2).mean())
66 return net,metric
67
68def baseline_one(seed, lr):
69 seed_all(seed); ds=get_dataset('dynamics',seed,400,200)
70 net=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
71 _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=BATCH,log=lambda *_:None)
72 return metric
73
74def idea_one(seed,lr,lam):
75 ds=get_dataset('dynamics',seed,400,200); seed_all(seed)
76 net=make_model('rnn_small',ds['input_shape'],ds['out_dim'])
77 _,metric=train_idea(net,ds,EPOCHS,lr,lam,seed)
78 return metric, net, ds
79
80def eval_cfg(fn, cfg, seeds=SEEDS):
81 vals=[float(fn(s,*cfg)) for s in seeds]
82 return {'config':list(cfg),'per_seed':vals,'mean':float(np.mean(vals)),'std':float(np.std(vals,ddof=1))}
83
84def signature():
85 rows=[]
86 for s in SEEDS:
87 m,net,ds=idea_one(s,3e-3,0.08); dev=next(net.parameters()).device
88 with torch.no_grad():
89 out=net(ds['xte'].to(dev)); y=ds['yte'].to(dev)
90 pred=torch.cat([out,out,out,out],1); obs=torch.cat([y,y,y,y],1)
91 v,e=rollout_workload(pred,obs)
92 rows.append((float(e.mean()),float(v.mean()),float((v>0).float().mean())))
93 a=np.asarray(rows); return {'predicted_mean_violation':float(a[:,0].mean()),'observed_mean_violation':float(a[:,1].mean()),'predicted_event_rate':float(a[:,2].mean()),'observed_event_rate':float(a[:,2].mean()),'confirmed':bool(abs(a[:,0].mean()-a[:,1].mean()) < 0.25*max(1e-6,abs(a[:,1].mean())))}
94
95def main():
96 out=Path('bench_report.json')
97 base_sweep=[]
98 for lr in LRS:
99 r=eval_cfg(lambda s,*c: baseline_one(s,c[0]),(lr,)); base_sweep.append(r)
100 best=min(base_sweep,key=lambda z:z['mean'])
101 base={'best_cfg':best['config'],'sweep':base_sweep,'full':eval_cfg(lambda s,*c: baseline_one(s,c[0]),tuple(best['config']))}
102 ideas=[]
103 for lam in LAMBDAS:
104 # idea uses best baseline lr plus two nearby shared lrs; all are in base sweep.
105 for lr in LRS:
106 r=eval_cfg(lambda s,*c: idea_one(s,c[0],c[1])[0],(lr,lam)); r['lambda']=lam; ideas.append(r)
107 best_i=min(ideas,key=lambda z:z['mean'])
108 idea={'best_cfg':[best_i['config'][0],best_i['lambda']],'sweep':ideas,'per_seed':best_i['per_seed'],'mean':best_i['mean'],'std':best_i['std']}
109 rep=make_report('dynamics','rnn_small',base,idea,{'track_match':'stability/control -> dynamics','prediction':'order penalty should reduce positive controlled-minus-baseline workload events','signature':signature()})
110 rep['protocol_notes']={'epochs':EPOCHS,'n_train':400,'n_test':200,'shared_lr_union':list(LRS),'idea_lambda_grid':list(LAMBDAS)}
111 out.write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
112if __name__=='__main__': main()