Order-Adaptive Integral Optimizer / stage2_bench.py
Failed on benchmark
1import json, math, os, sys
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
7from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10SWEEP_SEEDS = tuple(range(4))
11EPOCHS = 18
12BATCH = 128
13
14class OrderAdaptiveIntegral(torch.optim.Optimizer):
15 """Global-order nested integrated-gradient optimizer for this benchmark."""
16 def __init__(self, params, lr, max_order=2, beta=.95, rho=.98,
17 interval=2, patience=2, ramp_steps=8,
18 gains=(.05, .0005), state_clip=10.0, a0=1.0):
19 super().__init__(params, dict(lr=lr))
20 self.max_order, self.beta, self.rho = max_order, beta, rho
21 self.interval, self.patience, self.ramp_steps = interval, patience, ramp_steps
22 self.gains, self.state_clip, self.a0 = tuple(gains), state_clip, a0
23 self.order, self.ema, self.bad = 0, None, 0
24 self.step_count, self.activations = 0, []
25 for group in self.param_groups:
26 for p in group['params']:
27 self.state[p]['z'] = {}
28 self.state[p]['gamma'] = {}
29
30 @torch.no_grad()
31 def step(self):
32 ss, count = 0.0, 0
33 for group in self.param_groups:
34 for p in group['params']:
35 if p.grad is not None:
36 ss += float(p.grad.detach().pow(2).sum())
37 count += p.numel()
38 gn = math.sqrt(ss) / max(math.sqrt(count), 1.0)
39 self.ema = gn if self.ema is None else self.beta*self.ema + (1-self.beta)*gn
40 self.step_count += 1
41 if self.step_count % self.interval == 0 and self.order < self.max_order:
42 previous = getattr(self, '_last_decision_ema', None)
43 if previous is not None and self.ema > self.rho * previous:
44 self.bad += 1
45 else:
46 self.bad = 0
47 self._last_decision_ema = self.ema
48 if self.bad >= self.patience:
49 self.order += 1
50 self.bad = 0
51 self.activations.append({'step': self.step_count, 'order': self.order,
52 'ema': float(self.ema)})
53 for group in self.param_groups:
54 for p in group['params']:
55 self.state[p]['z'][self.order] = torch.zeros_like(p)
56 self.state[p]['gamma'][self.order] = 0.0
57 for group in self.param_groups:
58 lr = group['lr']
59 for p in group['params']:
60 if p.grad is None:
61 continue
62 st, g = self.state[p], p.grad
63 update = self.a0 * g
64 for r in range(1, self.order + 1):
65 z = st['z'][r]
66 z.add_(g if r == 1 else st['z'][r-1], alpha=lr)
67 zn = float(z.norm())
68 if np.isfinite(zn) and zn > self.state_clip:
69 z.mul_(self.state_clip / zn)
70 st['gamma'][r] = min(1.0, st['gamma'][r] + 1.0/max(1, self.ramp_steps))
71 update = update + st['gamma'][r] * self.gains[r-1] * z
72 p.add_(update, alpha=-lr)
73
74def train_one(track, model_name, cfg, seed, adaptive):
75 torch.manual_seed(seed)
76 np.random.seed(seed)
77 ds = get_dataset(track, seed, n_train=400, n_test=400)
78 model = make_model(model_name, ds['input_shape'], ds['out_dim'])
79 device = 'cuda' if torch.cuda.is_available() else 'cpu'
80 lossf = nn.MSELoss()
81 try:
82 model = model.to(device)
83 xtr, ytr = ds['xtr'].to(device), ds['ytr'].to(device)
84 if adaptive:
85 opt = OrderAdaptiveIntegral(model.parameters(), lr=cfg['lr'], max_order=2,
86 interval=2, patience=2, ramp_steps=8)
87 else:
88 opt = torch.optim.SGD(model.parameters(), lr=cfg['lr'], momentum=cfg['momentum'])
89 hist, grad_hist, order_hist = [], [], []
90 for _ in range(EPOCHS):
91 model.train(); perm = torch.randperm(len(xtr), device=device); total = 0.
92 for i in range(0, len(xtr), BATCH):
93 idx = perm[i:i+BATCH]
94 loss = lossf(model(xtr[idx]), ytr[idx])
95 opt.zero_grad(); loss.backward()
96 gn = math.sqrt(sum(float(p.grad.detach().pow(2).sum()) for p in model.parameters() if p.grad is not None))
97 opt.step(); total += float(loss) * len(idx)
98 grad_hist.append(gn)
99 order_hist.append(getattr(opt, 'order', 0))
100 hist.append(total / len(xtr))
101 model.eval()
102 with torch.no_grad(): metric = float(lossf(model(ds['xte'].to(device)), ds['yte'].to(device)))
103 return metric, {'model': model, 'history': hist, 'grad_hist': grad_hist,
104 'order_hist': order_hist, 'activations': getattr(opt, 'activations', [])}
105 except RuntimeError:
106 if device == 'cuda':
107 torch.cuda.empty_cache()
108 return train_one_cpu(track, model_name, cfg, seed, adaptive)
109 raise
110
111def train_one_cpu(track, model_name, cfg, seed, adaptive):
112 torch.manual_seed(seed); np.random.seed(seed)
113 ds = get_dataset(track, seed, n_train=400, n_test=400)
114 model = make_model(model_name, ds['input_shape'], ds['out_dim'])
115 opt = (OrderAdaptiveIntegral(model.parameters(), lr=cfg['lr'], max_order=2, interval=2, patience=2, ramp_steps=8)
116 if adaptive else torch.optim.SGD(model.parameters(), lr=cfg['lr'], momentum=cfg['momentum']))
117 lossf = nn.MSELoss(); hist=[]; grads=[]; orders=[]
118 for _ in range(EPOCHS):
119 perm=torch.randperm(len(ds['xtr'])); total=0.
120 for i in range(0,len(perm),BATCH):
121 idx=perm[i:i+BATCH]; loss=lossf(model(ds['xtr'][idx]),ds['ytr'][idx]); opt.zero_grad(); loss.backward()
122 grads.append(math.sqrt(sum(float(p.grad.detach().pow(2).sum()) for p in model.parameters() if p.grad is not None)))
123 opt.step(); orders.append(getattr(opt,'order',0)); total += float(loss)*len(idx)
124 hist.append(total/len(perm))
125 with torch.no_grad(): metric=float(lossf(model(ds['xte']),ds['yte']))
126 return metric, {'model':model,'history':hist,'grad_hist':grads,'order_hist':orders,'activations':getattr(opt,'activations',[])}
127
128def main():
129 track, model_name = 'tabular', 'mlp_tiny'
130 # Union of baseline and idea step sizes; baseline also sweeps its central momentum knob.
131 grid = [{'lr': lr, 'momentum': m} for lr in (0.003, 0.01, 0.03) for m in (0.0, 0.9)]
132 base = sweep_baseline(lambda c: lambda s: train_one(track, model_name, c, s, False)[0], grid, seeds=SWEEP_SEEDS)
133 best_lr, best_m = base['best_cfg']['lr'], base['best_cfg']['momentum']
134 idea_grid = [best_lr, 0.003 if best_lr != 0.003 else 0.01, 0.03 if best_lr != 0.03 else 0.01]
135 idea_cfgs = [{'lr': x, 'momentum': best_m} for x in idea_grid]
136 # Evaluate all three candidate idea settings on the same full paired seeds.
137 candidates=[]
138 for c in idea_cfgs:
139 vals=[]
140 for s in SEEDS: vals.append(train_one(track, model_name, c, s, True)[0])
141 candidates.append({'cfg':c,'result':{'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals)}})
142 best=min(candidates,key=lambda x:x['result']['mean'])
143 sig=[]
144 for s in SEEDS:
145 _, info=train_one(track, model_name, best['cfg'], s, True)
146 gh=np.asarray(info['grad_hist']); oh=np.asarray(info['order_hist'])
147 active=np.where(oh>0)[0]
148 pre=float(np.mean(gh[:max(1, len(gh)//3)])); post=float(np.mean(gh[-max(1, len(gh)//3):]))
149 sig.append({'seed':s,'activated':bool(len(active)),'first_activation':int(active[0]) if len(active) else None,'grad_ratio':post/max(pre,1e-12)})
150 idea_res=best['result']
151 baseline=base
152 # Signature prediction: easy tabular training should usually stay order zero; measure trained behavior.
153 frac=float(np.mean([x['activated'] for x in sig]))
154 confirmed=bool(frac < 0.5 and all(np.isfinite(x['grad_ratio']) for x in sig))
155 report=make_report(track, model_name, baseline, idea_res, {'mechanism_signature':{
156 'prediction':'on easy well-conditioned tabular training adaptive order remains mostly p=0 while residual contracts',
157 'observed_activation_fraction':frac,'per_seed':sig,'confirmed':confirmed},
158 'protocol_note':'baseline sweep uses all lr values tried by idea and sweeps momentum; idea uses three lr settings.'})
159 report['idea_candidates']=candidates
160 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
161 print(json.dumps(report,indent=2))
162
163if __name__ == '__main__': main()