Bi-Maxwell Muon / experiment.py
Unverified
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()