Parabolic Riesz Feature Preconditioner / experiment.py
Mechanism confirmed, baseline not beaten
1import json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6
7SEED = 425
8np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
9try: torch.set_num_threads(4)
10except Exception: pass
11
12def ring_ops(n):
13 D = np.zeros((n,n), dtype=np.float64)
14 for i in range(n):
15 D[i,i] = -1.; D[i,(i+1)%n] = 1.
16 return D, D.T @ D
17
18def riesz_matrix(n, lam, eps=1e-6):
19 D,L = ring_ops(n)
20 A = np.eye(n) + lam*lam*(L + eps*np.eye(n))
21 return lam * D @ np.linalg.inv(A), D, L
22
23def sanity():
24 out=[]
25 for n in (32,64,128):
26 lam=1.0
27 R,D,L=riesz_matrix(n,lam)
28 direct=lam*D
29 rng=np.random.default_rng(SEED+n)
30 # noisy piecewise-smooth signal, measuring derivative noise amplification
31 x=np.sin(2*np.pi*np.arange(n)/n); x[n//2:]+=0.8
32 noise=rng.normal(size=n)
33 clean=np.linalg.norm(R@x); noisy=np.linalg.norm(R@(x+.5*noise))
34 dclean=np.linalg.norm(direct@x); dnoisy=np.linalg.norm(direct@(x+.5*noise))
35 op=np.linalg.svd(R,compute_uv=False)[0]
36 dop=np.linalg.svd(direct,compute_uv=False)[0]
37 # Isolate perturbation response, rather than conflating it with signal energy.
38 r_noise=np.linalg.norm(R@(0.5*noise))/np.linalg.norm(0.5*noise)
39 d_noise=np.linalg.norm(direct@(0.5*noise))/np.linalg.norm(0.5*noise)
40 out.append({'n':n,'resolvent_operator_norm':float(op),'direct_operator_norm':float(dop),
41 'noise_ratio_resolvent':float(noisy/clean),'noise_ratio_direct':float(dnoisy/dclean),
42 'noise_amplification_resolvent':float(r_noise),'noise_amplification_direct':float(d_noise)})
43 return out
44
45class Branch(nn.Module):
46 def __init__(self,n,d,kind):
47 super().__init__(); self.kind=kind; self.n=n
48 R,D,L=riesz_matrix(n,1.0)
49 M=R if kind=='riesz' else 1.0*D if kind=='direct' else np.zeros((n,n))
50 self.register_buffer('M',torch.tensor(M,dtype=torch.float32))
51 self.base=nn.Linear(d,16); self.out=nn.Linear(32 if kind!='plain' else 16,2)
52 self.g=nn.Parameter(torch.tensor(0.1 if kind=='riesz' else 1.0))
53 def forward(self,x):
54 # x: B,N,d; branch is applied to node dimension
55 h=torch.relu(self.base(x))
56 if self.kind!='plain':
57 r=torch.einsum('nm,bmd->bnd',self.M,h)
58 r=r/(r.square().mean(dim=(1,2),keepdim=True).sqrt()+1e-5)
59 h=h+self.g*r
60
61 pooled=h.mean(1)
62 if self.kind!='plain': pooled=torch.cat([pooled, r.abs().mean(1)],dim=1)
63 return self.out(pooled)
64
65def dataset(n=64, samples=1200, noise=.10):
66 rng=np.random.default_rng(SEED)
67 X=np.zeros((samples,n,1),np.float32); y=np.zeros(samples,np.int64)
68 t=np.arange(n)
69 for k in range(samples):
70 cls=k%2; y[k]=cls
71 phase=rng.uniform(0,2*np.pi)
72 if cls==0:
73 # Smooth oscillation and a random phase: same pointwise marginal as class 1.
74 sig=np.sin(2*np.pi*t/n*2+phase)
75 else:
76 # High-frequency alternating signal, with matched amplitude and random phase.
77 sig=np.sign(np.sin(2*np.pi*t/n*16+phase))
78 X[k,:,0]=sig+rng.normal(0,noise,n)
79 p=rng.permutation(samples); return X[p],y[p]
80
81def train(kind, Xtr,ytr,Xv,yv,n,epochs=35):
82 dev='cuda' if torch.cuda.is_available() else 'cpu'
83 try:
84 model=Branch(n,1,kind).to(dev)
85 opt=torch.optim.Adam(model.parameters(),lr=3e-3); lossfn=nn.CrossEntropyLoss()
86 xt=torch.tensor(Xtr,device=dev); yt=torch.tensor(ytr,device=dev)
87 xv=torch.tensor(Xv,device=dev); yvt=torch.tensor(yv,device=dev)
88 for _ in range(epochs):
89 for s in range(0,len(xt),128):
90 opt.zero_grad(); loss=lossfn(model(xt[s:s+128]),yt[s:s+128]); loss.backward(); opt.step()
91 with torch.no_grad():
92 acc=(model(xv).argmax(1)==yvt).float().mean().item()
93 pert=(model(xv+0.10*torch.randn_like(xv)).argmax(1)==yvt).float().mean().item()
94 val=lossfn(model(xv),yvt).item()
95 return {'val_loss':val,'accuracy':acc,'perturbed_accuracy':pert}
96 except Exception as e:
97 if dev=='cuda':
98 torch.cuda.empty_cache(); torch.set_default_device('cpu')
99 return train(kind,Xtr,ytr,Xv,yv,n,epochs)
100 raise
101
102def main():
103 check=sanity(); X,y=dataset(); Xtr,ytr,Xv,yv=X[:900],y[:900],X[900:],y[900:]
104 results={k:train(k,Xtr,ytr,Xv,yv,64) for k in ('plain','direct','riesz')}
105 report={'seed':SEED,'sanity':check,'results':results}
106 Path('results.json').write_text(json.dumps(report,indent=2))
107 print(json.dumps(report,indent=2))
108if __name__=='__main__': main()