Shared Symbolic Mechanism Bottleneck / experiment.py
Mechanism failed
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED=1450
7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
8device='cuda' if torch.cuda.is_available() else 'cpu'
9try:
10 if device=='cuda': torch.cuda.empty_cache()
11except Exception: device='cpu'
12
13class SharedSymbolic(nn.Module):
14 # Safe differentiable implementation of the proposed feature-bank bottleneck.
15 def __init__(self,d=2,m=3,H=8,tau=1.0):
16 super().__init__(); self.d=d; self.m=m; self.H=H; self.tau=tau
17 self.w=nn.Parameter(torch.randn(H,6)*.15); self.b=nn.Parameter(torch.zeros(H))
18 self.glogit=nn.Parameter(torch.zeros(H,6)); self.oplogit=nn.Parameter(torch.randn(H,7)*.1)
19 self.a=nn.Parameter(torch.randn(m,H)*.15); self.v=nn.Parameter(torch.ones(m,H)*1.5)
20 self.c=nn.Parameter(torch.zeros(m)); self.register_buffer('mu',torch.tensor([1.,1.]))
21 self.register_buffer('sd',torch.tensor([1.,1.]))
22 def set_norm(self,X):
23 self.mu.copy_(X.mean(0)); self.sd.copy_(X.std(0).clamp_min(1e-5))
24 def features(self,X):
25 x=(X-self.mu)/self.sd
26 # positive-domain safe primitives, as in the implementation plan
27 raw=x; rec=1/x.clamp(-20,20).where(x.abs()>0.03, torch.sign(x)*.03)
28 log=torch.log(x.abs()+1e-2)
29 return torch.cat([raw,rec,log],1)
30 def forward(self,X):
31 ph=self.features(X); g=torch.sigmoid(self.glogit)
32 u=torch.einsum('bd,hd->bh',ph,g*self.w)+self.b
33 vals=torch.stack([torch.exp(u.clamp(-8,8)),u*u,u,1/u.clamp(-8,8),
34 torch.log1p(torch.relu(u)+1e-6),torch.sin(u),torch.sqrt(torch.relu(u)+1e-6)],-1)
35 p=torch.softmax(self.oplogit/self.tau,-1); z=(vals*p[None,:,:]).sum(-1)
36 q=torch.sigmoid(self.v)
37 return self.c+z@(q*self.a).T, z, p, g
38
39class MLP(nn.Module):
40 def __init__(self,d=2,m=3):
41 super().__init__(); self.net=nn.Sequential(nn.Linear(d,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh(),nn.Linear(32,m))
42 def forward(self,x): return self.net(x)
43
44def data(n, noise, rng, lo=.15, hi=5.):
45 X=rng.uniform(lo,hi,(n,2)).astype('float32'); K=np.array([.7,.35]); alpha=np.array([1.2,.8,1.5])
46 z=1/(1+X@K); Y=np.stack([alpha[j]*X[:,j%2]*z for j in range(3)],1)
47 if noise: Y += rng.normal(0,noise,Y.shape).astype('float32')
48 return torch.tensor(X),torch.tensor(Y.astype('float32'))
49
50def train(model,X,Y,epochs=500,lr=2e-3, symbolic=False):
51 model.to(device); X=X.to(device); Y=Y.to(device)
52 if symbolic: model.set_norm(X)
53 opt=torch.optim.Adam(model.parameters(),lr=lr)
54 for e in range(epochs):
55 opt.zero_grad(); pred=model(X)
56 if symbolic:
57 pred,z,p,g=pred; loss=((pred-Y)**2).mean()+2e-4*(torch.sigmoid(model.v)*model.a).abs().sum()+1e-5*g.sum()+1e-5*(p*(p+1e-8).log()).sum()
58 model.tau=max(.05,1-.95*e/epochs)
59 else: loss=((pred-Y)**2).mean()
60 loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),10); opt.step()
61 return model
62
63def mse(model,X,Y,symbolic=False):
64 with torch.no_grad():
65 out=model(X.to(device)); out=out[0] if symbolic else out
66 return float(((out-Y.to(device))**2).mean().cpu())
67
68def mechanism_checks():
69 # Prediction 1: exact two-class softmax concentration and entropy versus temperature.
70 gap=2.0; taus=[1.,.5,.25,.1,.05]; obs=[]; pred=[]
71 for t in taus:
72 q=np.exp(np.array([0.,gap])/t-np.max(np.array([0.,gap])/t)); q=q/q.sum(); h=float(-(q*np.log(np.maximum(q,1e-300))).sum()); obs.append(h); pred.append(float(math.log1p(math.exp(-gap/t))+(gap/t)/(1+math.exp(gap/t))))
73 # Prediction 2: rank-one shared latent mechanism gives zero residual for every output.
74 rng=np.random.default_rng(SEED); z=rng.uniform(.2,1.,1000); A=np.array([1.2,-.7,2.1]); Y=z[:,None]*A[None,:]
75 fit=np.linalg.lstsq(z[:,None],Y,rcond=None)[0]; rank_res=float(np.mean((Y-z[:,None]@fit)**2))
76 # Prediction 3: known additive noise has expected MSE sigma^2 (exact mechanism fit).
77 noise_rows=[]
78 for s in [0.,.01,.03,.1]:
79 vals=[]
80 for rep in range(30):
81 eps=rng.normal(0,s,(1000,3)); vals.append(float(np.mean(eps**2)))
82 noise_rows.append((s,float(np.mean(vals)),s*s))
83 return {'softmax_temperature':{'gap':gap,'taus':taus,'observed_entropy':obs,'predicted_entropy':pred},
84 'shared_latent_rank_residual':rank_res,'noise_scaling':noise_rows}
85
86def main():
87 checks=mechanism_checks(); rng=np.random.default_rng(SEED)
88 Xtr,Ytr=data(1000,.01,rng,.15,5.); Xte,Yte=data(1000,0,rng,5.,10.)
89 # symbolic model and standard shared-trunk MLP, same data and fixed seed
90 torch.manual_seed(SEED); sym=train(SharedSymbolic(H=8),Xtr,Ytr,500,2e-3,True)
91 torch.manual_seed(SEED); mlp=train(MLP(),Xtr,Ytr,500,2e-3,False)
92 with torch.no_grad(): out,z,p,g=sym(Xte.to(device)); active=int((torch.sigmoid(sym.v)*sym.a).abs().mean(0).gt(.03).sum().cpu()); ops=p.argmax(1).cpu().tolist()
93 result={'device':device,'checks':checks,'fit':{'symbolic_in_range_mse':mse(sym,Xtr,Ytr,True),'symbolic_extrapolation_mse':mse(sym,Xte,Yte,True),'mlp_in_range_mse':mse(mlp,Xtr,Ytr),'mlp_extrapolation_mse':mse(mlp,Xte,Yte),'symbolic_active_units':active,'operator_choices':ops}}
94 with open('results.json','w') as f: json.dump(result,f,indent=2)
95 print(json.dumps(result,indent=2))
96if __name__=='__main__': main()