Hilbert-Schmidt-scale KSD loss / ksd_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json
 2import numpy as np
 3from pathlib import Path
 4
 5
 6def stein_kernel(x, y, gamma):
 7    # Pairwise matrix for x[b,n,d], y[b,m,d], target N(0,I), RBF kernel.
 8    diff = x[:, :, None, :] - y[:, None, :, :]
 9    r2 = np.sum(diff * diff, axis=-1)
10    dot = np.einsum('bid,bjd->bij', x, y)
11    d = x.shape[-1]
12    return np.exp(-gamma * r2) * (dot + 2 * gamma * d - (2 * gamma + 4 * gamma**2) * r2)
13
14
15def stein_kernel_paired(x, y, gamma):
16    # Paired h(x_i,y_i), avoiding a large pairwise matrix.
17    diff = x - y
18    r2 = np.sum(diff * diff, axis=-1)
19    dot = np.sum(x * y, axis=-1)
20    d = x.shape[-1]
21    return np.exp(-gamma * r2) * (dot + 2 * gamma * d - (2 * gamma + 4 * gamma**2) * r2)
22
23
24def trace_c(d, gamma):
25    return (1 + 2 * gamma) * d
26
27
28def trace_c2(d, gamma):
29    a = 1 + 8 * gamma
30    return (a ** (-d / 2.0)) * d / (a * a) * ((64 * gamma**4 + 32 * gamma**3 + 4 * gamma**2) * d + 128 * gamma**4 + 128 * gamma**3 + 80 * gamma**2 + 16 * gamma + 1)
31
32
33def estimates(x, gamma):
34    h = stein_kernel(x, x, gamma)
35    n = x.shape[1]
36    v = np.sqrt(np.maximum(h.mean(axis=(1, 2)), 0))
37    u2 = (h.sum(axis=(1, 2)) - np.trace(h, axis1=1, axis2=2)) / (n * (n - 1))
38    u = np.sqrt(np.maximum(u2, 0))
39    return v, u, u2
40
41
42def replication(d, n, gamma, reps, rng, chunk=100):
43    vals_v, vals_u, vals_u2 = [], [], []
44    for start in range(0, reps, chunk):
45        b = min(chunk, reps - start)
46        x = rng.standard_normal((b, n, d))
47        v, u, u2 = estimates(x, gamma)
48        vals_v.append(v); vals_u.append(u); vals_u2.append(u2)
49    v = np.concatenate(vals_v); u = np.concatenate(vals_u); u2 = np.concatenate(vals_u2)
50    return {
51        'd': d, 'n': n, 'reps': reps,
52        'V_mean': float(v.mean()), 'V_sd': float(v.std(ddof=1)),
53        'U_mean': float(u.mean()), 'U_sd': float(u.std(ddof=1)),
54        'U2_mean': float(u2.mean()), 'U2_sd': float(u2.std(ddof=1)),
55        'U2_negative_frac': float(np.mean(u2 < 0)),
56        'trace_scale': float(np.sqrt(trace_c(d, gamma) / n)),
57        'hs_scale': float(np.sqrt(np.sqrt(trace_c2(d, gamma)) / n)),
58        'effective_rank': float(trace_c(d, gamma)**2 / trace_c2(d, gamma)),
59    }
60
61
62def main():
63    rng = np.random.default_rng(395)
64    gamma = 0.5
65    check = {}
66    for d in (2, 5, 10):
67        x = rng.standard_normal((200000, d)); y = rng.standard_normal((200000, d))
68        hxy = stein_kernel_paired(x, y, gamma)
69        hxx = np.sum(x*x, axis=1) + 2*gamma*d
70        exact2 = trace_c2(d, gamma)
71        emp2 = float(np.mean(hxy*hxy))
72        check[str(d)] = {
73            'diagonal_empirical': float(hxx.mean()), 'diagonal_exact': trace_c(d, gamma),
74            'second_moment_empirical': emp2, 'second_moment_exact': exact2,
75            'paired_kernel_mean': float(hxy.mean()),
76            'relative_diag_error': float(abs(hxx.mean()-trace_c(d,gamma))/trace_c(d,gamma)),
77            'relative_second_moment_error': float(abs(emp2-exact2)/exact2)
78        }
79    results = []
80    for d in (2, 5, 10):
81        for n in (32, 64, 128, 512):
82            reps = 1000 if n <= 128 else 300
83            results.append(replication(d, n, gamma, reps, rng))
84    d5 = [r for r in results if r['d'] == 5]
85    ns = np.array([r['n'] for r in d5], float)
86    slopes = {
87        'V_mean_vs_n': float(np.polyfit(np.log(ns), np.log([r['V_mean'] for r in d5]), 1)[0]),
88        'U_mean_vs_n': float(np.polyfit(np.log(ns), np.log([r['U_mean'] for r in d5]), 1)[0])
89    }
90    out = {'gamma': gamma, 'math_check': check, 'results': results, 'loglog_slopes_d5': slopes}
91    Path('results.json').write_text(json.dumps(out, indent=2))
92    print(json.dumps(out, indent=2))
93
94if __name__ == '__main__':
95    main()