Sparse Multiscale Kernel-Frame Operator / kernel_frame_experiment.py
Mechanism failed
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()