Exact Neural de Rham Backbone / run_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import sys, os, json, math, random
  2import numpy as np
  3import torch
  4from torch import nn
  5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  6from bench import train_model, evaluate, sweep_baseline, make_report
  7from bench.protocol import DEFAULT_SEEDS
  8import custom_exact_gradient_track as track
  9
 10SEED=1286
 11EPOCHS=35
 12NTRAIN=400
 13NTEST=400
 14NEXACT=32
 15NBASE=16
 16LRS=[0.003,0.01,0.03]
 17
 18def math_check():
 19    rng=np.random.default_rng(SEED); w=rng.normal(size=(7,2)); w/=np.linalg.norm(w,axis=1,keepdims=True)
 20    # 2D Koszul matrices K0: R -> R2 and K1: R2 -> R
 21    nil=0.
 22    for wi in w:
 23        K0=wi[:,None]
 24        K1=np.array([[-wi[1],wi[0]]])
 25        nil=max(nil,float(np.max(np.abs(K1@K0))))
 26    x=torch.tensor([[.37,-.21]],dtype=torch.float64,requires_grad=True)
 27    wt=torch.tensor(w[:1],dtype=torch.float64); b=torch.tensor([.41],dtype=torch.float64)
 28    s=x@wt.T+b; y=torch.relu(s)**3/6
 29    g=torch.autograd.grad(y.sum(),x)[0].detach().numpy()[0]
 30    rhs=(torch.relu(s.detach())**2/2).numpy()[0,0]*w[0]
 31    return {'nilpotence_max_abs':nil,'derivative_identity_max_abs':float(np.max(np.abs(g-rhs)))}
 32
 33def features(x,w,b,m):
 34    s=x@w.T+b
 35    return torch.relu(s)**m/math.factorial(m)
 36
 37class Exact(nn.Module):
 38    def __init__(self,w,b):
 39        super().__init__(); self.register_buffer('w',torch.tensor(w)); self.register_buffer('b',torch.tensor(b)); self.coef=nn.Parameter(torch.zeros(len(w)))
 40    def forward(self,x):
 41        f=features(x,self.w,self.b,1)
 42        return f @ (self.coef[:,None]*self.w)
 43
 44class Independent(nn.Module):
 45    def __init__(self,w,b):
 46        super().__init__(); self.register_buffer('w',torch.tensor(w)); self.register_buffer('b',torch.tensor(b)); self.coef=nn.Parameter(torch.zeros(len(w),2))
 47    def forward(self,x):
 48        return features(x,self.w,self.b,1) @ self.coef
 49
 50def make_system(kind,seed):
 51    rng=np.random.default_rng(SEED+int(seed)*7919)
 52    if kind=='idea':
 53        n=NEXACT
 54    else: n=NBASE
 55    w=rng.normal(size=(n,2)).astype(np.float32); w/=np.linalg.norm(w,axis=1,keepdims=True); b=rng.uniform(-1,1,size=n).astype(np.float32)
 56    return Exact(w,b) if kind=='idea' else Independent(w,b)
 57
 58def run(kind,lr,seed, want_sig=False):
 59    np.random.seed(seed); random.seed(seed); torch.manual_seed(seed)
 60    raw=track.get_dataset(seed,NTRAIN,NTEST)
 61    ds=dict(raw)
 62    for key in ('xtr','ytr','xte','yte'):
 63        ds[key]=torch.from_numpy(raw[key])
 64    net=make_system(kind,seed)
 65    net,metric,hist=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
 66    if net is None: return float('nan'), None
 67    sig=None
 68    if want_sig:
 69        dev=next(net.parameters()).device
 70        x=torch.from_numpy(raw['xte'][:128]).to(dev).requires_grad_(True)
 71        out=net(x)
 72        g0=torch.autograd.grad(out[:,0].sum(),x,retain_graph=True)[0]
 73        g1=torch.autograd.grad(out[:,1].sum(),x)[0]
 74        curl=g1[:,0]-g0[:,1]
 75        sig={'observed_curl_rms':float(torch.sqrt(torch.mean(curl.detach()**2))), 'n_probe':128,
 76             'predicted_curl_rms':0.0, 'confirmed':bool(float(torch.sqrt(torch.mean(curl.detach()**2))) < 1e-6)}
 77    return float(metric),sig
 78
 79def factory(kind):
 80    def make(cfg):
 81        return lambda seed: run(kind,float(cfg['lr']),seed)[0]
 82    return make
 83
 84if __name__=='__main__':
 85    mathres=math_check()
 86    # Baseline sweep includes every idea learning rate; same epochs and model budget.
 87    base=sweep_baseline(factory('baseline'),[{'lr':lr} for lr in LRS])
 88    # Explicit idea sweep at all same settings on full paired seeds.
 89    idea_trials=[]
 90    for lr in LRS:
 91        r=evaluate(lambda seed,lr=lr: run('idea',lr,seed)[0],DEFAULT_SEEDS)
 92        idea_trials.append({'cfg':{'lr':lr},'result':r})
 93    best=min(idea_trials,key=lambda z:z['result']['mean'])
 94    # Signature is measured on the selected trained systems, not an algebra-only toy.
 95    idea_sig=run('idea',best['cfg']['lr'],0,True)[1]
 96    rep=make_report('exact_gradient_harmonic_field','fixed_feature_mlp',base,best['result'],{
 97        'stage1_prediction':'gradient output has zero curl', **idea_sig,
 98        'baseline_observed_curl_rms':run('baseline',base['best_cfg']['lr'],0,True)[1]['observed_curl_rms'],
 99        'idea_sweep':idea_trials,'math_check':mathres})
100    rep['custom_track']={'name':track.META['name'],'file':'custom_exact_gradient_track.py','domain':track.META['domain']}
101    rep['math_check']=mathres
102    rep['idea_selected_cfg']=best['cfg']
103    with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
104    print(json.dumps(rep,indent=2))