Interval-Certified Equilibrium Layer / stage2_bench.py
Failed on benchmark
1import os, sys, json, time
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
7
8SEEDS=tuple(range(8))
9# Same union on both sides; baseline sweep uses all values tried by idea.
10GRID=[{'lr':1e-3,'iters':8},{'lr':3e-3,'iters':16},{'lr':6e-3,'iters':24}]
11
12class EquilibriumRNN(nn.Module):
13 """Small implicit tanh recurrent layer followed by a readout.
14 Both systems have exactly this architecture and parameters; only solve policy differs.
15 """
16 def __init__(self, in_dim=3, hidden=32, out_dim=1, iters=16, certified=False):
17 super().__init__(); self.hidden=hidden; self.iters=iters; self.certified=certified
18 self.in_proj=nn.Linear(in_dim,hidden); self.h_proj=nn.Linear(hidden,hidden)
19 self.head=nn.Linear(hidden,out_dim)
20 self.last_stats={'q':[], 'certified':0, 'fallback':0}
21 def _map(self,h, x): return torch.tanh(self.in_proj(x)+self.h_proj(h))
22 def _solve(self,x,n):
23 h=torch.zeros(x.shape[0],self.hidden,device=x.device)
24 for _ in range(n): h=self._map(h,x)
25 return h
26 def _certificate(self,x,h):
27 # Conservative local interval Jacobian bound on a box h +/- radius.
28 # For tanh(Ax+Bh), |d tanh| <= sech^2 of interval preactivation;
29 # use a conservative finite-radius analytic bound.
30 rad=.15
31 with torch.no_grad():
32 pre=self.in_proj(x)+self.h_proj(h)
33 br=rad*torch.sum(torch.abs(self.h_proj.weight),dim=1)
34 lo=pre-br; hi=pre+br
35 # max sech^2 over interval, exact for intervals crossing zero
36 near=torch.minimum(torch.abs(lo),torch.abs(hi))
37 near=torch.where((lo<=0)&(hi>=0),torch.zeros_like(near),near)
38 sech2=1/torch.cosh(near).clamp_min(1e-6)**2
39 # induced infinity norm of interval Jacobian
40 q=float(torch.max(sech2[:,None]*torch.abs(self.h_proj.weight).sum(dim=1)).item())
41 # residual and a simple Krawczyk enclosure radius; conservative scalar bound
42 res=torch.max(torch.abs(self._map(h,x)-h),dim=1).values
43 margin=(res/(1-max(q,0.0)) if q<1 else torch.full_like(res,float('inf')))
44 ok=(q<0.8) & (margin < rad*.5)
45 return q, bool(torch.all(ok).item())
46 def forward(self,x):
47 seq=x.view(x.shape[0],-1,3)
48 # equilibrium conditioning uses the complete observed sequence, not a new readout
49 u=seq.mean(dim=1)
50 self.last_stats={'q':[], 'certified':0, 'fallback':0}
51 if not self.certified:
52 h=self._solve(u,self.iters)
53 else:
54 # inexpensive approximate center, then interval certificate; fallback is standard solve
55 h0=self._solve(u,min(8,self.iters))
56 q,ok=self._certificate(u,h0)
57 self.last_stats={'q':[q], 'certified':int(ok), 'fallback':int(not ok)}
58 h=h0 if ok else self._solve(u,self.iters)
59 return self.head(h)
60
61def run_one(kind,cfg,seed,collect=False):
62 torch.manual_seed(seed); np.random.seed(seed)
63 d=get_dataset('dynamics',seed,n_train=400,n_test=200)
64 net=EquilibriumRNN(3,32,1,iters=int(cfg['iters']),certified=(kind=='idea'))
65 net,metric,_=train_model(net,d,epochs=10,lr=float(cfg['lr']),batch=128,log=lambda *_:None)
66 if net is None: return float('inf')
67 if collect:
68 with torch.no_grad():
69 dev=next(net.parameters()).device
70 _=net(d['xte'][:128].to(dev))
71 return float(metric), dict(net.last_stats)
72 return float(metric)
73
74def factory(kind):
75 return lambda cfg: (lambda seed: run_one(kind,cfg,seed))
76
77def main():
78 t=time.time()
79 base=sweep_baseline(factory('base'),GRID,seeds=(0,1,2,3))
80 idea_cfgs=[]
81 for cfg in GRID:
82 vals=evaluate(factory('idea')(cfg),SEEDS)
83 idea_cfgs.append({'cfg':cfg,'result':vals})
84 best=min(idea_cfgs,key=lambda z:z['result']['mean'])
85 sig=[]
86 for s in SEEDS:
87 val,st=run_one('idea',best['cfg'],s,collect=True)
88 sig.append({'seed':s,'q_observed':st['q'][0] if st['q'] else None,
89 'certified_batches':st['certified'],'fallback_batches':st['fallback']})
90 qs=[z['q_observed'] for z in sig if z['q_observed'] is not None]
91 certified=[z['certified_batches'] for z in sig]
92 # Prediction: q<0.8 should certify; this is measured on trained models.
93 low=[q for q in qs if q<.8]; high=[q for q in qs if q>=.8]
94 signature={'prediction':'trained-model local q below 0.8 predicts certification; q near/above 1 predicts fallback',
95 'n_models':len(qs),'predicted_low_q_count':len(low),'observed_certified_low_q_count':sum(c>0 for q,c in zip(qs,certified) if q<.8),
96 'mean_q':float(np.mean(qs)) if qs else None,'min_q':float(np.min(qs)) if qs else None,
97 'max_q':float(np.max(qs)) if qs else None,'q_values':qs,'certified_flags':certified,
98 'confirmed':bool(low and all(c>0 for q,c in zip(qs,certified) if q<.8) and all(c==0 for q,c in zip(qs,certified) if q>=.8))}
99 idea=dict(best['result']); idea['best_cfg']=best['cfg']; idea['sweep']=[{'cfg':z['cfg'],'mean':z['result']['mean']} for z in idea_cfgs]
100 rep=make_report('dynamics','custom_equilibrium_rnn',base,idea,{'mechanism_signature':signature,
101 'architecture_note':'Both arms train the same implicit tanh recurrent layer; baseline always iterates, idea certifies then falls back.',
102 'runtime_sec':time.time()-t})
103 json.dump(rep,open('bench_report.json','w'),indent=2)
104 print(json.dumps(rep,indent=2))
105if __name__=='__main__': main()