import json import numpy as np import torch from torch import nn SEED=1273 N=32 DEPTH=16 DT=.08 J=1.0 def device(): return 'cuda' if torch.cuda.is_available() else 'cpu' class PropNet(nn.Module): def __init__(self, mode, gamma=.5, target=.15, alpha=.5): super().__init__(); self.mode=mode; self.gamma0=gamma; self.target=target; self.alpha=alpha # two real fields implement damped coherent propagation; parameter count is identical. self.readout=nn.Sequential(nn.Linear(2*N,16),nn.Tanh(),nn.Linear(16,2)) def forward(self,x, return_stats=False): # x: batch,N; q and p are the propagated fields. q=x; p=torch.zeros_like(x); gam=torch.full((x.shape[0],),self.gamma0,device=x.device) gs=[]; rs=[] for _ in range(DEPTH): lapq=torch.roll(q,1,1)+torch.roll(q,-1,1)-2*q lapp=torch.roll(p,1,1)+torch.roll(p,-1,1)-2*p q=q+DT*J*lapp p=p-DT*J*lapq-DT*gam[:,None]*p # fixed controller is detached from the optimization graph. corr=(q[:,:-1]*q[:,1:]).mean(1).abs() var=(q*q).mean(1)+1e-6 r=corr/var if self.mode=='adaptive': gam=torch.clamp(gam*torch.exp(self.alpha*(r.detach()-self.target)),.02,16.) gs.append(gam.mean().item()); rs.append(r.mean().item()) logits=self.readout(torch.cat([q,p],1)) if return_stats: return logits, float(np.mean(gs)), float(np.mean(rs)), float(gam.mean()) return logits def batch(n, dev): # Position of a pulse encodes the class; distractor noise stresses propagation. y=torch.randint(0,2,(n,),device=dev); x=.10*torch.randn(n,N,device=dev) pos=torch.where(y==0, torch.randint(3,N//2,(n,),device=dev), torch.randint(N//2,N-3,(n,),device=dev)) x[torch.arange(n,device=dev),pos]=1. return x,y def run(mode,gamma): torch.manual_seed(SEED); np.random.seed(SEED); dev=device() try: net=PropNet(mode,gamma).to(dev); opt=torch.optim.Adam(net.parameters(),lr=3e-3); lossfn=nn.CrossEntropyLoss() for _ in range(160): x,y=batch(96,dev); loss=lossfn(net(x),y); opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): x,y=batch(512,dev); logits,g,r,gf=net(x,True); acc=(logits.argmax(1)==y).float().mean().item(); loss=lossfn(logits,y).item() return {'loss':loss,'accuracy':acc,'mean_gamma':g,'mean_ratio':r,'final_gamma':gf,'device':dev} except Exception as e: if dev=='cuda': torch.cuda.empty_cache(); # CPU fallback is explicit and reproducible. torch.set_default_device('cpu'); return run(mode,gamma) raise def main(): out={'fixed_ballistic':run('fixed',.5),'fixed_diffusive':run('fixed',4.),'adaptive':run('adaptive',.5)} open('mini_results.json','w').write(json.dumps(out,indent=2)); print(json.dumps(out,indent=2)) if __name__=='__main__': main()