Covariance-Adaptive Hermite Latent Bottleneck / hermite_bottleneck.py
Mechanism confirmed, baseline not beaten
1import itertools
2import math
3import numpy as np
4from scipy.special import eval_hermitenorm, roots_hermitenorm
5
6
7def multi_indices(d, degree):
8 out = []
9 def rec(i, left, cur):
10 if i == d - 1:
11 out.append(tuple(cur + [left]))
12 return
13 for k in range(left + 1):
14 rec(i + 1, left - k, cur + [k])
15 for total in range(degree + 1):
16 rec(0, total, [])
17 return out
18
19
20def hermite_basis(x, indices):
21 x = np.asarray(x)
22 vals = []
23 for a in indices:
24 h = np.ones(x.shape[0])
25 for j, k in enumerate(a):
26 if k:
27 h *= eval_hermitenorm(k, x[:, j]) / math.sqrt(math.factorial(k))
28 vals.append(h)
29 return np.stack(vals, axis=1)
30
31
32def density_ratio(x, cov, mean=None):
33 x = np.asarray(x)
34 d = x.shape[1]
35 mean = np.zeros(d) if mean is None else np.asarray(mean)
36 sign, logdet = np.linalg.slogdet(cov)
37 if sign <= 0:
38 raise ValueError('covariance must be positive definite')
39 inv = np.linalg.inv(cov)
40 delta = x - mean
41 # M_cov,mean / M_I,0; clipping only prevents rare MC overflow.
42 logf = -.5 * logdet - .5*np.einsum('bi,ij,bj->b', delta, inv, delta) + .5*np.sum(x*x, axis=1)
43 return np.exp(np.clip(logf, -80, 80))
44
45
46def project_density(samples_w, cov, degree, mean=None):
47 idx = multi_indices(samples_w.shape[1], degree)
48 B = hermite_basis(samples_w, idx)
49 f = density_ratio(samples_w, cov, mean)
50 # Orthonormal Hermites: coefficient is E_w[f phi] = E_M[phi].
51 coef = np.mean(B * f[:, None], axis=0)
52 return idx, coef
53
54
55def reconstruct_error(samples_w, cov, idx, coef, mean=None):
56 f = density_ratio(samples_w, cov, mean)
57 pred = hermite_basis(samples_w, idx) @ coef
58 return float(np.sqrt(np.mean((f - pred)**2))), float(np.sqrt(np.mean(f*f)))
59
60
61def heat_identity_check(rng, d=2, n=12, nz=12):
62 # Product Gauss-Hermite quadrature over Z gives a deterministic check of the
63 # expectation in the heat-semigroup identity at independent test points v.
64 u = np.array([.35, -.25])[:d]
65 t = .31
66 nodes, weights = roots_hermitenorm(nz)
67 weights = weights / math.sqrt(2 * math.pi)
68 grid = np.array(list(itertools.product(nodes, repeat=d)))
69 wg = np.array([np.prod([weights[k] for k in ix])
70 for ix in itertools.product(range(nz), repeat=d)])
71 v = rng.normal(size=(n, d))
72 uu = u[None, :] + math.sqrt(2*t) * grid
73 lhs = np.exp(uu @ v.T - .5*np.sum(uu*uu, axis=1)[:, None]).T @ wg
74 s = 1 + 2*t
75 rhs = s**(-d/2) * np.exp((v @ u)/s - np.sum(u*u)/(2*s) + t*np.sum(v*v, axis=1)/s)
76 return float(np.max(np.abs(lhs-rhs) / (np.abs(rhs)+1e-12)))
77
78
79def run(seed=7):
80 rng = np.random.default_rng(seed)
81 heat_rel = heat_identity_check(rng)
82 d = 4
83 # Shared samples make comparisons low-noise. All eigenvalues lie in q<1 regime.
84 nw, nm = 160000, 120000
85 w = rng.normal(size=(nw, d))
86 covs = [np.diag([1.04, .97, 1.01, .99]),
87 np.diag([1.20, .82, 1.08, .94]),
88 np.diag([1.55, .62, 1.18, .88])]
89 qs = [float(np.max(np.abs(np.linalg.eigvalsh(c)-1))) for c in covs]
90 maxdeg = 6
91 results = []
92 # Empirical C is calibrated only on the mildest covariance, as a practical tail scale.
93 calib = covs[1]
94 tails_cal = []
95 projections = {}
96 for N in range(maxdeg + 1):
97 idx, co = project_density(w, calib, N)
98 err, norm = reconstruct_error(w, calib, idx, co)
99 tails_cal.append(err)
100 projections[(id(calib), N)] = (idx, co)
101 # C_hat is deliberately a held-out-style empirical scale, not a claimed universal constant.
102 C_hat = max(tails_cal[N] / (qs[1] ** ((N+1)/2)) for N in range(maxdeg+1))
103 target = .08
104 fixed_N = 4
105 for ci, cov in enumerate(covs):
106 q = qs[ci]
107 rows = []
108 for N in range(maxdeg + 1):
109 idx, co = project_density(w, cov, N)
110 err, norm = reconstruct_error(w, cov, idx, co)
111 rows.append((N, len(idx), err / norm))
112 chosen = next((N for N in range(maxdeg+1) if C_hat*q**((N+1)/2) <= target), maxdeg)
113 aidx, aco = project_density(w, cov, chosen)
114 actual, norm = reconstruct_error(w, cov, aidx, aco)
115 fidx, fco = project_density(w, cov, fixed_N)
116 fixed, _ = reconstruct_error(w, cov, fidx, fco)
117 results.append({'q': q, 'adaptive_degree': chosen,
118 'adaptive_coefficients': len(aidx),
119 'adaptive_rel_error': actual/norm,
120 'fixed_degree': fixed_N, 'fixed_coefficients': len(fidx),
121 'fixed_rel_error': fixed/norm, 'curve': rows})
122 return {'heat_relative_error': heat_rel, 'C_hat': C_hat, 'target': target,
123 'dimension': d, 'results': results}
124
125if __name__ == '__main__':
126 import json
127 print(json.dumps(run(), indent=2))