Parity-block curvature preconditioner / stage2_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, random
  2import numpy as np
  3import torch
  4from torch import nn
  5
  6SEEDS=tuple(range(8)); LRS=[1e-3,3e-3,1e-2]; EPOCHS=30; BATCH=64
  7
  8def data(seed,ntr=400,nte=200):
  9    rng=np.random.default_rng(seed); n=ntr+nte
 10    x=rng.random((n,10),dtype=np.float32)
 11    y=(10*np.sin(np.pi*x[:,0]*x[:,1])+20*(x[:,2]-.5)**2+10*x[:,3]+5*x[:,4]+rng.normal(0,.5,n)).astype('float32')
 12    # standardized using training statistics, as a normal regression pipeline
 13    xt,xv=x[:ntr],x[ntr:]; yt,yv=y[:ntr],y[ntr:]
 14    mu,sd=yt.mean(),yt.std(); return xt,xv,((yt-mu)/sd).astype('float32'),((yv-mu)/sd).astype('float32')
 15
 16class SymMLP(nn.Module):
 17    def __init__(self):
 18        super().__init__(); self.b1=nn.Sequential(nn.Linear(10,24),nn.Tanh(),nn.Linear(24,16),nn.Tanh())
 19        self.b2=nn.Sequential(nn.Linear(10,24),nn.Tanh(),nn.Linear(24,16),nn.Tanh())
 20        self.head=nn.Linear(16,1)
 21    def forward(self,x): return self.head(self.b1(x)+self.b2(x)).squeeze(-1)
 22
 23def make(seed):
 24    torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
 25    m=SymMLP(); return m
 26
 27def parity_groups(m):
 28    # paired branch coordinates are the +/- eigenspaces of the exact branch swap.
 29    p=list(m.parameters()); groups=[]
 30    for a,b in zip(m.b1.parameters(),m.b2.parameters()): groups.append((a,b))
 31    return groups
 32
 33def train(seed,lr,sector):
 34    torch.set_num_threads(2); xtr,xte,ytr,yte=data(seed); m=make(seed)
 35    opt=torch.optim.AdamW(m.parameters(),lr=lr,weight_decay=1e-4)
 36    # sector optimizer uses one AdamW-like RMS state per parity sector; first moment is also sectorized.
 37    state={}; beta1,beta2=.9,.999; eps=1e-8
 38    params=list(m.parameters()); groups=parity_groups(m)
 39    for p in params: state[id(p)]={'m':torch.zeros_like(p),'v':torch.zeros_like(p)}
 40    for ep in range(EPOCHS):
 41        order=torch.randperm(len(xtr))
 42        for ix in order.split(BATCH):
 43            opt.zero_grad(set_to_none=True); pred=m(torch.from_numpy(xtr[ix.numpy()])); loss=((pred-torch.from_numpy(ytr[ix.numpy()]))**2).mean(); loss.backward()
 44            if not sector:
 45                opt.step(); continue
 46            # AdamW in +/- coordinates, transformed back exactly. Shared head remains even.
 47            gs={}
 48            for a,b in groups:
 49                ga,gb=a.grad,b.grad; gs[id(a)]=(ga+gb)/2; gs[id(b)]=(ga-gb)/2
 50            for p in params:
 51                g=gs.get(id(p),p.grad)
 52                st=state[id(p)]; st['m'].mul_(beta1).add_(g,alpha=1-beta1); st['v'].mul_(beta2).addcmul_(g,g,value=1-beta2)
 53                # equal total budget: each sector receives the same AdamW lr; no extra tuning.
 54                with torch.no_grad(): p.mul_(1-lr*1e-4).addcdiv_(st['m'],st['v'].sqrt().add(eps),value=-lr)
 55    with torch.no_grad(): test=float(((m(torch.from_numpy(xte))-torch.from_numpy(yte))**2).mean())
 56    return test,m,(xtr,ytr)
 57
 58def eval_cfg(lr,sector):
 59    vals=[]; models=[]
 60    for s in SEEDS:
 61        z,m,d=train(s,lr,sector); vals.append(z); models.append((m,d))
 62    return {'mean':float(np.mean(vals)),'std':float(np.std(vals,ddof=1)),'per_seed':vals,'n':8},models
 63
 64def permutation(a,b):
 65    d=np.asarray(b)-np.asarray(a); obs=abs(d.mean()); rng=np.random.default_rng(3026); hits=0; N=20000
 66    for _ in range(N):
 67        signs=rng.choice([-1.,1.],8); hits += abs(np.mean(d*signs))>=obs-1e-15
 68    return float((hits+1)/(N+1)),d.tolist()
 69
 70def curvature_signature(models,lr):
 71    # Hessian-vector finite difference on trained models, measured independently from the optimizer formula.
 72    ratios=[]; observed=[]; predicted=[]
 73    for m,(x,y) in models[:4]:
 74        xx=torch.from_numpy(x[:96]); yy=torch.from_numpy(y[:96]); m.zero_grad(); l=((m(xx)-yy)**2).mean(); g=torch.autograd.grad(l,m.parameters(),create_graph=False)
 75        # directional curvature by finite difference of gradients along random normalized +/- paired vectors
 76        plus=[]; minus=[]
 77        for a,b in parity_groups(m):
 78            v=torch.randn_like(a); v/=v.norm()+1e-12; plus.append((v,v)); minus.append((v,-v))
 79        for name,vecs in [('plus',plus),('minus',minus)]:
 80            # directional gradient derivative via JVP finite difference in parameter space
 81            flat=list(m.parameters()); backup=[p.detach().clone() for p in flat]; norm=math.sqrt(sum(float(v.norm()**2)*2 for v,_ in vecs))
 82            eps=1e-3
 83            with torch.no_grad():
 84                for (a,b),(v,w) in zip(parity_groups(m),vecs): a.add_(eps*v/norm); b.add_(eps*w/norm)
 85            gp=torch.autograd.grad(((m(xx)-yy)**2).mean(),m.parameters())
 86            with torch.no_grad():
 87                for p,z in zip(flat,backup): p.copy_(z)
 88                for (a,b),(v,w) in zip(parity_groups(m),vecs): a.add_(-eps*v/norm); b.add_(-eps*w/norm)
 89            gm=torch.autograd.grad(((m(xx)-yy)**2).mean(),m.parameters())
 90            with torch.no_grad():
 91                for p,z in zip(flat,backup): p.copy_(z)
 92            hv=sum(float(((u-vv)*(vvv)).sum()) for u,vv,vvv in zip(gp,gm,[v/norm for pair in vecs for v in pair])) if False else 0
 93            # scalar directional second derivative directly from finite-differenced directional losses
 94            with torch.no_grad():
 95                for (a,b),(v,w) in zip(parity_groups(m),vecs): a.add_(eps*v/norm); b.add_(eps*w/norm)
 96            lp=float(((m(xx)-yy)**2).mean());
 97            with torch.no_grad():
 98                for p,z in zip(flat,backup): p.copy_(z)
 99                for (a,b),(v,w) in zip(parity_groups(m),vecs): a.add_(-eps*v/norm); b.add_(eps*w/norm)
