Global-statistics context shortcut / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, os, random, time
2import numpy as np
3import torch
4import torch.nn as nn
5from torch.utils.data import DataLoader, TensorDataset
6
7SEED = 2536
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(min(12, os.cpu_count() or 1))
10DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
11try:
12 if DEVICE == "cuda": torch.zeros(1, device=DEVICE)
13except Exception:
14 DEVICE = "cpu"
15
16class GlobalNorm(nn.Module):
17 def __init__(self, channels, eps=1e-8):
18 super().__init__(); self.gamma=nn.Parameter(torch.ones(channels)); self.beta=nn.Parameter(torch.zeros(channels)); self.eps=eps
19 def forward(self, x, mask=None):
20 if mask is None: mask=torch.ones(x.shape[:2],device=x.device,dtype=x.dtype)
21 m=mask[...,None]; n=m.sum(1,keepdim=True).clamp_min(1)
22 mu=(m*x).sum(1,keepdim=True)/n; var=(m*(x-mu).square()).sum(1,keepdim=True)/n
23 return ((x-mu)/torch.sqrt(var+self.eps)*self.gamma+self.beta)*m
24
25def math_checks():
26 torch.manual_seed(SEED+1); n,c=37,2
27 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
28 mu=x.mean(0); var=((x-mu)**2).mean(0); sig=torch.sqrt(var+eps); z=gamma*(x-mu)/sig+beta
29 t,s,ch=11,29,0
30 jac=torch.autograd.functional.jacobian(lambda q: (gamma*(q-q.mean(0))/torch.sqrt(((q-q.mean(0))**2).mean(0)+eps)+beta),x)
31 empirical=jac[t,ch,s,ch].item(); hat=(x.detach()-mu.detach())/sig.detach()
32 pred=(gamma[ch]/sig[ch]*(-(1+hat[t,ch]*hat[s,ch])/n)).item()
33 diagonal_pred=(gamma[ch]/sig[ch]*(1-(1+hat[t,ch]**2)/n)).item()
34 off_err=abs(empirical-pred); diag_err=abs(jac[t,ch,t,ch].item()-diagonal_pred)
35 # Predictions: off-diagonal magnitude scales as |gamma| and approximately 1/n; zero gamma removes influence.
36 scaling=[]
37 # Deterministic alternating +/-1 inputs: sigma=1 and two same-sign positions
38 # have hat{x}_t hat{x}_s=1, so the prediction is exactly 2/n.
39 for nn in [16,32,64,128]:
40 xx=torch.tensor([1.0 if i%2==0 else -1.0 for i in range(nn)],dtype=torch.float64)[:,None]
41 mm=xx.mean(); ss=torch.sqrt(((xx-mm)**2).mean()+eps); hh=(xx-mm)/ss
42 coeff=abs((-1.0/ss*(1+hh[0,0]*hh[nn-2,0])/nn).item()); scaling.append((nn,coeff,coeff*nn))
43 gammas=[0.0,0.5,1.0,2.0]; gvals=[]
44 xx=torch.randn(64,1,dtype=torch.float64); mm=xx.mean(); ss=torch.sqrt(((xx-mm)**2).mean()+eps); hh=(xx-mm)/ss
45 base=abs((-1/ss*(1+hh[3,0]*hh[50,0])/64).item())
46 for g in gammas: gvals.append((g,g*base))
47 return {'offdiag_abs_error':off_err,'diag_abs_error':diag_err,'offdiag_empirical':empirical,'offdiag_predicted':pred,'n_sweep':scaling,'gamma_sweep':gvals}
48
49class Model(nn.Module):
50 def __init__(self, width=16, kernel=9, norm=False):
51 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)
52 def forward(self,x):
53 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)
54
55def data(n, run, seed):
56 g=np.random.default_rng(seed); X=np.zeros((n,256,2),np.float32); Y=np.zeros((n,256),np.int64)
57 for i in range(n):
58 labels=np.zeros(256,np.int64); p=0; v=int(g.integers(0,2))
59 while p<256:
60 r=max(1,int(g.geometric(1/run))); labels[p:min(256,p+r)]=v; p+=r; v=int(g.integers(0,2))
61 X[i,:,0]=labels + g.normal(0,.8,256); X[i,:,1]=g.normal(0,1,256); Y[i]=labels
62 return torch.tensor(X),torch.tensor(Y)
63
64def train_eval(norm, kernel, run, steps=180):
65 torch.manual_seed(SEED+run+kernel+int(norm)); x,y=data(384,run,SEED+run); xv,yv=data(128,run,SEED+900+run)
66 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)
67 it=iter(ds); t0=time.perf_counter(); model.train()
68 for _ in range(steps):
69 try: xb,yb=next(it)
70 except StopIteration: it=iter(ds); xb,yb=next(it)
71 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()
72 elapsed=time.perf_counter()-t0; model.eval()
73 with torch.no_grad():
74 pred=model(xv.to(DEVICE)).argmax(-1).cpu(); acc=(pred==yv).float().mean().item()
75 return acc,elapsed, sum(p.numel() for p in model.parameters())
76
77def main():
78 global DEVICE
79 checks=math_checks(); results=[]
80 try:
81 for run in [2,8,32]:
82 local=train_eval(False,9,run); shortcut=train_eval(True,9,run); wide=train_eval(False,33,run)
83 results.append({'run_length':run,'local':local,'global_norm':shortcut,'wide_context':wide})
84 except Exception as exc:
85 if DEVICE != 'cuda': raise
86 print('CUDA failed; falling back to CPU:', repr(exc))
87 DEVICE='cpu'; results=[]
88 for run in [2,8,32]:
89 local=train_eval(False,9,run); shortcut=train_eval(True,9,run); wide=train_eval(False,33,run)
90 results.append({'run_length':run,'local':local,'global_norm':shortcut,'wide_context':wide})
91 out={'device':DEVICE,'math':checks,'benchmark':results}
92 with open('results.json','w') as f: json.dump(out,f,indent=2)
93 print(json.dumps(out,indent=2))
94if __name__=='__main__': main()