Shared Symbolic Mechanism Bottleneck / experiment.py

Mechanism failed

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