Nonreciprocal Brownian Optimizer / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, random
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 1505
  6rng = np.random.default_rng(SEED)
  7
  8
  9def matrix(k0, k1, k2, gamma=1.0):
 10    return np.array([[k0+k1, -k1], [-k2, k0+k2]], dtype=float) / gamma
 11
 12
 13def lyapunov_2x2(A, D):
 14    # Solve A S + S A.T = 2D for symmetric S=(a,b;b,c).
 15    M = np.array([[2*A[0,0], 2*A[0,1], 0],
 16                  [A[1,0], A[0,0]+A[1,1], A[0,1]],
 17                  [0, 2*A[1,0], 2*A[1,1]]])
 18    rhs = 2*np.array([D[0,0], D[0,1], D[1,1]])
 19    a,b,c = np.linalg.solve(M, rhs)
 20    return np.array([[a,b],[b,c]])
 21
 22
 23def ep_rate(A, D):
 24    S = lyapunov_2x2(A, D)
 25    invS = np.linalg.inv(S)
 26    B = D @ invS - A
 27    # E[(Bx)^T D^-1 (Bx)] = tr(B S B.T D^-1)
 28    ep = np.trace(B @ S @ B.T @ np.linalg.inv(D))
 29    return float(ep), S
 30
 31
 32def stability_checks():
 33    out = {}
 34    # Prediction 1: continuous stability boundary k0+k1+k2=0.
 35    k0, gamma = 1.0, 1.0
 36    vals = np.linspace(-1.4, 0.4, 19)
 37    rows = []
 38    for s in vals:
 39        A = matrix(k0, s/2, s/2, gamma)
 40        stable_pred = (k0+s) > 0
 41        stable_obs = np.min(np.real(np.linalg.eigvals(A))) > 0
 42        rows.append([float(s), bool(stable_pred), bool(stable_obs)])
 43    # discrete boundary for a representative asymmetric pair, spectral radius=1.
 44    eta_grid = np.linspace(.01, 1.5, 1500)
 45    k1,k2 = .75,.25
 46    A = matrix(k0,k1,k2,gamma)
 47    rhos = np.array([max(abs(np.linalg.eigvals(np.eye(2)-e*A))) for e in eta_grid])
 48    idx = np.where(rhos >= 1)[0][0]
 49    eta_obs = eta_grid[idx]
 50    eig = np.linalg.eigvals(A)
 51    eta_pred = min(2*np.real(z)/(abs(z)**2) for z in eig)
 52    # Direct deterministic rollout: classify as divergent if norm grows by 1e6.
 53    rollout = []
 54    for eta in np.linspace(.1, 1.4, 14):
 55        M = np.eye(2)-eta*A
 56        x=np.array([1., -1.]); initial=np.linalg.norm(x)
 57        for _ in range(300): x=M@x
 58        rollout.append({'eta':float(eta), 'rho':float(max(abs(np.linalg.eigvals(M)))), 'diverged':bool(np.linalg.norm(x)>1e6*initial)})
 59    out['stability'] = {'continuous_rows': rows, 'discrete_pred_eta': float(eta_pred), 'discrete_obs_eta': float(eta_obs), 'rho_at_pred': float(max(abs(np.linalg.eigvals(np.eye(2)-eta_pred*A)))), 'rollout':rollout}
 60
 61    # Prediction 2: EP vanishes at delta=0 and is quadratic near reciprocity.
 62    D=np.eye(2) * .3
 63    s=.8
 64    deltas=np.linspace(-.7,.7,15)
 65    eps=[]
 66    for d in deltas:
 67        e,_=ep_rate(matrix(k0,(s+d)/2,(s-d)/2),D)
 68        eps.append(e)
 69    small=np.abs(deltas)<=.3
 70    coeff=np.polyfit(deltas[small]**2, np.array(eps)[small], 1)[0]
 71    zero=float(eps[len(eps)//2])
 72    # Ratio EP/delta^2 over small nonzero values.
 73    ratios=[eps[i]/deltas[i]**2 for i in range(len(deltas)) if small[i] and abs(deltas[i])>.05]
 74    out['ep']={'deltas':deltas.tolist(),'rates':eps,'zero_rate':zero,'quadratic_coeff':float(coeff),'small_ratio_mean':float(np.mean(ratios)),'small_ratio_cv':float(np.std(ratios)/np.mean(ratios))}
 75
 76    # Prediction 3: covariance is finite only on stable side and grows toward boundary.
 77    cov_rows=[]
 78    for s in [-.8,-.5,0,.5,1.0]:
 79        A=matrix(k0,s/2,s/2)
 80        S=lyapunov_2x2(A,D)
 81        cov_rows.append({'s':s,'max_variance':float(np.max(np.linalg.eigvalsh(S))), 'pred_stable':bool(k0+s>0)})
 82    out['covariance']=cov_rows
 83    return out
 84
 85
 86def ml_experiment():
 87    # Tiny nonlinear problem with fixed full-batch gradients; explicit Langevin noise
 88    # makes the optimizer comparison deterministic and compute-matched.
 89    import torch
 90    torch.set_num_threads(4)
 91    torch.manual_seed(SEED)
 92    n=512
 93    x=torch.randn(n,2)
 94    y=((x[:,0]*x[:,1])>0).long()
 95    # fixed MLP parameter vector, functional forward
 96    shapes=[(2,24),(24,24),(24,2)]
 97    sizes=[a*b for a,b in shapes]+[24,24,2]
 98    total=sum(sizes)
 99    def unpack(v):
100        p=[]; q=0
101        for (a,b),sz in zip(shapes,sizes[:3]): p.append(v[q:q+sz].reshape(a,b)); q+=sz
102        for sz in sizes[3:]: p.append(v[q:q+sz]); q+=sz
103        return p
104    def loss(v):
105        w1,w2,w3,b1,b2,b3=unpack(v)
106        z=torch.tanh(x@w1+b1); z=torch.tanh(z@w2+b2); logits=z@w3+b3
107        return torch.nn.functional.cross_entropy(logits,y)
108    def grad(v):
109        v=v.detach().requires_grad_(True); l=loss(v); return l.detach(),torch.autograd.grad(l,v)[0].detach()
110    init=torch.randn(total)*.15
111    steps=250; eta=.08; T=.0008
112    def run(kind):
113        a=init.clone(); b=init.clone(); losses=[]
114        for t in range(steps):
115            if kind=='baseline':
116                l,g=grad(a); a=a-eta*g
117                losses.append(float(l))
118            else:
119                l1,g1=grad(a); l2,g2=grad(b)
120                if kind=='reciprocal': k1=k2=.8
121                else: k1,k2=.95,.25
122                noise1=torch.randn_like(a)*math.sqrt(2*eta*T); noise2=torch.randn_like(a)*math.sqrt(2*eta*T)
123                a=a-eta*(g1+k1*(a-b))-eta*.02*a+noise1
124                b=b-eta*(g2+k2*(b-a))-eta*.02*b+noise2
125                losses.append(float(loss((a+b)/2)))
126        return float(loss((a+b)/2) if kind!='baseline' else loss(a)), float((losses[0]-losses[-1])), losses
127    result={}
128    for k in ['baseline','reciprocal','nonreciprocal']:
129        final,drop,trace=run(k); result[k]={'final_loss':final,'loss_drop':drop,'trace_tail':trace[-10:]}
130    return result
131
132if __name__=='__main__':
133    result={'seed':SEED,'checks':stability_checks()}
134    try:
135        result['ml']=ml_experiment()
136    except Exception as e:
137        result['ml_error']=repr(e)
138    Path('results.json').write_text(json.dumps(result,indent=2))
139    c=result['checks']
140    print(json.dumps({'stability':c['stability'],'ep':c['ep'],'covariance':c['covariance'],'ml':result.get('ml',result.get('ml_error'))},indent=2))