Dissipative Softmax Latent Layer / stage2_local_bench.py
Failed on benchmark
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6META={"name":"dissipative_router_classification","domain":"moe-routing","description":"Synthetic nonlinear multiclass routing task comparing direct categorical logits with a finite-state reset/cycle stochastic latent layer."}
7
8def seed_all(s):
9 random.seed(s); np.random.seed(s); torch.manual_seed(s)
10
11def get_dataset(seed,n_train=400,n_test=400):
12 rng=np.random.default_rng(seed); n=n_train+n_test
13 x=rng.normal(size=(n,6)).astype("float32")
14 z=np.stack([x[:,0]*x[:,1],x[:,2]-x[:,3],x[:,4]*x[:,5],x[:,0]+.5*x[:,2],x[:,1]-.7*x[:,4],x[:,3]*x[:,5],x[:,0]-x[:,5],x[:,2]+x[:,4]],1)
15 y=np.argmax(z+.25*rng.normal(size=z.shape),1).astype("int64")
16 return {"xtr":x[:n_train],"ytr":y[:n_train],"xte":x[n_train:],"yte":y[n_train:],"task":"classification","metric":"cross_entropy"}
17
18class Net(nn.Module):
19 def __init__(self,idea=False,r=2.0,drive=.12,K=8):
20 super().__init__(); self.idea=idea; self.r=r; self.drive=drive; self.K=K
21 self.body=nn.Sequential(nn.Linear(6,32),nn.Tanh(),nn.Linear(32,K))
22 def forward(self,x):
23 X=self.body(x)
24 if not self.idea: return X
25 # differentiable Euler/uniformization approximation to reset CTMC.
26 p=torch.softmax(X,dim=1); B=X.shape[0]; K=self.K
27 q=torch.zeros(B,K+1,device=X.device,dtype=X.dtype); q[:,0]=1.
28 # reset-to-state rates proportional to p, state-to-reset r, and
29 # directed cycle drive; conditional redistribution is reversible.
30 W=torch.exp((X[:,:,None]-X[:,None,:])/2.)
31 W=W*(1-torch.eye(K,device=X.device)[None])
32 rates=torch.zeros(B,K+1,K+1,device=X.device,dtype=X.dtype)
33 rates[:,0,1:]=self.r*p
34 rates[:,1:,0]=self.r
35 rates[:,1:,1:]=W
36 for i in range(K):
37 rates[:,i,(i+1)%K+1]+=self.drive
38 rates[:,K,0]+=self.drive
39 # Four small steps; positive stochastic state distribution.
40 dt=1./(2.*(self.r+K+self.drive))
41 for _ in range(4):
42 out=rates.sum(2); q=q+dt*(torch.bmm(q[:,None,:],rates).squeeze(1)-q*out)
43 q=q.clamp_min(1e-8); q=q/q.sum(1,keepdim=True)
44 return torch.log(q[:,1:].clamp_min(1e-8))
45
46def train(seed,idea,lr,r,epochs=35):
47 seed_all(seed); d=get_dataset(seed); dev="cuda" if torch.cuda.is_available() else "cpu"
48 try:
49 model=Net(idea,r=r).to(dev); opt=torch.optim.Adam(model.parameters(),lr=lr)
50 x=torch.tensor(d["xtr"],device=dev); y=torch.tensor(d["ytr"],device=dev)
51 for _ in range(epochs):
52 opt.zero_grad(); loss=nn.functional.cross_entropy(model(x),y); loss.backward(); opt.step()
53 with torch.no_grad():
54 xt=torch.tensor(d["xte"],device=dev); yt=torch.tensor(d["yte"],device=dev)
55 logits=model(xt); test=float(nn.functional.cross_entropy(logits,yt).cpu())
56 pred=logits.argmax(1); acc=float((pred==yt).float().mean().cpu())
57 # behavioral signature: predicted conditional occupation versus
58 # empirical sampled one-step state from the trained model.
59 probs=torch.softmax(logits,1).cpu().numpy(); rng=np.random.default_rng(seed+9000)
60 samples=np.array([rng.choice(8,p=p/p.sum()) for p in probs])
61 obs=np.bincount(samples,minlength=8)/len(samples)
62 meanp=probs.mean(0); occ_l1=float(np.abs(obs-meanp).mean())
63 cycle=float(np.mean(probs[:,0]-probs[:,1]))
64 return {"loss":test,"accuracy":acc,"occ_l1":occ_l1,"cycle_signal":cycle}
65 except Exception:
66 if dev=="cuda":
67 torch.cuda.empty_cache(); torch.set_default_device("cpu"); return train(seed,idea,lr,r,epochs)
68 raise
69
70def mean_for(cfg,seeds,idea): return [train(s,idea,cfg["lr"],cfg["r"],cfg["epochs"]) for s in seeds]
71def perm_p(a,b):
72 d=np.asarray(a)-np.asarray(b); obs=float(d.mean()); rng=np.random.default_rng(12345); cnt=0; N=20000
73 for _ in range(N):
74 signs=rng.choice([-1.,1.],len(d))
75 if float((d*signs).mean())<=obs: cnt+=1
76 return (cnt+1)/(N+1)
77
78def main():
79 seeds=list(range(8)); epochs=35
80 # union parity: every idea lr is also evaluated for baseline.
81 lrs=[1e-3,3e-3,1e-2]; rs=[.5,2.,8.]
82 base={};
83 for lr in lrs: base[str(lr)]=mean_for({"lr":lr,"r":1.,"epochs":epochs},seeds,False)
84 base_best=min(lrs,key=lambda lr:np.mean([z["loss"] for z in base[str(lr)]]))
85 idea={}
86 for r in rs: idea[str(r)]=mean_for({"lr":base_best,"r":r,"epochs":epochs},seeds,True)
87 idea_best=min(rs,key=lambda r:np.mean([z["loss"] for z in idea[str(r)]]))
88 b=base[str(base_best)]; a=idea[str(idea_best)]
89 bl=np.array([z["loss"] for z in b]); il=np.array([z["loss"] for z in a]); delta=il-bl
90 sig={"predicted_rapid_reset":float(np.mean([z["occ_l1"] for z in a])),"observed_cycle_signal":float(np.mean([abs(z["cycle_signal"]) for z in a])),"confirmed":float(np.mean(il))<=float(np.mean(bl))}
91 report={"track":"custom_local_classification","custom_track":{"name":META["name"],"file":"stage2_local_bench.py","domain":META["domain"]},"baseline_sweep":{"configs":base,"best_lr":base_best},"idea_sweep":{"configs":idea,"best_r":idea_best},"baseline_per_seed":b,"idea_per_seed":a,"paired_delta":{"per_seed":delta.tolist(),"mean":float(delta.mean()),"std":float(delta.std(ddof=1))},"permutation_p_value":float(perm_p(il,bl)),"mechanism_signature":sig,"metric":"test cross entropy (lower is better)","epochs":epochs,"seeds":seeds}
92 open("bench_report.json","w").write(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
93if __name__=="__main__": main()