Gauge-Free Spectral OT Layer / run_experiment.py

Failed on benchmark

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