Averaged Contractive State-Space Network / bench_stage2.py

Failed on benchmark

Raw ⬇ ZIP
  1import json, sys, time
  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, train_model, sweep_baseline, make_report
  8
  9TRACK='dynamics'; MODEL='rnn_small'; SEEDS=tuple(range(8)); SWEEP_SEEDS=(0,1,2,3)
 10# The union is shared by baseline and idea; no hidden hyperparameter is used.
 11GRID=[{'lr':1e-3},{'lr':3e-3},{'lr':1e-2}]
 12EPOCHS=5; NTRAIN=400; NTEST=160
 13
 14class AveragedContractiveRNN(nn.Module):
 15    """Discrete Euler sample of h'=F(t/eps,h,x), using quadrature over phase.
 16    The recurrent matrix is spectrally bounded and the linear drift is strictly
 17    contractive; phase averaging is the sole architectural intervention.
 18    """
 19    def __init__(self, hidden=16, phases=3, alpha=2.0, q=.30, dt=.05):
 20        super().__init__(); self.hidden=hidden; self.phases=phases
 21        self.alpha=alpha; self.q=q; self.dt=dt
 22        self.w_raw=nn.Parameter(torch.randn(hidden,hidden)*.05)
 23        self.inp=nn.Linear(3,hidden); self.bias=nn.Parameter(torch.zeros(hidden))
 24        self.head=nn.Linear(hidden,1)
 25        pattern=torch.ones(hidden); pattern[1::2]=-1
 26        self.register_buffer('pattern',pattern)
 27    def w_bound(self):
 28        # A differentiable, uniform spectral bound, so tanh Jacobian <= ||W||.
 29        return self.w_raw / (torch.linalg.matrix_norm(self.w_raw,2)+1e-6) * .35
 30    def vector_field(self,h,x,phase):
 31        W=self.w_bound(); a=-self.alpha + self.q*torch.sin(2*torch.pi*torch.as_tensor(phase, device=h.device, dtype=h.dtype))*self.pattern
 32        return a*h + torch.tanh(h@W.T + self.inp(x) + self.bias)
 33    def rollout(self,x, averaged=True, eps=1/8, return_states=False):
 34        z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device)
 35        states=[]
 36        for k in range(z.shape[1]):
 37            if averaged:
 38                ps=torch.arange(self.phases,device=x.device,dtype=x.dtype)/self.phases
 39                f=sum(self.vector_field(h,z[:,k],p) for p in ps)/self.phases
 40            else:
 41                phase=(k*self.dt/eps)
 42                f=self.vector_field(h,z[:,k],phase)
 43            h=h+self.dt*f; states.append(h)
 44        out=self.head(h)
 45        return (out,torch.stack(states,1)) if return_states else out
 46    def forward(self,x): return self.rollout(x, averaged=True)
 47
 48def seed_all(seed):
 49    np.random.seed(seed); torch.manual_seed(seed)
 50    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
 51
 52def ds(seed): return get_dataset(TRACK,seed,n_train=NTRAIN,n_test=NTEST)
 53
 54def baseline_fn(cfg):
 55    def run(seed):
 56        seed_all(seed); d=ds(seed); m=make_model(MODEL,d['input_shape'],d['out_dim'])
 57        _,metric,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=128)
 58        return metric
 59    return run
 60
 61def idea_fn(cfg, retain=False):
 62    def run(seed):
 63        seed_all(seed); d=ds(seed); m=AveragedContractiveRNN()
 64        _,metric,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=128)
 65        return metric
 66    return run
 67
 68def signature():
 69    # Re-test the stage-1 prediction on trained benchmark models, not a toy graph.
 70    seed=0; seed_all(seed); d=ds(seed); m=AveragedContractiveRNN()
 71    _,_,_=train_model(m,d,epochs=EPOCHS,lr=3e-3,batch=128)
 72    device=next(m.parameters()).device; m.eval(); x=d['xte'][:96].to(device)
 73    rows=[]
 74    with torch.no_grad():
 75        _,ha=m.rollout(x,averaged=True,return_states=True)
 76        for eps in [0.5,.25,.125,.0625]:
 77            _,hf=m.rollout(x,averaged=False,eps=eps,return_states=True)
 78            e=torch.sqrt(((hf-ha)**2).mean()).item()
 79            rows.append({'eps':eps,'rms_state_error':e})
 80    slope=float(np.polyfit(np.log([r['eps'] for r in rows]),np.log([r['rms_state_error']+1e-12 for r in rows]),1)[0])
 81    # Empirical matrix measure of the trained field at sampled hidden/input points.
 82    mus=[]
 83    for i in range(12):
 84        h=torch.randn(1,m.hidden,device=device,requires_grad=True); u=x[i:i+1].view(1,-1,3)[:,0]
 85        phase=float(i)/12
 86        J=torch.autograd.functional.jacobian(lambda hh:m.vector_field(hh,u,phase),h).squeeze(0).squeeze(1)
 87        mus.append(float(torch.linalg.eigvalsh((J+J.T)/2).max().detach().cpu()))
 88    max_mu=max(mus); predicted_bound=-m.alpha+m.q+.35
 89    return {'prediction':'trained averaged/fast state error decreases with eps; mu2 is negative',
 90            'rows':rows,'observed_loglog_slope':slope,'predicted_error_slope':1.0,
 91            'predicted_mu2_upper_bound':predicted_bound,'observed_max_mu2':max_mu,
 92            'confirmed':bool(slope>0.5 and max_mu<0)}
 93
 94def main():
 95    t=time.time()
 96    # Baseline sweep on four seeds, then canonical full eight-seed reevaluation.
 97    base=sweep_baseline(baseline_fn,GRID,seeds=SWEEP_SEEDS)
 98    # Explicitly evaluate idea at best baseline lr and two nearby/shared settings.
 99    idea_trials=[]
100    for cfg in GRID:
101        r=__import__('bench').evaluate(idea_fn(cfg),seeds=SEEDS)
102        idea_trials.append({'cfg':cfg,'result':r})
103    best=min(idea_trials,key=lambda z:z['result']['mean'])
104    rep=make_report(TRACK,MODEL,base,best['result'],signature())
105    rep['idea_sweep']=idea_trials; rep['runtime_seconds']=time.time()-t
106    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
107    print(json.dumps(rep,indent=2))
108if __name__=='__main__': main()