Critical-Tail Multiscale Mixer / stage2_bench.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()