Parabolic Riesz Feature Preconditioner / experiment.py

Mechanism confirmed, baseline not beaten

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