Bi-Maxwell Muon / verify.py
Unverified
1import json, math, time
2import numpy as np
3from bi_maxwell_muon import BiMaxwellMuon, polar_newton_schulz
4
5SEED = 2693
6
7def first_half_step(beta):
8 n = 0
9 while 1-beta**(n+1) < .5:
10 n += 1
11 return n
12
13def run(kind, X, Y, steps=350, lr=.025, beta=.95, bf=.9, bs=.99, w=.5):
14 W=np.zeros((16,8)); m=np.zeros_like(W); mf=np.zeros_like(W); ms=np.zeros_like(W)
15 v=np.zeros_like(W); losses=[]; cs=[]; start=time.perf_counter()
16 for _ in range(steps):
17 pred=X@W; g=X.T@(pred-Y)/len(X)
18 if kind=='adamw':
19 m=.9*m+.1*g; v=.999*v+.001*g*g; upd=m/(np.sqrt(v)+1e-8)
20 elif kind=='muon':
21 m=beta*m+(1-beta)*g; upd=polar_newton_schulz(m)
22 else:
23 mf=bf*mf+(1-bf)*g; ms=bs*ms+(1-bs)*g
24 upd=polar_newton_schulz(w*mf+(1-w)*ms)
25 W -= lr*upd
26 losses.append(float(np.mean((pred-Y)**2)/2))
27 cs.append(float((g*upd).sum()/(np.linalg.norm(g)*np.linalg.norm(upd)+1e-12)))
28 return {'final_loss':losses[-1], 'loss_step_100':losses[99],
29 'mean_gradient_update_cosine':float(np.mean(cs)),
30 'seconds':time.perf_counter()-start}
31
32def main():
33 eta=1.; tf=2.; ts=20.; bf=math.exp(-eta/tf); bs=math.exp(-eta/ts)
34 half=[]
35 for b in [bf,bs,.8,.9,.99]:
36 pred=math.log(.5)/math.log(b)-1
37 half.append({'beta':b, 'predicted_continuous_index':pred,
38 'predicted_first_integer':math.ceil(pred),
39 'observed_first_integer':first_half_step(b)})
40 slopes=[]
41 for b in [.5,.8,.9,.95,.99]:
42 n=np.arange(1,51); y=b**n
43 obs=float(np.polyfit(n,np.log(y),1)[0])
44 slopes.append({'beta':b,'predicted_log_slope':math.log(b),'observed_log_slope':obs})
45 def gain(b,o): return (1-b)/math.sqrt(1+b*b-2*b*math.cos(o))
46 freqs=[]
47 for o in [.1,.5,1.5,2.5]:
48 n=6000; t=np.arange(n); z=np.cos(o*t); a=c=0.; out=[]
49 for q in z:
50 a=bf*a+(1-bf)*q; c=bs*c+(1-bs)*q; out.append(.5*a+.5*c)
51 tail=slice(1000,None); tt=t[tail]; yy=np.asarray(out)[tail]
52 observed=2*abs(np.mean(yy*np.exp(-1j*o*tt)))
53 predicted=.5*gain(bf,o)+.5*gain(bs,o)
54 freqs.append({'omega':o,'predicted_gain':predicted,'observed_gain':float(observed),
55 'relative_error':float(abs(observed-predicted)/predicted)})
56 r=np.random.default_rng(SEED); X=r.normal(size=(256,16)); true=r.normal(size=(16,8)); Y=X@true+.05*r.normal(size=(256,8))
57 runs={'Muon':run('muon',X,Y), 'BiMaxwell_(.90,.99,.5)':run('bi',X,Y,bf=.9,bs=.99),
58 'BiMaxwell_(.95,.995,.5)':run('bi',X,Y,bf=.95,bs=.995), 'AdamW':run('adamw',X,Y)}
59 print(json.dumps({'predictions':{'half_life':half,'impulse_log_slope':slopes,'frequency_gain':freqs},'mini_experiment':runs},indent=2))
60if __name__=='__main__': main()