Thermodynamic Two-State Expert Gate / experiment.py
Mechanism failed
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()