Dissipative Softmax Latent Layer / stage2_local_bench.py

Failed on benchmark

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