FMM-Accelerated Polyharmonic Neural Field Head / phs_experiment.py

✓✓ Beats tuned baseline

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