Bi-Maxwell Muon / experiment.py

Unverified

Raw ⬇ ZIP
  1import json, math, time
  2import numpy as np
  3
  4SEED = 2693
  5rng = np.random.default_rng(SEED)
  6
  7
  8def relax_impulse(beta, n):
  9    # X_0=1, X_t=0 afterwards; state at integer n is beta**n.
 10    return beta ** n
 11
 12
 13def relax_constant(beta, n):
 14    # X_t=1 from t=0, M_-1=0.
 15    return 1.0 - beta ** (n + 1)
 16
 17
 18def polar_ns(a, iters=5):
 19    # Newton-Schulz approximation to the polar factor, normalized for stability.
 20    x = a.astype(np.float64, copy=True)
 21    norm = np.linalg.norm(x, 2)
 22    if norm == 0:
 23        return np.zeros_like(x)
 24    x /= norm
 25    for _ in range(iters):
 26        x = 0.5 * x @ (3.0 * np.eye(x.shape[1]) - x.T @ x)
 27    return x
 28
 29
 30def update_mode(m, g, beta):
 31    return beta * m + (1.0 - beta) * g
 32
 33
 34def optimizer_run(kind, x, y, steps=350, lr=0.025, beta=.95, bf=.9, bs=.99, w=.5):
 35    W = np.zeros((x.shape[1], y.shape[1]), dtype=np.float64)
 36    m = np.zeros_like(W); mf = np.zeros_like(W); ms = np.zeros_like(W)
 37    losses=[]; cosines=[]
 38    t0=time.perf_counter()
 39    for _ in range(steps):
 40        pred=x@W; g=x.T@(pred-y)/len(x)
 41        if kind == 'adamw':
 42            # Included as a standard reference, with fixed hyperparameters.
 43            m = .9*m + .1*g
 44            v = .999*(v if 'v' in locals() else np.zeros_like(W)) + .001*g*g
 45            W -= lr * m/(np.sqrt(v)+1e-8)
 46            upd = m/(np.sqrt(v)+1e-8)
 47        elif kind == 'muon':
 48            m = update_mode(m,g,beta); upd=polar_ns(m)
 49            W -= lr*upd
 50        else:
 51            mf=update_mode(mf,g,bf); ms=update_mode(ms,g,bs)
 52            mix=w*mf+(1-w)*ms; upd=polar_ns(mix)
 53            W -= lr*upd
 54        losses.append(float(np.mean((pred-y)**2)/2))
 55        gn=np.linalg.norm(g); un=np.linalg.norm(upd)
 56        cosines.append(float(np.sum(g*upd)/(gn*un+1e-12)))
 57    return {'final_loss':losses[-1], 'best_loss':min(losses), 'loss_100':losses[99],
 58            'mean_update_cos':float(np.mean(cosines)), 'seconds':time.perf_counter()-t0}
 59
 60
 61def main():
 62    # Stage-1 prediction A: tau ordering implies beta ordering and measured relaxation ordering.
 63    eta=1.0; tauf=2.0; taus=20.0
 64    bf=math.exp(-eta/tauf); bs=math.exp(-eta/taus)
 65    ngrid=np.arange(0,1000)
 66    hf=int(ngrid[np.argmin(np.abs(np.array([relax_constant(bf,n) for n in ngrid])-.5))])
 67    hs=int(ngrid[np.argmin(np.abs(np.array([relax_constant(bs,n) for n in ngrid])-.5))])
 68    pred_hf=math.log(.5)/math.log(bf)-1
 69    pred_hs=math.log(.5)/math.log(bs)-1
 70
 71    # Prediction B: impulse residual is exactly beta^n; sweep beta and compare fitted log slope.
 72    beta_sweep=[.5,.8,.9,.95,.99]
 73    impulse=[]
 74    for b in beta_sweep:
 75        n=np.arange(1,51)
 76        slope=np.polyfit(n,np.log(np.array([relax_impulse(b,int(k)) for k in n])),1)[0]
 77        impulse.append({'beta':b,'predicted_log_slope':math.log(b),'observed_log_slope':float(slope),
 78                        'residual_at_20':relax_impulse(b,20),'predicted_residual_at_20':b**20})
 79
 80    # Prediction C: sinusoidal response. For input cos(omega t), |H| for EMA is analytic.
 81    def gain(b,om):
 82        return (1-b)/math.sqrt(1+b*b-2*b*math.cos(om))
 83    freq_rows=[]
 84    for om in [0.1,0.5,1.5,math.pi]:
 85        N=4000; t=np.arange(N); signal=np.cos(om*t)
 86        mf=ms=0.; out=[]
 87        for z in signal:
 88            mf=bf*mf+(1-bf)*z; ms=bs*ms+(1-bs)*z; out.append(.5*mf+.5*ms)
 89        out=np.asarray(out); tail=slice(1000,None)
 90        # projection amplitude avoids phase sensitivity
 91        obs=2*abs(np.mean(out[tail]*np.exp(-1j*om*t[tail])))
 92        pred=.5*gain(bf,om)+.5*gain(bs,om)
 93        freq_rows.append({'omega':om,'predicted_gain':pred,'observed_gain':float(obs)})
 94
 95    # Deterministic low-dimensional matrix regression; same polar routine for Muon variants.
 96    local=np.random.default_rng(SEED)
 97    X=local.normal(size=(256,16)); trueW=local.normal(size=(16,8)); Y=X@trueW + .05*local.normal(size=(256,8))
 98    runs={
 99      'Muon':optimizer_run('muon',X,Y,beta=.95),
100      'BiMaxwell_(.90,.99,.5)':optimizer_run('bi',X,Y,bf=.90,bs=.99,w=.5),
101      'BiMaxwell_(.95,.995,.5)':optimizer_run('bi',X,Y,bf=.95,bs=.995,w=.5),
102      'AdamW':optimizer_run('adamw',X,Y),
103    }
104    result={'seed':SEED,'predictions':{
105      'A_half_life':{'beta_fast':bf,'beta_slow':bs,'predicted_steps_fast':pred_hf,'observed_steps_fast':hf,'predicted_steps_slow':pred_hs,'observed_steps_slow':hs},
106      'B_impulse_decay':impulse,'C_frequency_gain':freq_rows},'mini_experiment':runs}
107    print(json.dumps(result,indent=2))
108
109if __name__=='__main__': main()