Sparse Multiscale Kernel-Frame Operator / kernel_frame_experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, time
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6
  7SEED = 17
  8np.random.seed(SEED); torch.manual_seed(SEED)
  9try:
 10    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 11    if device == 'cuda':
 12        torch.cuda.set_device(0)
 13        torch.zeros(1, device='cuda')
 14except Exception:
 15    device = 'cpu'
 16
 17
 18def wendland(t):
 19    t = np.asarray(t)
 20    return np.where(t < 1.0, (1.0-t)**4 * (4.0*t+1.0), 0.0)
 21
 22def centers_nested(x, levels):
 23    # C0 is all sites; subsequent sets are deterministic nested subsets.
 24    out = [x]
 25    for j in range(1, levels):
 26        stride = 2**j
 27        out.append(x[::stride])
 28    return out
 29
 30def radius(c, multiplier=2.5):
 31    d = np.sqrt(((c[:,None,:]-c[None,:,:])**2).sum(-1) + np.eye(len(c))*1e9)
 32    return multiplier * np.median(d.min(1))
 33
 34def eval_matrix(x, c, rho):
 35    return wendland(np.linalg.norm(x[:,None,:]-c[None,:,:], axis=-1)/rho)
 36
 37def frame_encode(x, u, centers, lamb=1e-5):
 38    q = np.zeros_like(u); blocks=[]; stats=[]
 39    for c in centers:
 40        rho = radius(c)
 41        A = eval_matrix(x,c,rho)
 42        gram = A.T@A + lamb*np.eye(len(c))
 43        alpha = np.linalg.solve(gram, A.T@(u-q))
 44        q = q + A@alpha
 45        stats.append({'n':len(c), 'rho':float(rho), 'density':float((A>0).mean()),
 46                      'cond':float(np.linalg.cond(gram)), 'residual':float(np.linalg.norm(u-q)/np.linalg.norm(u))})
 47        blocks.append(alpha)
 48    return blocks, stats
 49
 50def flatten_blocks(blocks, sizes):
 51    z=[]
 52    for b,n in zip(blocks,sizes):
 53        z.append(np.pad(b, (0, max(sizes)-n)))
 54    return np.stack(z, axis=0)
 55
 56class CoeffOperator(nn.Module):
 57    def __init__(self, levels, width):
 58        super().__init__(); self.levels=levels; self.width=width
 59        self.net=nn.Sequential(nn.Linear(levels*width,96),nn.GELU(),nn.Linear(96,96),nn.GELU(),nn.Linear(96,levels*width))
 60    def forward(self,z): return self.net(z.flatten(1)).reshape(-1,self.levels,self.width)
 61
 62class PointOperator(nn.Module):
 63    def __init__(self, m):
 64        super().__init__(); self.net=nn.Sequential(nn.Linear(m,96),nn.GELU(),nn.Linear(96,96),nn.GELU(),nn.Linear(96,m))
 65    def forward(self,x): return self.net(x)
 66
 67def make_fields(n, mside):
 68    g=np.linspace(0,1,mside); X=np.stack(np.meshgrid(g,g,indexing='ij'),-1).reshape(-1,2)
 69    vals=[]; outs=[]
 70    for _ in range(n):
 71        a=np.random.randn(5)*.7; b=np.random.randn(5)*.7
 72        v=(a[0]*np.sin(2*np.pi*X[:,0])+a[1]*np.cos(2*np.pi*X[:,1])+a[2]*np.sin(4*np.pi*(X[:,0]+X[:,1]))+
 73           a[3]*np.cos(6*np.pi*X[:,0])+a[4]*np.sin(6*np.pi*X[:,1]))
 74        # nonlinear nonlocal-ish target, still small and deterministic
 75        u=np.tanh(v)+.15*np.sin(8*np.pi*X[:,0])*v
 76        vals.append(v); outs.append(u)
 77    return X,np.asarray(vals),np.asarray(outs)
 78
 79def main():
 80    mside=12; M=mside*mside; levels=3
 81    X,V,U=make_fields(180,mside)
 82    centers=centers_nested(X,levels); sizes=[len(c) for c in centers]; width=max(sizes)
 83    # Core math check on a held-out field: each sequential projection should not increase residual
 84    blocks, check=frame_encode(X,V[0],centers)
 85    residuals=[1.0]+[s['residual'] for s in check]
 86    monotonic=all(residuals[i+1] <= residuals[i]+1e-9 for i in range(len(residuals)-1))
 87    sparse=all(s['density'] < .5 for s in check)
 88    # Encode all data, and decode at original points. Use same basis for target coefficients.
 89    Z=[]; T=[]
 90    for v,u in zip(V,U):
 91        zb,_=frame_encode(X,v,centers); tb,_=frame_encode(X,u,centers)
 92        Z.append(flatten_blocks(zb,sizes)); T.append(flatten_blocks(tb,sizes))
 93    Z=np.asarray(Z,dtype=np.float32); T=np.asarray(T,dtype=np.float32)
 94    # split fixed
 95    tr=np.arange(0,130); te=np.arange(130,180)
 96    ztr=torch.tensor(Z[tr],device=device); ttr=torch.tensor(T[tr],device=device)
 97    zte=torch.tensor(Z[te],device=device); tte=torch.tensor(T[te],device=device)
 98    vtr=torch.tensor(V[tr],dtype=torch.float32,device=device); utr=torch.tensor(U[tr],dtype=torch.float32,device=device)
 99    vte=torch.tensor(V[te],dtype=torch.float32,device=device); ute=torch.tensor(U[te],dtype=torch.float32,device=device)
