Saturation-Adaptive Prefill Chunking / bench_run.py

Failed on benchmark

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