Saturation-Adaptive Prefill Chunking / bench_run.py
Failed on benchmark
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
7from bench.protocol import evaluate
8from bench.protocol import DEFAULT_SEEDS
9
10SEEDS = tuple(range(8))
11# Union of all lrs is used on both sides (search-space parity).
12LRS = [1e-3, 3e-3, 1e-2]
13EPOCHS = 8
14NTR, NTE = 400, 160
15
16class AdaptiveChunkRNN(torch.nn.Module):
17 """Same GRU/head as rnn_small, but schedules sequence tokens in chunks.
18
19 A chunk boundary is selected from input-derived saturation (mean absolute
20 angular velocity) and whale/load proxy (fraction of large controls). Hidden
21 state is carried across chunks, so the intervention is scheduling only.
22 """
23 def __init__(self, base, cmin=2, cmax=8, ks=3.0, kw=2.0):
24 super().__init__()
25 self.rnn, self.head = base.rnn, base.head
26 self.cmin, self.cmax, self.ks, self.kw = cmin, cmax, ks, kw
27 self.last_trace = None
28
29 def forward(self, x):
30 seq = x.view(x.shape[0], -1, 3)
31 # saturation normalized by a robust pendulum scale; whale proxy is
32 # high-control occupancy, matching the serving controller's inputs.
33 sat = (seq[:, :, 1].abs().mean(1) / 2.0).clamp(0, 1)
34 whale = (seq[:, :, 2].abs() > 1.0).float().mean(1)
35 chunks = (self.cmax - self.ks * sat - self.kw * whale).clamp(self.cmin, self.cmax).round().long()
36 # Per-example chunk sizes are batched by using the minimum size in a
37 # quantum; this is conservative and deterministic.
38 c = int(chunks.min().item())
39 h = None; states = []
40 for j in range(0, seq.shape[1], c):
41 out, h = self.rnn(seq[:, j:j+c], h)
42 states.append(out[:, -1].detach())
43 self.last_trace = {'chunks': chunks.detach(), 'states': states}
44 return self.head(h[-1])
45
46def seed_all(seed):
47 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
48 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
49
50def baseline_one(lr, seed, return_net=False):
51 seed_all(seed); d=get_dataset('dynamics', seed, NTR, NTE)
52 net=make_model('rnn_small', d['input_shape'], d['out_dim'])
53 net, m, hist=train_model(net,d,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
54 return (m, net, d) if return_net else m
55
56def idea_one(lr, seed, return_net=False):
57 seed_all(seed); d=get_dataset('dynamics', seed, NTR, NTE)
58 base=make_model('rnn_small', d['input_shape'], d['out_dim'])
59 net=AdaptiveChunkRNN(base)
60 net, m, hist=train_model(net,d,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
61 return (m, net, d) if return_net else m
62
63def main():
64 # Baseline sweep and idea sweep use identical lr union and 4 tuning seeds.
65 grid=[{'lr':lr} for lr in LRS]
66 base=sweep_baseline(lambda cfg: (lambda seed: baseline_one(cfg['lr'],seed)), grid=grid)
67 # sweep_baseline expects a factory; inspect result shape defensively.
68 best_cfg=base.get('best_cfg', grid[1])
69 if isinstance(best_cfg, dict) and 'lr' in best_cfg: best_lr=best_cfg['lr']
70 else: best_lr=3e-3
71 idea_grid=[{'lr':lr} for lr in LRS]
72 idea_sweep=sweep_baseline(lambda cfg: (lambda seed: idea_one(cfg['lr'],seed)), grid=idea_grid)
73 # Explicit paired eight-seed results at the idea sweep winner.
74 idea_lr=idea_sweep.get('best_cfg', {'lr':best_lr}).get('lr',best_lr)
75 base_res=[baseline_one(idea_lr,s) for s in SEEDS]
76 idea_res=[idea_one(idea_lr,s) for s in SEEDS]
77 # Model-derived NN-scale mechanism signature: activation state ramps on
78 # the trained systems, not an analytical identity.
79 sig=[]
80 for s in SEEDS:
81 bm,bn,d=baseline_one(idea_lr,s,True); im,inn,_=idea_one(idea_lr,s,True)
82 with torch.no_grad():
83 # The training ladder may leave weights on a shared GPU. The
84 # signature is only a diagnostic, so run it robustly on CPU.
85 bn=bn.cpu(); inn=inn.cpu(); xb=d['xte'][:128].cpu()
86 old_cudnn=torch.backends.cudnn.enabled; torch.backends.cudnn.enabled=False
87 bn.eval(); inn.eval()
88 try:
89 z=xb.view(len(xb),-1,3); out,h=bn.rnn(z); bstates=out
90 _=inn(xb); astates=torch.stack(inn.last_trace['states'],1)
91 finally:
92 torch.backends.cudnn.enabled=old_cudnn
93 br=float(torch.diff(bstates,dim=1).abs().mean())
94 ar=float(torch.diff(astates,dim=1).abs().mean()) if astates.shape[1]>1 else 0.
95 sig.append({'seed':s,'baseline_state_ramp':br,'adaptive_state_ramp':ar})
96 br=float(np.mean([x['baseline_state_ramp'] for x in sig])); ar=float(np.mean([x['adaptive_state_ramp'] for x in sig]))
97 report=make_report('dynamics','rnn_small',base,idea_sweep['full'],extra={
98 'prediction':'adaptive chunking lowers high-load sequential-state ramp while preserving task function',
99 'observed_baseline_state_ramp':br,'observed_adaptive_state_ramp':ar,
100 'relative_reduction':1-ar/br if br else 0.,'confirmed': bool(ar < br),
101 'note':'activation ramp is a hardware-independent proxy; no NVML power available'})
102 report['custom_track']=None
103 report['idea_sweep']=idea_sweep
104 report['idea_lr']=idea_lr
105 report['paired_baseline_at_idea_lr']=base_res
106 print(json.dumps(report,indent=2))
107
108if __name__=='__main__': main()