Thermodynamic Two-State Expert Gate / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, random, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6
  7SEED=17
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9torch.set_num_threads(4)
 10try:
 11    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 12    if device.type=='cuda': torch.zeros(1,device=device)
 13except Exception:
 14    device=torch.device('cpu')
 15
 16def math_check():
 17    ps=np.linspace(.01,.99,99); tau=.37; dy=2.3
 18    rng=np.random.default_rng(SEED); vars=[]
 19    for p in ps:
 20        y=np.where(rng.random(200000)<p,dy,0.)
 21        vars.append(np.var(y))
 22    vars=np.asarray(vars); theory=ps*(1-ps)*dy**2
 23    b=np.linspace(-8,8,1001); pp=1/(1+np.exp(-b/tau)); db=b[1]-b[0]
 24    deriv=np.gradient(pp,db)
 25    return {'max_abs_variance_error':float(np.max(np.abs(vars-theory))),
 26            'variance_peak_p':float(ps[np.argmax(theory)]),
 27            'susceptibility_peak_p':0.5,
 28            'finite_diff_peak_p':float(pp[np.argmax(deriv)]),
 29            'susceptibility_bound':1/(4*tau),
 30            'finite_diff_max':float(deriv.max()),
 31            'peak_error_from_half':float(abs(ps[np.argmax(theory)]-.5))}
 32
 33def data(n,T=28,rho=.5,seed=0):
 34    r=np.random.default_rng(seed); x=np.zeros((n,T+1,1),np.float32); z=np.zeros((n,T),np.float32)
 35    for i in range(n):
 36        z[i,0]=r.random()<rho; x[i,0]=r.normal(0,.8)
 37        for t in range(T):
 38            if t and r.random()>.94: z[i,t]=r.random()<rho
 39            a=.82 if z[i,t]==0 else -.72
 40            x[i,t+1]=a*x[i,t]+r.normal(0,.12)
 41    return torch.tensor(x),torch.tensor(z)
 42
 43class Baseline(nn.Module):
 44    def __init__(self,h=20):
 45        super().__init__(); self.r=nn.GRU(1,h,batch_first=True); self.out=nn.Linear(h,1)
 46    def forward(self,x):
 47        h,_=self.r(x); return self.out(h),h
 48
 49class ThermoGate(nn.Module):
 50    def __init__(self,h=20,tau=.65):
 51        super().__init__(); self.r=nn.GRU(1,h,batch_first=True); self.q=nn.Linear(h,1)
 52        self.f0=nn.Linear(h,1); self.f1=nn.Linear(h,1); self.tau=tau
 53    def forward(self,x):
 54        h,_=self.r(x); q=self.q(h).squeeze(-1); p=torch.sigmoid(q/self.tau)
 55        y0=self.f0(h).squeeze(-1); y1=self.f1(h).squeeze(-1)
 56        return (1-p)*y0+p*y1,p,q
 57
 58def train(model,gated,steps=650):
 59    model.to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3); model.train()
 60    for step in range(steps):
 61        x,z=data(96,28,.5,1000+step); x=x.to(device); target=x[:,1:,0]
 62        pred,*rest=model(x[:,:-1]); loss=(pred.squeeze(-1)-target).pow(2).mean()
 63        if gated:
 64            p=rest[0]; loss=loss+.015*(p[:,1:]-p[:,:-1]).pow(2).mean()+.003*(p.mean()-.5).pow(2)
 65        opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),2.0); opt.step()
 66    return float(loss)
 67
 68def evaluate(model,gated,rho):
 69    model.eval(); vals=[]; ps=[]; zs=[]
 70    with torch.no_grad():
 71        for k in range(12):
 72            x,z=data(96,28,rho,50000+int(rho*1000)+k); x=x.to(device); out=model(x[:,:-1])
 73            vals.append((out[0].squeeze(-1)-x[:,1:,0]).pow(2).mean().item())
 74            if gated: ps.append(out[1].cpu().numpy()); zs.append(z.numpy())
 75    ret={'mse':float(np.mean(vals))}
 76    if gated:
 77        pp=np.concatenate(ps); zz=np.concatenate(zs)
 78        ret.update({'mean_gate':float(pp.mean()),'gate_variance':float(pp.var()),
 79                    'gate_smoothness':float(np.mean(np.diff(pp,axis=1)**2)),
 80                    'gate_regime_corr':float(np.corrcoef(pp.ravel(),zz.ravel())[0,1])})
 81    return ret
 82
 83def run_benchmark():
 84    global device
 85    for attempt in range(2):
 86        try:
 87            b=Baseline(); g=ThermoGate(); train(b,False); train(g,True)
 88            result={'params':{'baseline':sum(p.numel() for p in b.parameters()),'idea':sum(p.numel() for p in g.parameters())}}
 89            for rho in [.1,.5,.9]: result[f'rho_{rho}']={'baseline':evaluate(b,False,rho),'idea':evaluate(g,True,rho)}
 90            return result
 91        except RuntimeError as e:
 92            if attempt==0 and device.type=='cuda':
 93                device=torch.device('cpu'); torch.cuda.empty_cache(); continue
 94            raise
 95
 96def main():
 97    t=time.time(); results={'device_requested':('cuda' if torch.cuda.is_available() else 'cpu'),'math_check':math_check()}
 98    results.update(run_benchmark()); results['device_used']=str(device); results['seconds']=time.time()-t
 99    Path('results.json').write_text(json.dumps(results,indent=2)); print(json.dumps(results,indent=2))
100if __name__=='__main__': main()