Exact Neural de Rham Backbone / run_bench_registered.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys,json,random,math
 2import numpy as np
 3import torch
 4from torch import nn
 5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
 7from bench.protocol import DEFAULT_SEEDS
 8
 9TRACK='poisson_boundary'; MODEL='scalar_potential_backbone'; EPOCHS=35; NTRAIN=400; NTEST=400
10LRS=[0.003,0.01,0.03]
11
12def math_check():
13    rng=np.random.default_rng(1286); w=rng.normal(size=(7,2)); w/=np.linalg.norm(w,axis=1,keepdims=True)
14    nil=0.0
15    for wi in w:
16        k0=wi[:,None]; k1=np.array([[-wi[1],wi[0]]]); nil=max(nil,float(np.max(np.abs(k1@k0))))
17    x=torch.tensor([[.37,-.21]],dtype=torch.float64,requires_grad=True); wt=torch.tensor(w[:1],dtype=torch.float64); b=torch.tensor([.41],dtype=torch.float64)
18    s=x@wt.T+b; y=torch.relu(s)**3/6; g=torch.autograd.grad(y.sum(),x)[0].detach().numpy()[0]
19    rhs=(torch.relu(s.detach())**2/2).numpy()[0,0]*w[0]
20    return {'nilpotence_max_abs':nil,'derivative_identity_max_abs':float(np.max(np.abs(g-rhs)))}
21
22class ExactPotential(nn.Module):
23    def __init__(self,w,b,k=3):
24        super().__init__(); self.register_buffer('w',torch.tensor(w)); self.register_buffer('b',torch.tensor(b)); self.coef=nn.Parameter(torch.zeros(len(w))); self.k=k
25    def forward(self,x):
26        s=x@self.w.T+self.b
27        return (torch.relu(s)**self.k/math.factorial(self.k) @ self.coef)[:,None]
28
29class StandardMLP(nn.Module):
30    def __init__(self):
31        super().__init__(); self.net=nn.Sequential(nn.Linear(2,16),nn.Tanh(),nn.Linear(16,16),nn.Tanh(),nn.Linear(16,1))
32    def forward(self,x): return self.net(x)
33
34def run(kind,lr,seed,signature=False):
35    np.random.seed(seed); random.seed(seed); torch.manual_seed(seed)
36    d=get_dataset(TRACK,seed,NTRAIN,NTEST); ds=dict(d)
37    for q in ('xtr','ytr','xte','yte'):
38        ds[q]=d[q] if torch.is_tensor(d[q]) else torch.from_numpy(d[q])
39    rng=np.random.default_rng(9001+seed*7919); w=rng.normal(size=(32,2)).astype('float32'); w/=np.linalg.norm(w,axis=1,keepdims=True); b=rng.uniform(-1,1,32).astype('float32')
40    net=ExactPotential(w,b) if kind=='idea' else StandardMLP()
41    net,metric,_=train_model(net,ds,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
42    if net is None:return float('nan'),None
43    sig=None
44    if signature:
45        dev=next(net.parameters()).device; x=torch.as_tensor(d['xte'][:128],device=dev).requires_grad_(True); out=net(x)
46        grad=torch.autograd.grad(out.sum(),x,create_graph=True)[0]
47        c0=torch.autograd.grad(grad[:,0].sum(),x,retain_graph=True)[0]; c1=torch.autograd.grad(grad[:,1].sum(),x)[0]
48        curl=c1[:,0]-c0[:,1]; val=float(torch.sqrt(torch.mean(curl.detach()**2)))
49        sig={'predicted_gradient_curl_rms':0.0,'observed_gradient_curl_rms':val,'n_probe':128,'confirmed':val<1e-5}
50    return float(metric),sig
51
52def factory(kind):
53    return lambda cfg: (lambda seed: run(kind,float(cfg['lr']),seed)[0])
54
55if __name__=='__main__':
56    mc=math_check()
57    base=sweep_baseline(factory('baseline'),[{'lr':x} for x in LRS])
58    trials=[]
59    for x in LRS:
60        r=evaluate(lambda seed,x=x:run('idea',x,seed)[0],DEFAULT_SEEDS); trials.append({'cfg':{'lr':x},'result':r})
61    best=min(trials,key=lambda z:z['result']['mean']); idea=best['result']
62    sig=run('idea',best['cfg']['lr'],0,True)[1]
63    bsig=run('baseline',base['best_cfg']['lr'],0,True)[1]
64    rep=make_report(TRACK,MODEL,base,idea,{'stage1_prediction':'d squared equals zero for the learned potential gradient','baseline_gradient_curl_rms':bsig['observed_gradient_curl_rms'],**sig,'idea_sweep':trials,'math_check':mc})
65    rep['math_check']=mc; rep['idea_selected_cfg']=best['cfg']
66    with open('bench_report_registered.json','w') as f: json.dump(rep,f,indent=2)
67    print(json.dumps(rep,indent=2))