Interval-Certified Equilibrium Layer / stage2_bench.py

Failed on benchmark

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