Covariance-Adaptive Hermite Latent Bottleneck / hermite_bottleneck.py

Mechanism confirmed, baseline not beaten

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