Minimum-motion curvature-targeted preconditioner / experiment.py

Mechanism failed

Raw ⬇ ZIP
 1import json, math, random, time
 2import numpy as np
 3import torch
 4
 5SEED = 17
 6np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
 7torch.set_num_threads(4)
 8
 9# Diagonal specialization: G=diag(exp(g)), H=diag(exp(h)).
10def plan_metric(g0, hhat, K=2.0, T=6, alpha=0.15, beta=12.0, iters=30):
11    g0t = torch.tensor(g0, dtype=torch.float64)
12    ht = torch.tensor(hhat, dtype=torch.float64)
13    z = g0t.repeat(T, 1).clone().detach().requires_grad_(True)
14    opt = torch.optim.Adam([z], lr=0.08)
15    for _ in range(iters):
16        opt.zero_grad()
17        spread = (ht-z[-1]).max() - (ht-z[-1]).min()
18        viol = torch.relu(spread - math.log(K))
19        kinetic = ((z[1:]-z[:-1])**2).sum()/(2*T)
20        loss = beta*viol**2 + alpha*kinetic
21        loss.backward(); opt.step()
22    return z.detach().numpy()
23
24def affine_dist(A, B):
25    w,V=np.linalg.eigh(A); Ai=V@np.diag(w**-.5)@V.T
26    q=np.linalg.eigvalsh(Ai@B@Ai)
27    return float(np.linalg.norm(np.log(q)))
28
29def math_check():
30    g=np.array([-.7,.2,1.1]); g2=np.array([.4,-.3,.8])
31    exact=np.linalg.norm(g2-g)
32    got=affine_dist(np.diag(np.exp(g)),np.diag(np.exp(g2)))
33    A=np.array([[2.0,.4],[.4,1.0]]); B=np.array([[1.3,.2],[.2,2.4]])
34    Q=np.linalg.qr(np.random.randn(2,2))[0]; scale=3.7
35    inv_err=abs(affine_dist(A,B)-affine_dist(scale*Q@[email protected],scale*Q@[email protected]))
36    T=6; path=np.linspace(g,g2,T+1); ds=[np.linalg.norm(path[i+1]-path[i]) for i in range(T)]
37    kinetic=sum(x*x for x in ds)/(2*T); length=sum(ds)
38    identity_err=abs(kinetic-length*length/(2*T*T))
39    # Controller should reduce terminal condition number toward K from a bad metric.
40    h=np.log(np.array([1., 10., 100., 1000.]))
41    g0=np.zeros(4); planned=plan_metric(g0,h,K=3,T=6)
42    before=math.exp((h-g0).max()-(h-g0).min())
43    after=math.exp((h-planned[-1]).max()-(h-planned[-1]).min())
44    return dict(diagonal_distance_error=abs(exact-got), affine_invariance_error=inv_err,
45                kinetic_identity_error=identity_err, condition_before=before,
46                condition_after=after, target=3.0, math_pass=(abs(exact-got)<1e-10 and inv_err<1e-9 and identity_err<1e-10 and after<before))
47
48def run(kind, n=10, steps=500, lr=.07, noise=.35):
49    # Fixed ill-conditioned quadratic with noisy gradients; metric updates every R steps.
50    eig=np.geomspace(1.,100.,n); h=np.log(eig); x=np.ones(n)*2.; gmetric=np.zeros(n)
51    ema=np.zeros(n); prev_metric=gmetric.copy(); jumps=[]; conds=[]; losses=[]; gradvars=[]
52    t0=time.perf_counter(); window=[]
53    for step in range(steps):
54        grad=eig*x
55        noisy=grad + noise*np.sqrt(eig)*np.random.randn(n)
56        window.append(noisy.copy())
57        if kind=='idea' and step%20==0:
58            # curvature estimate from noisy diagonal gradient/parameter statistics
59            hhat=h + .10*np.random.randn(n)
60            plan=plan_metric(gmetric,hhat,K=8,T=4,alpha=.25,beta=8,iters=18)
61            newg=plan[0]
62            jumps.append(float(np.linalg.norm(newg-gmetric)))
63            gmetric=newg
64        elif kind=='ema' and step%20==0:
65            # standard EMA diagonal second moment, normalized only for scale
66            if len(window):
67                sq=np.mean(np.array(window[-20:])**2,axis=0)
68                ema=.95*ema+.05*sq
69            newg=.5*np.log(ema+1e-4); newg-=newg.mean()
70            jumps.append(float(np.linalg.norm(newg-gmetric)))
71            gmetric=newg
72        # Use normalized metrics so the comparison is about shape, not global scale.
73        inv=np.exp(-gmetric); inv=inv/np.mean(inv)
74        x -= lr*inv*noisy/np.sqrt(eig.mean())
75        losses.append(float(.5*np.sum(eig*x*x)))
76        M=eig*np.exp(-gmetric); conds.append(float(M.max()/M.min()))
77        gradvars.append(float(np.var(noisy)))
78    return dict(final_loss=losses[-1], best_loss=min(losses), final_condition=conds[-1],
79                mean_condition=float(np.mean(conds[-100:])), mean_jump=float(np.mean(jumps)),
80                gradient_variance=float(np.mean(gradvars[-100:])), seconds=time.perf_counter()-t0,
81                losses=losses, conditions=conds)
82
83def main():
84    check=math_check()
85    # Reset randomness before fair paired runs.
86    np.random.seed(SEED+1); base=run('ema')
87    np.random.seed(SEED+1); idea=run('idea')
88    repeats=[]
89    for seed in [3, 11, 29, 47]:
90        np.random.seed(seed); b=run('ema')
91        np.random.seed(seed); q=run('idea')
92        repeats.append({'seed':seed, 'baseline':{k:b[k] for k in ('final_loss','best_loss','mean_condition','mean_jump','gradient_variance','seconds')}, 'idea':{k:q[k] for k in ('final_loss','best_loss','mean_condition','mean_jump','gradient_variance','seconds')}})
93    out={'math_check':check,'baseline':base,'idea':idea,'repeats':repeats}
94    with open('results.json','w') as f: json.dump(out,f,indent=2)
95    print(json.dumps({'math_check':check,'baseline':{k:v for k,v in base.items() if k not in ('losses','conditions')},'idea':{k:v for k,v in idea.items() if k not in ('losses','conditions')},'repeats':repeats},indent=2))
96
97if __name__=='__main__': main()