Critical-Tail Multiscale Mixer / stage2_bench.py
Beats tuned baseline
1import sys, os, json, math, random
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10# Union of all tried lrs is shared by both systems; baseline's decisive knob is attention type,
11# and the standard Transformer has no additional exposed method knob in this fixed harness.
12GRID = [{'lr': 1e-3, 'epochs': 18}, {'lr': 3e-3, 'epochs': 18}, {'lr': 6e-3, 'epochs': 18}]
13
14class DyadicMixer(nn.Module):
15 def __init__(self, d, nhead=2):
16 super().__init__()
17 self.d = d
18 self.nhead = nhead
19 self.qkv = nn.Linear(d, 3*d)
20 self.out = nn.Linear(d, d)
21 self.norm = nn.LayerNorm(d)
22 self.ff = nn.Sequential(nn.Linear(d, 128), nn.ReLU(), nn.Linear(128, d))
23 self.ffnorm = nn.LayerNorm(d)
24 self.alpha = nn.Parameter(torch.tensor(0.1))
25 self.beta = nn.Parameter(torch.tensor(1.0))
26
27 def forward(self, x):
28 # x is B,L,d. Boundary ranges are clipped and divided by valid counts.
29 h = self.norm(x)
30 v = self.qkv(h).chunk(3, dim=-1)[2]
31 B,L,D = v.shape
32 pref = torch.cat((torch.zeros(B,1,D,device=v.device,dtype=v.dtype), v.cumsum(1)), 1)
33 out = torch.zeros_like(v)
34 K = max(1, int(math.log2(L)))
35 idx = torch.arange(L, device=v.device)
36 for k in range(K):
37 lo, hi = 2**k, min(2**(k+1), L)
38 # left indices [i-hi, i-lo), right [i+lo, i+hi)
39 la = (idx-hi).clamp(0,L); lb = (idx-lo).clamp(0,L)
40 ra = (idx+lo).clamp(0,L); rb = (idx+hi).clamp(0,L)
41 lc = (lb-la).clamp(min=1).to(v.dtype)[None,:,None]
42 rc = (rb-ra).clamp(min=1).to(v.dtype)[None,:,None]
43 left = (pref[:,lb]-pref[:,la])/lc
44 right = (pref[:,rb]-pref[:,ra])/rc
45 # k'=k+1 gives the stated 1/[k'(k'+1) log 2] mass.
46 a = 1.0/((k+1)*(k+2)*math.log(2.0))
47 out = out + a*(left+right)
48 z = x + self.alpha * self.out(out)
49 z = z + self.beta * self.ff(self.ffnorm(z))
50 return z
51
52class TailTransformer(nn.Module):
53 def __init__(self, win=32, d=64, depth=2, idea=True):
54 super().__init__(); self.idea = idea
55 self.inp = nn.Linear(1,d)
56 self.pos = nn.Parameter(torch.zeros(1,win,d)); nn.init.normal_(self.pos,std=.02)
57 if idea:
58 self.layers = nn.ModuleList([DyadicMixer(d) for _ in range(depth)])
59 else:
60 layer = nn.TransformerEncoderLayer(d, nhead=2, dim_feedforward=128,
61 batch_first=True, dropout=0.0, activation='relu')
62 self.enc = nn.TransformerEncoder(layer, depth)
63 self.head = nn.Linear(win*d,1)
64 def forward(self,x):
65 h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]
66 if self.idea:
67 for layer in self.layers: h=layer(h)
68 else: h=self.enc(h)
69 return self.head(h.reshape(x.shape[0],-1))
70
71def seed_all(s):
72 random.seed(s); np.random.seed(s); torch.manual_seed(s)
73 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
74
75def run_one(kind, cfg, seed, capture=False):
76 seed_all(seed)
77 ds=get_dataset('sequence', seed, n_train=400, n_test=400)
78 net=TailTransformer(win=ds['input_shape'][0], idea=(kind=='idea'))
79 net, metric, hist=train_model(net, ds, epochs=cfg['epochs'], lr=cfg['lr'], batch=128, log=lambda *_: None)
80 if capture:
81 # Trained-model signature: compare observed influence of distant versus near input
82 # coordinates by finite perturbations on the held-out task model.
83 dev=next(net.parameters()).device; x=ds['xte'][:128].to(dev)
84 with torch.no_grad():
85 base=net(x).squeeze(-1)
86 near=x.clone(); near[:,1]+=0.1
87 far=x.clone(); far[:,-1]+=0.1
88 dn=(net(near).squeeze(-1)-base).abs().mean().item()
89 df=(net(far).squeeze(-1)-base).abs().mean().item()
90 return float(metric), {'near_influence':dn,'far_influence':df,'far_near_ratio':df/(dn+1e-12)}
91 return float(metric)
92
93def main():
94 # Math sanity check is included before training: tail and dyadic-mass ratios.
95 rs=np.arange(1,2000000,dtype=np.float64); kval=1/((rs+2)*np.log(rs+2)**2)
96 math_rows=[]
97 for R in [8,32,128,512]:
98 tail=float(kval[R:].sum()+1/np.log(2000002)); math_rows.append({'R':R,'ratio':tail/(1/np.log(R+2))})
99 def base_factory(cfg): return lambda s: run_one('baseline',cfg,s)
100 baseline=sweep_baseline(base_factory, GRID)
101 # Idea uses exactly the same 3-point union and is evaluated on all paired seeds.
102 idea_trials=[]
103 for cfg in GRID:
104 r=evaluate(lambda s,cfg=cfg: run_one('idea',cfg,s), seeds=tuple(range(4)))
105 idea_trials.append({'cfg':cfg,'mean':r['mean']})
106 best=min(idea_trials,key=lambda z:z['mean'])['cfg']
107 idea=evaluate(lambda s: run_one('idea',best,s), seeds=SEEDS)
108 # Signature from two independently trained systems at their selected settings.
109 _, bsig=run_one('baseline',baseline['best_cfg'],0,True)
110 _, isig=run_one('idea',best,0,True)
111 signature={'prediction':'dyadic tail should preserve more distant influence than standard local/attention-free mixing',
112 'baseline_observed':bsig,'idea_observed':isig,
113 'predicted_far_near_ratio_idea_gt_baseline':True,
114 'confirmed': bool(isig['far_near_ratio'] > bsig['far_near_ratio'])}
115 report=make_report('sequence','transformer_tiny',baseline,idea,{'mechanism_signature':signature,
116 'math_sanity':math_rows,'idea_sweep':idea_trials,
117 'track_justification':'Sequence forecast has multi-token correlations and is the mandated structural match for attention/SSM ideas.'})
118 with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
119 print(json.dumps(report,indent=2))
120if __name__=='__main__': main()