Gauge-Free Spectral OT Layer / run_experiment.py
Failed on benchmark
1import json, math, os
2import numpy as np
3import torch
4from scipy.special import logsumexp
5
6SEED = 2227
7np.random.seed(SEED)
8torch.manual_seed(SEED)
9torch.set_num_threads(min(8, os.cpu_count() or 1))
10
11
12def sinkhorn(cost, eps=0.15, iters=250, a=None, b=None):
13 # Log-domain Sinkhorn avoids false gauge violations from kernel clipping.
14 n, m = cost.shape
15 if a is None: a = np.ones(n) / n
16 if b is None: b = np.ones(m) / m
17 logK = -cost / eps
18 logu = np.zeros(n); logv = np.zeros(m)
19 for _ in range(iters):
20 logu = np.log(a) - logsumexp(logK + logv[None, :], axis=1)
21 logv = np.log(b) - logsumexp(logK + logu[:, None], axis=0)
22 return np.exp(logu[:, None] + logK + logv[None, :])
23
24
25def gauge_matrix(k):
26 # r_i + s_j, with one target potential removed.
27 G = np.zeros((k*k, 2*k-1))
28 for i in range(k):
29 for j in range(k):
30 q = i*k+j
31 G[q, i] = 1
32 if j < k-1: G[q, k+j] = 1
33 return G
34
35
36def gauge_project(Phi, G):
37 # Orthogonal projection in pair/cost space.
38 Q = np.eye(G.shape[0]) - G @ np.linalg.pinv(G)
39 A = Q @ Phi
40 U0, s, Vt = np.linalg.svd(A, full_matrices=False)
41 rank = int(np.sum(s > max(A.shape)*np.finfo(float).eps*s[0])) if s[0] else 0
42 U = Vt[:rank].T
43 return Q, A, U, s, rank
44
45
46def make_features(k=8, scales=(1.0, .1, .01), gauge_amp=1.0, seed=0):
47 rng = np.random.default_rng(seed)
48 G = gauge_matrix(k)
49 # Three identifiable random pair features with prescribed anisotropy.
50 R = rng.normal(size=(k*k, 3)) @ np.diag(scales)
51 # Remove accidental gauge components so the quotient is controlled.
52 Q = np.eye(k*k) - G @ np.linalg.pinv(G)
53 R = Q @ R
54 # Exact row/column gauge directions, mixed into columns to mimic nuisance biases.
55 B = G[:, :5] * gauge_amp
56 Phi = np.concatenate([B, R], axis=1)
57 return Phi, G
58
59
60def jacobian_plan(Phi, theta, k, eps=.15):
61 # Finite difference is intentionally independent of the whitening code.
62 base = sinkhorn((-Phi @ theta).reshape(k,k), eps)
63 J = np.zeros((k*k, Phi.shape[1]))
64 h = 1e-5
65 for f in range(Phi.shape[1]):
66 t = theta.copy(); t[f] += h
67 p = sinkhorn((-Phi @ t).reshape(k,k), eps)
68 J[:, f] = ((p-base)/h).ravel()
69 return J
70
71
72def prediction_sweeps():
73 # P1: exact gauge invariance over several amplitudes.
74 k=8; eps=.15
75 Phi, G = make_features(k, gauge_amp=1.0, seed=4)
76 rng=np.random.default_rng(5); theta=rng.normal(size=Phi.shape[1])
77 gauge_theta=np.zeros(Phi.shape[1]); gauge_theta[:5]=rng.normal(size=5)
78 inv=[]
79 for amp in [0., .1, 1., 10., 100.]:
80 p=sinkhorn((-Phi@(theta+amp*gauge_theta)).reshape(k,k),eps)
81 p0=sinkhorn((-Phi@theta).reshape(k,k),eps)
82 inv.append(float(np.max(np.abs(p-p0))))
83
84 # P2: quotient removes exactly the gauge-induced zero Jacobian spectrum.
85 J=jacobian_plan(Phi,theta,k,eps)
86 sv=np.linalg.svd(J,compute_uv=False)
87 Q,A,U,s,rank=gauge_project(Phi,G)
88 Jr=jacobian_plan(A@U, np.linalg.lstsq(A@U, Q@Phi@theta,rcond=None)[0], k, eps)
89 svr=np.linalg.svd(Jr,compute_uv=False)
90 zero_count=int(np.sum(sv < 1e-8))
91
92 # P3: covariance conditioning scales quadratically with feature anisotropy,
93 # while whitening gives approximately unit covariance (delta regularizer).
94 rows=[]
95 for alpha in [1., .3, .1, .03, .01]:
96 P,G2=make_features(k, scales=(1.,alpha,alpha*alpha), gauge_amp=3., seed=8)
97 Q,A,U,_,rank=gauge_project(P,G2)
98 C=(A@U).T@(A@U)/(k*k)
99 ev=np.linalg.eigvalsh(C)
100 cond=float(ev[-1]/max(ev[0],1e-15))
101 delta=1e-6
102 # Symmetric eigendecomposition implements (C+delta I)^(-1/2).
103 ce, cv=np.linalg.eigh(C+delta*np.eye(rank))
104 W=(cv*(1/np.sqrt(ce)))@cv.T
105 Cw=W.T@C@W
106 ew=np.linalg.eigvalsh(Cw)
107 ridge_pred=float(ev[-1]*(ev[0]+delta)/(ev[0]*(ev[-1]+delta)))
108 rows.append({'anisotropy_alpha':alpha,'raw_cov_condition':cond,
109 'predicted_condition_scaling':float(1/(alpha**4)),
110 'ridge_aware_whitening_prediction':ridge_pred,
111 'whitened_cov_condition':float(ew[-1]/ew[0]),
112 'cov_eigenvalues':ev.tolist(),'rank':rank})
113 return {'gauge_invariance_max_abs_plan_change':inv,
114 'gauge_jacobian_singular_values':sv.tolist(),
115 'gauge_jacobian_near_zero_count':zero_count,
116 'quotient_rank_predicted':rank,
117 'quotient_jacobian_singular_values':svr.tolist(),
118 'conditioning_sweep':rows}
119
120
121def train_compare(steps=180, trials=6):
122 # Small supervised matching: target is diagonal assignment, uniform marginals.
123 k=8; eps=.12
124 losses={'raw':[], 'gauge_free_whitened':[]}
125 final={'raw':[], 'gauge_free_whitened':[]}
126 for trial in range(trials):
127 Phi,G=make_features(k, scales=(1.,.15,.02), gauge_amp=5., seed=100+trial)
128 Q,A,U,_,rank=gauge_project(Phi,G)
129 C=(A@U).T@(A@U)/(k*k); delta=1e-5
130 # symmetric eig whitening, with theta=U W z and cost Phi theta.
131 ew,ev=np.linalg.eigh(C+delta*np.eye(rank)); W=(ev*(1/np.sqrt(ew)))@ev.T
132 target=torch.tensor(np.eye(k)/k,dtype=torch.float32)
133 def run(kind):
134 z=torch.zeros(rank if kind=='gauge_free_whitened' else Phi.shape[1], requires_grad=True)
135 opt=torch.optim.Adam([z],lr=.08)
136 vals=[]
137 Pt=torch.tensor(Phi,dtype=torch.float32)
138 if kind=='gauge_free_whitened':
139 T=torch.tensor(U@W,dtype=torch.float32)
140 for st in range(steps):
141 opt.zero_grad()
142 theta=(T@z if kind=='gauge_free_whitened' else z)
143 cost=-(Pt@theta).reshape(k,k)
144 # log-domain-ish stabilized Sinkhorn in torch.
145 logK=-cost/eps
146 logu=torch.zeros(k); logv=torch.zeros(k)
147 for _ in range(35):
148 logu=-torch.logsumexp(logK+logv[None,:],dim=1)-math.log(k)
149 logv=-torch.logsumexp(logK+logu[:,None],dim=0)-math.log(k)
150 plan=torch.exp(logu[:,None]+logK+logv[None,:])
151 loss=((plan-target)**2).mean()
152 loss.backward(); opt.step()
153 vals.append(float(loss.detach()))
154 return vals
155 for kind in losses: losses[kind].append(run(kind))
156 for kind in losses: final[kind].append(losses[kind][-1][-1])
157 mean_curves={k:np.mean(v,axis=0).tolist() for k,v in losses.items()}
158 return {'final_loss_mean':{k:float(np.mean(v)) for k,v in final.items()},
159 'final_loss_std':{k:float(np.std(v)) for k,v in final.items()},
160 'loss_at_steps':{k:{str(s):float(np.mean([x[s-1] for x in v])) for s in [10,30,60,120,180]} for k,v in losses.items()},
161 'curves':mean_curves}
162
163
164def main():
165 out={'seed':SEED,'predictions':prediction_sweeps(),'training':train_compare()}
166 with open('results.json','w') as f: json.dump(out,f,indent=2)
167 print(json.dumps(out,indent=2))
168
169if __name__=='__main__': main()