Cross-Channel Scattering Front End / experiment.py

Mechanism confirmed, baseline not beaten

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