Hilbert-Schmidt-scale KSD loss / ksd_experiment.py
Mechanism confirmed, baseline not beaten
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()