Cross-Channel Scattering Front End / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5import torch.nn.functional as F
6
7SEED = 17
8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
9torch.set_num_threads(4)
10DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
11
12class CrossChannelScattering(nn.Module):
13 """Fixed complex Morlet analysis, cross-channel amplitude products, averaging."""
14 def __init__(self, channels, edges, bands=(0.10, 0.20), kernel=25, pool=32, eps=1e-8):
15 super().__init__()
16 self.channels, self.edges, self.bands = channels, edges, tuple(bands)
17 self.pool, self.eps = pool, eps
18 t = torch.arange(kernel, dtype=torch.float32) - (kernel-1)/2
19 sigma = kernel / 6.0
20 ws = []
21 for f in bands:
22 g = torch.exp(-0.5*(t/sigma)**2)
23 # zero-mean approximately analytic Morlet pair
24 ws.append(torch.stack([g*torch.cos(2*math.pi*f*t), g*torch.sin(2*math.pi*f*t)]))
25 w = torch.stack(ws) # bands, 2, kernel
26 self.register_buffer('wr', w[:,0].view(len(bands),1,kernel))
27 self.register_buffer('wi', w[:,1].view(len(bands),1,kernel))
28
29 def forward(self, x):
30 B,C,T = x.shape
31 xr = x.reshape(B*C,1,T)
32 pad = self.wr.shape[-1]//2
33 zr = F.conv1d(xr, self.wr, padding=pad).reshape(B,C,len(self.bands),T)
34 zi = F.conv1d(xr, self.wi, padding=pad).reshape(B,C,len(self.bands),T)
35 out=[]
36 for m,n in self.edges:
37 # q = (zr_m+i zi_m) conj(zr_n+i zi_n)
38 qr = zr[:,m]*zr[:,n] + zi[:,m]*zi[:,n]
39 qi = zi[:,m]*zr[:,n] - zr[:,m]*zi[:,n]
40 u = torch.sqrt(qr.square()+qi.square()+self.eps)
41 out.append(F.avg_pool1d(u.reshape(B,len(self.bands),T), self.pool, stride=self.pool))
42 return torch.cat(out, dim=1)
43
44def verify():
45 torch.manual_seed(SEED)
46 z1 = torch.randn(3, 2, 20) + 1j*torch.randn(3,2,20)
47 z2 = torch.randn(3, 2, 20) + 1j*torch.randn(3,2,20)
48 q=z1*z2.conj()
49 identity=(q.abs()-(z1.abs()*z2.abs())).abs().max().item()
50 phase=torch.rand(3,2,20)*2*math.pi
51 phase_err=((z1*torch.exp(1j*phase)).abs()-z1.abs()).abs().max().item()
52 # end-to-end layer agrees with direct magnitude product (apart from eps)
53 layer=CrossChannelScattering(2,[(0,1)],bands=(.12,),kernel=17,pool=8)
54 x=torch.randn(2,2,64)
55 got=layer(x)
56 # independently reproduce its convolution and pooling
57 pad=8; zr=F.conv1d(x.reshape(4,1,64),layer.wr,padding=pad).reshape(2,2,1,64)
58 zi=F.conv1d(x.reshape(4,1,64),layer.wi,padding=pad).reshape(2,2,1,64)
59 direct=F.avg_pool1d(torch.sqrt((zr[:,0]*zr[:,1]+zi[:,0]*zi[:,1]).square()+(zi[:,0]*zr[:,1]-zr[:,0]*zi[:,1]).square()+1e-8).reshape(2,1,64),8,stride=8)
60 implementation_err=(got-direct).abs().max().item()
61 return {'complex_magnitude_identity_max_error':identity,'phase_invariance_max_error':phase_err,'layer_direct_formula_max_error':implementation_err}
62
63def make_data(n, T=256):
64 # Class 1: neighboring sensors share a slowly varying amplitude envelope;
65 # phases are independent, so coupling is amplitude- rather than phase-based.
66 t=np.arange(T)/T; X=np.zeros((n,4,T),np.float32); y=np.arange(n)%2
67 for i,c in enumerate(y):
68 freq=.12 + .008*np.random.randn()
69 envs=[]
70 if c==1:
71 env=0.55+0.45*(np.sin(2*np.pi*(1.5*t+np.random.rand()))**2)
72 envs=[env,env, .6+.4*np.random.rand()*np.ones(T), .6+.4*np.random.rand()*np.ones(T)]
73 else:
74 envs=[0.55+0.45*np.sin(2*np.pi*(1.5*t+np.random.rand()))**2 for _ in range(4)]
75 for ch in range(4):
76 ph=np.random.rand()*2*np.pi
77 X[i,ch]=envs[ch]*np.sin(2*np.pi*freq*np.arange(T)+ph)+.65*np.random.randn(T)
78 return torch.tensor(X),torch.tensor(y,dtype=torch.long)
79
80class RawCNN(nn.Module):
81 def __init__(self):
82 super().__init__(); self.net=nn.Sequential(nn.Conv1d(4,12,15,padding=7),nn.ReLU(),nn.AvgPool1d(8),nn.Conv1d(12,12,9,padding=4),nn.ReLU(),nn.AdaptiveAvgPool1d(1))
83 self.fc=nn.Linear(12,2)
84 def forward(self,x): return self.fc(self.net(x).squeeze(-1))
85class SNSTModel(nn.Module):
86 def __init__(self):
87 super().__init__(); self.s=CrossChannelScattering(4,[(0,1),(1,2),(2,3)],bands=(.10,.20),kernel=25,pool=32); self.fc=nn.Sequential(nn.Flatten(),nn.Linear(3*2*8,16),nn.ReLU(),nn.Linear(16,2))
88 def forward(self,x): return self.fc(self.s(x))
89
90def train(model, Xtr,ytr,Xte,yte, epochs=35):
91 model=model.to(DEVICE); Xtr=Xtr.to(DEVICE); ytr=ytr.to(DEVICE); Xte=Xte.to(DEVICE); yte=yte.to(DEVICE)
92 opt=torch.optim.Adam(model.parameters(),lr=3e-3,weight_decay=1e-3)
93 for _ in range(epochs):
94 model.train(); p=torch.randperm(len(Xtr),device=DEVICE)
95 for j in range(0,len(Xtr),64):
96 ix=p[j:j+64]; loss=F.cross_entropy(model(Xtr[ix]),ytr[ix]); opt.zero_grad(); loss.backward(); opt.step()
97 model.eval()
98 with torch.no_grad():
99 pred=model(Xte).argmax(1); acc=(pred==yte).float().mean().item()
100 return acc
101
102def experiment():
103 # 20% labels: deliberately small training set, fixed held-out test.
104 X,y=make_data(800); Xtr,ytr=X[:160],y[:160]; Xte,yte=X[160:],y[160:]
105 # normalize using training statistics, as required for a fair front end
106 mu=Xtr.mean((0,2),keepdim=True); sd=Xtr.std((0,2),keepdim=True)+1e-5
107 Xtr=(Xtr-mu)/sd; Xte=(Xte-mu)/sd
108 results=[]
109 for kind in ('raw','snst'):
110 vals=[]
111 for seed in (3,7,11):
112 torch.manual_seed(seed); np.random.seed(seed)
113 model=RawCNN() if kind=='raw' else SNSTModel()
114 vals.append(train(model,Xtr,ytr,Xte,yte))
115 results.append((kind,vals,float(np.mean(vals)),float(np.std(vals))))
116 return results
117
118if __name__=='__main__':
119 try:
120 if DEVICE=='cuda': torch.cuda.init()
121 except Exception as e:
122 DEVICE='cpu'; print('CUDA fallback:',repr(e))
123 check=verify(); print('DEVICE',DEVICE); print('VERIFY',json.dumps(check))
124 try:
125 results=experiment()
126 except Exception as e:
127 if DEVICE != 'cpu':
128 print('Runtime fallback to CPU:',repr(e))
129 DEVICE='cpu'
130 results=experiment()
131 else:
132 raise
133 print('RESULTS',json.dumps(results))