import json, math, os, random, time import numpy as np import torch import torch.nn as nn from torch.utils.data import DataLoader, TensorDataset SEED = 2536 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_num_threads(min(12, os.cpu_count() or 1)) DEVICE = "cuda" if torch.cuda.is_available() else "cpu" try: if DEVICE == "cuda": torch.zeros(1, device=DEVICE) except Exception: DEVICE = "cpu" class GlobalNorm(nn.Module): def __init__(self, channels, eps=1e-8): super().__init__(); self.gamma=nn.Parameter(torch.ones(channels)); self.beta=nn.Parameter(torch.zeros(channels)); self.eps=eps def forward(self, x, mask=None): if mask is None: mask=torch.ones(x.shape[:2],device=x.device,dtype=x.dtype) m=mask[...,None]; n=m.sum(1,keepdim=True).clamp_min(1) mu=(m*x).sum(1,keepdim=True)/n; var=(m*(x-mu).square()).sum(1,keepdim=True)/n return ((x-mu)/torch.sqrt(var+self.eps)*self.gamma+self.beta)*m def math_checks(): torch.manual_seed(SEED+1); n,c=37,2 x=torch.randn(n,c,dtype=torch.float64,requires_grad=True); gamma=torch.tensor([1.7,-.8],dtype=torch.float64); beta=torch.tensor([.2,-.1],dtype=torch.float64); eps=1e-12 mu=x.mean(0); var=((x-mu)**2).mean(0); sig=torch.sqrt(var+eps); z=gamma*(x-mu)/sig+beta t,s,ch=11,29,0 jac=torch.autograd.functional.jacobian(lambda q: (gamma*(q-q.mean(0))/torch.sqrt(((q-q.mean(0))**2).mean(0)+eps)+beta),x) empirical=jac[t,ch,s,ch].item(); hat=(x.detach()-mu.detach())/sig.detach() pred=(gamma[ch]/sig[ch]*(-(1+hat[t,ch]*hat[s,ch])/n)).item() diagonal_pred=(gamma[ch]/sig[ch]*(1-(1+hat[t,ch]**2)/n)).item() off_err=abs(empirical-pred); diag_err=abs(jac[t,ch,t,ch].item()-diagonal_pred) # Predictions: off-diagonal magnitude scales as |gamma| and approximately 1/n; zero gamma removes influence. scaling=[] # Deterministic alternating +/-1 inputs: sigma=1 and two same-sign positions # have hat{x}_t hat{x}_s=1, so the prediction is exactly 2/n. for nn in [16,32,64,128]: xx=torch.tensor([1.0 if i%2==0 else -1.0 for i in range(nn)],dtype=torch.float64)[:,None] mm=xx.mean(); ss=torch.sqrt(((xx-mm)**2).mean()+eps); hh=(xx-mm)/ss coeff=abs((-1.0/ss*(1+hh[0,0]*hh[nn-2,0])/nn).item()); scaling.append((nn,coeff,coeff*nn)) gammas=[0.0,0.5,1.0,2.0]; gvals=[] xx=torch.randn(64,1,dtype=torch.float64); mm=xx.mean(); ss=torch.sqrt(((xx-mm)**2).mean()+eps); hh=(xx-mm)/ss base=abs((-1/ss*(1+hh[3,0]*hh[50,0])/64).item()) for g in gammas: gvals.append((g,g*base)) return {'offdiag_abs_error':off_err,'diag_abs_error':diag_err,'offdiag_empirical':empirical,'offdiag_predicted':pred,'n_sweep':scaling,'gamma_sweep':gvals} class Model(nn.Module): def __init__(self, width=16, kernel=9, norm=False): super().__init__(); self.conv=nn.Conv1d(2,width,kernel,padding=kernel//2); self.norm=GlobalNorm(width) if norm else nn.Identity(); self.head=nn.Conv1d(width,2,1) def forward(self,x): y=self.conv(x.transpose(1,2)).transpose(1,2); y=torch.relu(y); y=self.norm(y); return self.head(y.transpose(1,2)).transpose(1,2) def data(n, run, seed): g=np.random.default_rng(seed); X=np.zeros((n,256,2),np.float32); Y=np.zeros((n,256),np.int64) for i in range(n): labels=np.zeros(256,np.int64); p=0; v=int(g.integers(0,2)) while p<256: r=max(1,int(g.geometric(1/run))); labels[p:min(256,p+r)]=v; p+=r; v=int(g.integers(0,2)) X[i,:,0]=labels + g.normal(0,.8,256); X[i,:,1]=g.normal(0,1,256); Y[i]=labels return torch.tensor(X),torch.tensor(Y) def train_eval(norm, kernel, run, steps=180): torch.manual_seed(SEED+run+kernel+int(norm)); x,y=data(384,run,SEED+run); xv,yv=data(128,run,SEED+900+run) model=Model(16,kernel,norm).to(DEVICE); opt=torch.optim.Adam(model.parameters(),lr=3e-3); ds=DataLoader(TensorDataset(x,y),batch_size=64,shuffle=True) it=iter(ds); t0=time.perf_counter(); model.train() for _ in range(steps): try: xb,yb=next(it) except StopIteration: it=iter(ds); xb,yb=next(it) xb,yb=xb.to(DEVICE),yb.to(DEVICE); loss=nn.functional.cross_entropy(model(xb).reshape(-1,2),yb.reshape(-1)); opt.zero_grad(); loss.backward(); opt.step() elapsed=time.perf_counter()-t0; model.eval() with torch.no_grad(): pred=model(xv.to(DEVICE)).argmax(-1).cpu(); acc=(pred==yv).float().mean().item() return acc,elapsed, sum(p.numel() for p in model.parameters()) def main(): global DEVICE checks=math_checks(); results=[] try: for run in [2,8,32]: local=train_eval(False,9,run); shortcut=train_eval(True,9,run); wide=train_eval(False,33,run) results.append({'run_length':run,'local':local,'global_norm':shortcut,'wide_context':wide}) except Exception as exc: if DEVICE != 'cuda': raise print('CUDA failed; falling back to CPU:', repr(exc)) DEVICE='cpu'; results=[] for run in [2,8,32]: local=train_eval(False,9,run); shortcut=train_eval(True,9,run); wide=train_eval(False,33,run) results.append({'run_length':run,'local':local,'global_norm':shortcut,'wide_context':wide}) out={'device':DEVICE,'math':checks,'benchmark':results} with open('results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps(out,indent=2)) if __name__=='__main__': main()