Error-budgeted local log-signature tokens / stage2_bench.py
Failed on benchmark
1import sys, json, math, time, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, sweep_baseline, evaluate, make_report
8
9SEEDS = tuple(range(8))
10EPOCHS = 8
11NTR, NTE = 800, 300
12
13def tail(z, n):
14 term, partial = 1.0, 1.0
15 for k in range(1, n+1):
16 term *= z/k; partial += term
17 return math.exp(z)-partial
18
19def lie_dim(k, d):
20 def mu(q):
21 x=q; p=2; nd=0
22 while p*p <= x:
23 if x%p==0:
24 x//=p; nd+=1
25 if x%p==0: return 0
26 while x%p==0: x//=p
27 p+=1
28 if x>1: nd+=1
29 return -1 if nd%2 else 1
30 return sum(mu(q)*d**(k//q) for q in range(1,k+1) if k%q==0)//k
31
32def sanity():
33 # Cheap numerical verification before any neural training.
34 deg = [[tail(z,n) for n in range(1,6)] for z in (.5, 1., 2.)]
35 mono_n = all(a>b for row in deg for a,b in zip(row,row[1:]))
36 ms = [m*tail(6/m,2) for m in (1,2,4,8,16)]
37 mono_m = all(a>b for a,b in zip(ms,ms[1:]))
38 costs = [m*sum(lie_dim(k,2) for k in range(1,n+1)) for m,n in ((1,3),(2,2),(4,2),(8,1))]
39 return {'tail_decreases_with_degree': mono_n, 'equal_variation_error_decreases_with_m': mono_m,
40 'tail_values': deg, 'equal_m_values': ms, 'lie_dims_d2': [lie_dim(k,2) for k in range(1,6)],
41 'cost_examples': costs}
42
43class TokenTransformer(nn.Module):
44 def __init__(self, feat, tokens, d=48):
45 super().__init__(); self.inp=nn.Linear(feat,d)
46 self.pos=nn.Parameter(torch.zeros(1,tokens,d)); nn.init.normal_(self.pos,std=.02)
47 layer=nn.TransformerEncoderLayer(d,nhead=2,dim_feedforward=96,batch_first=True,dropout=0.)
48 self.enc=nn.TransformerEncoder(layer,2); self.head=nn.Linear(tokens*d,1)
49 def forward(self,x):
50 h=self.inp(x)+self.pos[:,:x.shape[1]]
51 return self.head(self.enc(h).reshape(x.shape[0],-1))
52
53def local_sig(x, m, adaptive=True):
54 # Add time as a second control channel; degree-2 log signature is its signed area.
55 n=len(x); t=np.linspace(0,1,n,dtype=np.float32)
56 path=np.stack([t,x],1); ds=np.linalg.norm(np.diff(path,axis=0),axis=1)
57 v=np.r_[0.,np.cumsum(ds)]; total=v[-1]
58 if adaptive and total>1e-9:
59 cuts=[0]+[int(np.searchsorted(v,q*total/m,'left')) for q in range(1,m)]+[n-1]
60 else: cuts=list(np.linspace(0,n-1,m+1).astype(int))
61 cuts=np.maximum.accumulate(np.asarray(cuts)); cuts[-1]=n-1
62 out=[]; ells=[]
63 for a,b in zip(cuts[:-1],cuts[1:]):
64 dx=np.diff(path[a:b+1],axis=0); inc=path[b]-path[a]
65 pref=np.zeros(2); area=0.
66 for u in dx: area += .5*(pref[0]*u[1]-u[0]*pref[1]); pref += u
67 ell=float(v[b]-v[a]); ells.append(ell)
68 out.append([inc[0],inc[1],area,float(b-a)/max(1,n-1)])
69 return np.asarray(out,np.float32),np.asarray(ells),cuts
70
71def raw_tokens(x,m=8):
72 # Fixed standard patch interface, same four-channel width as idea.
73 n=len(x); out=[]
74 for a,b in zip(np.linspace(0,n,m+1)[:-1].astype(int),np.linspace(0,n,m+1)[1:].astype(int)):
75 b=max(b,a+1); p=x[a:b]; out.append([p.mean(),p[-1]-p[0],p.std(),float(b-a)/n])
76 return np.asarray(out,np.float32)
77
78def transform_ds(ds, kind, m):
79 fn=(lambda x: raw_tokens(x,m)) if kind=='baseline' else (lambda x: local_sig(x,m,True)[0])
80 tr=np.stack([fn(x) for x in ds['xtr'].numpy()]); te=np.stack([fn(x) for x in ds['xte'].numpy()])
81 return {'xtr':torch.tensor(tr), 'ytr':ds['ytr'].float(), 'xte':torch.tensor(te), 'yte':ds['yte'].float(), 'task':'regression','metric':'mse','input_shape':tr.shape[1:],'out_dim':1}
82
83def run_one(seed, kind, lr, m):
84 torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
85 ds=get_dataset('sequence',seed,n_train=NTR,n_test=NTE)
86 z=transform_ds(ds,kind,m); net=TokenTransformer(z['input_shape'][1],z['input_shape'][0])
87 _, metric, _=train_model(net,z,epochs=EPOCHS,lr=lr,batch=128,log=lambda *_:None)
88 return float(metric)
89
90def main():
91 sanity_result=sanity(); print('SANITY',json.dumps(sanity_result))
92 # Union parity: every idea lr is in baseline grid; baseline method knob m is also swept.
93 grid=[{'lr':1e-3,'m':4},{'lr':3e-3,'m':8},{'lr':6e-3,'m':8}]
94 def maker(cfg): return lambda seed: run_one(seed,'baseline',cfg['lr'],cfg['m'])
95 base=sweep_baseline(maker,grid,seeds=(0,1,2,3))
96 # Evaluate idea at all three shared settings, then select best by the same four-seed tuning split.
97 idea_trials=[]
98 for cfg in grid:
99 r=evaluate(lambda seed,cfg=cfg: run_one(seed,'idea',cfg['lr'],cfg['m']),seeds=(0,1,2,3))
100 idea_trials.append({'cfg':cfg,'mean':r['mean']})
101 best=min(idea_trials,key=lambda q:q['mean'])['cfg']
102 idea=evaluate(lambda seed: run_one(seed,'idea',best['lr'],best['m']),seeds=SEEDS)
103 # Trained-model behavior signature: observed residual versus the representation's predicted tail budget.
104 vals=[]
105 for s in SEEDS:
106 ds=get_dataset('sequence',s,n_train=NTR,n_test=NTE); z=transform_ds(ds,'idea',best['m'])
107 # proxy uses observed local variation, while residuals are from the trained system.
108 _,metric,_=train_model(TokenTransformer(z['input_shape'][1],z['input_shape'][0]),z,epochs=EPOCHS,lr=best['lr'],batch=128,log=lambda *_:None)
109 e=[]
110 for x in ds['xte'].numpy():
111 _,ells,_=local_sig(x,best['m'],True); e.append(sum(tail(float(q),2) for q in ells))
112 vals.append((float(np.mean(e)),float(metric)))
113 corr=float(np.corrcoef(np.asarray(vals).T)[0,1]) if len(vals)>1 else float('nan')
114 sig={'prediction':'larger local Taylor tail should proxy larger CDE approximation error',
115 'predicted_tail_mean':float(np.mean([v[0] for v in vals])),
116 'observed_test_mse_mean':float(np.mean([v[1] for v in vals])),
117 'seed_level_correlation':corr,'confirmed':False,
118 'note':'Both quantities are measured on trained benchmark systems; eight seed points do not establish the claimed quantitative relation.'}
119 report=make_report('sequence','transformer_tiny',base,idea,{'math_sanity':sanity_result,'idea_trials':idea_trials,'selected_cfg':best,'mechanism_signature':sig})
120 report['baseline']['idea_union_grid']=grid
121 Path('bench_report.json').write_text(json.dumps(report,indent=2))
122 print(json.dumps(report,indent=2))
123if __name__=='__main__': main()