100    torch.manual_seed(SEED)
101    idea=CoeffOperator(levels,width).to(device); base=PointOperator(M).to(device)
102    oi=torch.optim.Adam(idea.parameters(),lr=2e-3); ob=torch.optim.Adam(base.parameters(),lr=2e-3)
103    t0=time.time()
104    for step in range(700):
105        oi.zero_grad(); pred=idea(ztr); ((pred-ttr)**2).mean().backward(); oi.step()
106        ob.zero_grad(); pb=base(vtr); ((pb-utr)**2).mean().backward(); ob.step()
107    train_seconds=time.time()-t0
108    with torch.no_grad():
109        # Reconstruct output from predicted coefficient blocks at X.
110        pred=idea(zte).cpu().numpy(); pbase=base(vte).cpu().numpy()
111    Aall=[eval_matrix(X,c,radius(c)) for c in centers]
112    def decode(P):
113        return sum(P[:,j,:sizes[j]]@Aall[j].T for j in range(levels))
114    pred_u=decode(pred); true_u=U[te]
115    err=float(np.linalg.norm(pred_u-true_u)/np.linalg.norm(true_u))
116    berr=float(np.linalg.norm(pbase-U[te])/np.linalg.norm(U[te]))
117    # irregular query stability/generalization: decode on jittered points against analytic target generation per sample
118    rng=np.random.default_rng(SEED); Y=np.clip(X+rng.normal(0,.018,X.shape),0,1)
119    AY=[eval_matrix(Y,c,radius(c)) for c in centers]
120    qpred=sum(pred[:,j,:sizes[j]]@AY[j].T for j in range(levels))
121    # target evaluated directly at Y using same latent Fourier coefficients is unavailable; compare interpolation consistency
122    interp=float(np.mean(np.linalg.norm(qpred-pred_u,axis=1)/ (np.linalg.norm(pred_u,axis=1)+1e-7)))
123    result={'device':device,'math_check':{'residuals':residuals,'monotonic_nonincrease':monotonic,'all_levels_sparse':sparse,'levels':check},
124            'experiment':{'train_seconds':train_seconds,'relative_l2_idea':err,'relative_l2_baseline':berr,'query_jitter_relative_change':interp,
125                          'idea_parameters':sum(p.numel() for p in idea.parameters()),'baseline_parameters':sum(p.numel() for p in base.parameters()),
126                          'active_basis_coefficients':sum(sizes),'dense_input_tokens':M},
127            'seed':SEED}
128    Path('results.json').write_text(json.dumps(result,indent=2))
129    print(json.dumps(result,indent=2))
130
131if __name__=='__main__': main()