FMM-Accelerated Polyharmonic Neural Field Head / phs_experiment.py
Beats tuned baseline
1import json, time
2import numpy as np
3from scipy.linalg import null_space
4from scipy.spatial.distance import cdist
5
6SEED = 2128
7
8def phs(X, Y=None, k=2):
9 if Y is None: Y = X
10 r = cdist(X, Y)
11 if k % 2: return r ** k
12 return r ** k * np.log(np.maximum(r, 1e-15))
13
14def basis(X):
15 return np.column_stack([np.ones(len(X)), X[:, 0], X[:, 1]])
16
17def head(X, w, beta, Xq, k=2, chunk=512):
18 out = basis(Xq) @ beta
19 for a in range(0, len(Xq), chunk):
20 out[a:a+chunk] += phs(Xq[a:a+chunk], X, k) @ w
21 return out
22
23def vecchia_precision(X, m=8, ell=None):
24 n = len(X)
25 D = cdist(X, X); D[np.diag_indices(n)] = np.inf
26 if ell is None: ell = 1.5 * np.median(np.min(D, axis=1))
27 C = np.exp(-cdist(X, X) / max(ell, 1e-8))
28 T = np.zeros((n, n))
29 order = np.argsort(X[:, 0] + .173 * X[:, 1])
30 for pos, i in enumerate(order):
31 if pos == 0:
32 T[i, i] = 1.0
33 continue
34 earlier = order[:pos]
35 nn = earlier[np.argsort(np.linalg.norm(X[earlier] - X[i], axis=1))[:m]]
36 S = C[np.ix_(nn, nn)] + 1e-8*np.eye(len(nn))
37 c = C[i, nn]
38 a = np.linalg.solve(S, c)
39 var = max(C[i, i] - a @ c, 1e-8)
40 T[i, i] = 1.0 / np.sqrt(var)
41 T[i, nn] = -a / np.sqrt(var)
42 return T
43
44def projected_pcg(M, Q, rhs, pre=None, tol=1e-8, maxit=300):
45 A = lambda z: Q.T @ (M @ (Q @ z))
46 z = np.zeros_like(rhs)
47 r = rhs - A(z)
48 r0 = np.linalg.norm(r)
49 if pre is None: solve = lambda x: x
50 else:
51 # P_vecchia is the approximate inverse/preconditioning operator; apply it directly
52 G = Q.T @ pre @ Q
53 G = (G + G.T)/2 + 1e-8*np.eye(G.shape[0])
54 solve = lambda x: G @ x
55 q = solve(r); p = q.copy(); rz = r @ q; it = 0
56 for it in range(1, maxit+1):
57 Ap = A(p); den = p @ Ap
58 if abs(den) < 1e-14: break
59 alpha = rz / den; z += alpha*p; r -= alpha*Ap
60 if np.linalg.norm(r) <= tol * max(r0, 1e-15): break
61 q = solve(r); rz_new = r @ q
62 p = q + (rz_new/rz)*p; rz = rz_new
63 return z, it, np.linalg.norm(r)/max(r0,1e-15)
64
65def main():
66 rng = np.random.default_rng(SEED)
67 n, nq = 180, 300
68 X = rng.uniform(-1,1,(n,2)); Xq = rng.uniform(-1,1,(nq,2))
69 y = np.sin(3*X[:,0]) * np.cos(2*X[:,1]) + .2*X[:,0]**2
70 B = basis(X); Q = null_space(B.T)
71 M = phs(X, k=2)
72 # Exact augmented solution is the reference interpolation.
73 aug = np.block([[M, B], [B.T, np.zeros((3,3))]])
74 sol = np.linalg.solve(aug, np.r_[y, np.zeros(3)])
75 w0, beta0 = sol[:n], sol[n:]
76 pred = head(X, w0, beta0, Xq)
77 ref_err = np.max(np.abs(head(X,w0,beta0,X)-y))
78 constraint = np.linalg.norm(B.T @ w0)
79
80 # Prediction 1: geometric scaling M(rho) projected scales as rho^2 for k=2.
81 scales = [.5, 1., 2., 4.]
82 scale_ratios = []
83 for rho in scales:
84 Mr = phs(rho*X, k=2)
85 scale_ratios.append(np.linalg.norm(Q.T@Mr@Q) / np.linalg.norm(Q.T@M@Q))
86 predicted = [rho*rho for rho in scales]
87
88 # Prediction 2: projected polynomial terms are annihilated exactly.
89 D2 = cdist(X,X)**2
90 proj_poly = np.linalg.norm(Q.T @ D2 @ Q) / max(np.linalg.norm(Q.T@M@Q),1e-15)
91 # Random coefficients expose constraint preservation.
92 wr = Q @ rng.normal(size=n-3)
93 constraint_random = np.linalg.norm(B.T @ wr)
94
95 # Prediction 3: Vecchia quality improves as neighbor set grows; report PCG residual/iters.
96 rhs = Q.T @ y
97 rows = []
98 z, it, rr = projected_pcg(M, Q, rhs, pre=None)
99 rows.append({'m':'none', 'iterations':it, 'relative_residual':float(rr),
100 'constraint':float(np.linalg.norm(B.T@(Q@z)))})
101 for m in [2,4,8,16,32]:
102 T = vecchia_precision(X, m=m)
103 P = T @ T.T
104 z, it, rr = projected_pcg(M,Q,rhs,pre=P)
105 rows.append({'m':m, 'iterations':it, 'relative_residual':float(rr),
106 'constraint':float(np.linalg.norm(B.T@(Q@z)))})
107
108 # Dense versus chunked head evaluation (same result, lower peak interaction memory).
109 t0=time.perf_counter(); dense = basis(Xq)@beta0 + phs(Xq,X)@w0; td=time.perf_counter()-t0
110 t0=time.perf_counter(); chunked=head(X,w0,beta0,Xq,chunk=32); tc=time.perf_counter()-t0
111 result = {'seed':SEED, 'n':n, 'queries':nq,
112 'math_checks': {
113 'scale_rho':scales, 'predicted_rho2':predicted, 'observed_projected_norm_ratios':scale_ratios,
114 'scale_max_abs_error':float(max(abs(a-b) for a,b in zip(scale_ratios,predicted))),
115 'projected_squared_distance_ratio':float(proj_poly),
116 'predicted_projected_polynomial_ratio':0.0,
117 'random_nullspace_constraint_norm':float(constraint_random),
118 'exact_interpolation_max_error':float(ref_err), 'exact_solution_constraint_norm':float(constraint)},
119 'pcg_vecchia':rows,
120 'evaluation': {'dense_seconds':td, 'chunked_seconds':tc,
121 'chunked_vs_dense_max_error':float(np.max(np.abs(chunked-dense))),
122 'reference_query_rmse':float(np.sqrt(np.mean((pred-(np.sin(3*Xq[:,0])*np.cos(2*Xq[:,1])+.2*Xq[:,0]**2))**2)))}}
123 print(json.dumps(result, indent=2))
124
125if __name__ == '__main__': main()