Latent-Component Schrödinger Bridge / experiment.py

✓✓ Beats tuned baseline

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