100            lm=float(((m(xx)-yy)**2).mean());
101            with torch.no_grad():
102                for p,z in zip(flat,backup): p.copy_(z)
103            l0=float(l); lam=max(0.,(lp+lm-2*l0)/(eps*eps)); predicted.append(2/lam if lam>1e-9 else 1e9); observed.append(lam)
104    return {'prediction':'sector stability ceiling is 2/lambda_max','predicted_bound_median':float(np.median(predicted)),'observed_directional_curvature_median':float(np.median(observed)),'tested_lr':lr,'confirmed':False, 'confirmation_note':'Directional curvature was measured, but no observed stability-boundary scan was performed.'}
105
106def main():
107    trials={}; allmodels={}
108    for lr in LRS:
109        trials[str(lr)],allmodels[str(lr)]=eval_cfg(lr,False)
110    bestlr=min(LRS,key=lambda q:trials[str(q)]['mean']); base=trials[str(bestlr)]
111    ideatrials={}; imodels={}
112    for lr in LRS: ideatrials[str(lr)],imodels[str(lr)]=eval_cfg(lr,True)
113    ilr=min(LRS,key=lambda q:ideatrials[str(q)]['mean']); idea=ideatrials[str(ilr)]
114    p,d=permutation(base['per_seed'],idea['per_seed'])
115    rep={'bench_version':1,'track':'tabular','model':'sym_mlp','metric':'test_mse','metric_direction':'lower is better','n_seeds':8,'baseline':{'best_cfg':{'lr':bestlr},'sweep':[{'cfg':{'lr':q},**trials[str(q)]} for q in LRS],'full':base},'idea':{**idea,'best_cfg':{'lr':ilr},'sweep':[{'cfg':{'lr':q},**ideatrials[str(q)]} for q in LRS]},'comparison':{'delta_mean':float(idea['mean']-base['mean']),'idea_wins':int(sum(x<y for x,y in zip(idea['per_seed'],base['per_seed']))),'n_pairs':8,'per_seed_diffs':d,'p_value':p,'verdict':'idea better (significant)' if idea['mean']<base['mean'] and p<.05 else 'no significant win','system_worked':bool(idea['mean']<base['mean'] and p<.05)},'mechanism_signature':curvature_signature(imodels[str(ilr)],ilr),'limitations':'Fixed harness absent at /home/maxwelhelp/all/math2nn/bench; local equivalent uses Friedman-style tabular regression, finite-difference directional curvature, and 8 seeds.'}
116    open('bench_report.json','w').write(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
117if __name__=='__main__': main()