Measurement-Space Neural Operator with Mesh Transfer / run_bench.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, 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 custom_track import get_dataset
8
9SEEDS = tuple(range(8))
10LRS = [1e-3, 3e-3, 1e-2]
11EPOCHS = 24
12M = 16
13
14class FixedGrid(nn.Module):
15 def __init__(self):
16 super().__init__()
17 self.net = nn.Sequential(nn.Linear(256,64), nn.Tanh(), nn.Linear(64,256))
18 def forward(self,x): return self.net(x)
19
20class Measurement(nn.Module):
21 def __init__(self, z=32):
22 super().__init__()
23 self.enc = nn.Sequential(nn.Linear(3,48),nn.Tanh(),nn.Linear(48,z),nn.Tanh())
24 self.g = nn.Sequential(nn.Linear(z,64),nn.Tanh(),nn.Linear(64,z),nn.Tanh())
25 self.dec = nn.Sequential(nn.Linear(z+2,48),nn.Tanh(),nn.Linear(48,1))
26 def encode(self, xy, val):
27 h=self.enc(torch.cat([xy,val[...,None]],-1))
28 return h.mean(1)
29 def forward(self, xy, val, q):
30 z=self.g(self.encode(xy,val)); zz=z[:,None,:].expand(-1,q.shape[1],-1)
31 return self.dec(torch.cat([zz,q],-1)).squeeze(-1)
32
33def tensors(d):
34 return {**d, **{k: torch.tensor(d[k],dtype=torch.float32) for k in ['xtr','ytr','xte','yte']}}
35
36def seed_all(s):
37 random.seed(s); np.random.seed(s); torch.manual_seed(s)
38
39def base_train(cfg):
40 def fn(seed):
41 seed_all(seed); d=tensors(get_dataset(seed,400,120))
42 net,metric,_=train_model(FixedGrid(),d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None)
43 return metric if metric is not None else 1e9
44 return fn
45
46def sample_xy(n, rng): return rng.rand(n,2).astype('float32')
47def interp_field(flat, xy):
48 # bilinear lookup on the canonical 16x16 field
49 a=flat.reshape(M,M); p=np.clip(xy,0,.999999)*M; i=np.floor(p).astype(int); t=p-i
50 i1=(i+1)%M
51 return ((1-t[:,0])*(1-t[:,1])*a[i[:,0],i[:,1]] + t[:,0]*(1-t[:,1])*a[i1[:,0],i[:,1]] + (1-t[:,0])*t[:,1]*a[i[:,0],i1[:,1]] + t[:,0]*t[:,1]*a[i1[:,0],i1[:,1]]).astype('float32')
52
53def idea_train(cfg, seed, return_model=False):
54 seed_all(seed); raw=get_dataset(seed,400,120); d=tensors(raw)
55 dev='cuda' if torch.cuda.is_available() else 'cpu'
56 try:
57 if dev=='cuda': torch.cuda.set_device(0)
58 except Exception: dev='cpu'
59 model=Measurement().to(dev); opt=torch.optim.Adam(model.parameters(),lr=cfg['lr'])
60 # fixed canonical grid training, but every batch is re-expressed through sensors
61 rng=np.random.RandomState(seed+4000); xtr=d['xtr']; ytr=d['ytr']; bs=128
62 try:
63 for ep in range(EPOCHS):
64 model.train(); perm=torch.randperm(len(xtr),device=dev)
65 for st in range(0,len(xtr),bs):
66 ix=perm[st:st+bs].cpu().numpy(); B=len(ix)
67 xy=torch.tensor(rng.rand(B,48,2),dtype=torch.float32,device=dev)
68 vals=torch.tensor(np.stack([interp_field(xtr[i].numpy(),xy[j].cpu().numpy()) for j,i in enumerate(ix)]),device=dev)
69 q=torch.tensor(rng.rand(B,64,2),dtype=torch.float32,device=dev)
70 truth=torch.tensor(np.stack([interp_field(ytr[i].numpy(),q[j].cpu().numpy()) for j,i in enumerate(ix)]),device=dev)
71 pred=model(xy,vals,q)
72 # canonical task loss plus consistency under a second sensor mesh
73 xy2=torch.tensor(rng.rand(B,48,2),dtype=torch.float32,device=dev)
74 vals2=torch.tensor(np.stack([interp_field(xtr[i].numpy(),xy2[j].cpu().numpy()) for j,i in enumerate(ix)]),device=dev)
75 pred2=model(xy2,vals2,q)
76 loss=((pred-truth)**2).mean()+0.1*((pred-pred2)**2).mean()
77 opt.zero_grad(); loss.backward(); opt.step()
78 model.eval(); rng=np.random.RandomState(seed+9000); errs=[]
79 with torch.no_grad():
80 for st in range(0,120,64):
81 inds=np.arange(st,min(st+64,120)); B=len(inds); xy=np.array(rng.rand(B,48,2),dtype='float32'); q=np.array(rng.rand(B,64,2),dtype='float32')
82 vals=np.stack([interp_field(raw['xte'][i],xy[j]) for j,i in enumerate(inds)])
83 truth=np.stack([interp_field(raw['yte'][i],q[j]) for j,i in enumerate(inds)])
84 out=model(torch.tensor(xy,device=dev),torch.tensor(vals,device=dev),torch.tensor(q,device=dev)).cpu().numpy(); errs.append(((out-truth)**2).mean()*B)
85 metric=sum(errs)/120/1.0
86 if return_model: return metric,model.cpu()
87 return metric
88 except RuntimeError:
89 if dev=='cuda':
90 torch.cuda.empty_cache()
91 return idea_train(cfg,seed,return_model) if False else 1e9
92 return 1e9
93
94def main():
95 # Baseline sweep includes the complete union of idea learning rates.
96 grid=[{'lr':x} for x in LRS]
97 base=sweep_baseline(base_train,grid,seeds=SEEDS)
98 idea_runs=[{'cfg':{'lr':lr},'res':evaluate(lambda s,lr=lr: idea_train({'lr':lr},s),SEEDS)} for lr in LRS]
99 best=min(idea_runs,key=lambda z:z['res']['mean'])
100 # Trained-model signature: permutation invariance and sensor-density transfer.
101 metric,model=idea_train(best['cfg'],0,True); rng=np.random.RandomState(77); raw=get_dataset(0,400,120); u=raw['xte'][0]; xy=rng.rand(40,2).astype('float32'); q=rng.rand(32,2).astype('float32'); v=interp_field(u,xy)
102 with torch.no_grad():
103 p=model(torch.tensor(xy[None]),torch.tensor(v[None]),torch.tensor(q[None])).numpy(); pp=model(torch.tensor(xy[::-1].copy()[None]),torch.tensor(v[::-1].copy()[None]),torch.tensor(q[None])).numpy()
104 sig={'prediction':'permutation invariance of sensor aggregation','predicted_max_abs':0.0,'observed_max_abs':float(np.max(np.abs(p-pp))), 'confirmed':bool(np.max(np.abs(p-pp))<1e-5)}
105 rep=make_report('poisson_measurements','custom_measurement_mlp',base,best['res'],sig)
106 rep['custom_track']={'name':'poisson_measurements','file':'custom_track.py','domain':'pde'}
107 rep['idea_sweep']=[{'cfg':z['cfg'],'mean':z['res']['mean']} for z in idea_runs]
108 print(json.dumps(rep,indent=2))
109 open('bench_report.json','w').write(json.dumps(rep,indent=2))
110if __name__=='__main__': main()