Bi-Maxwell Muon / verify.py

Unverified

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