import sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import get_dataset, train_model, sweep_baseline, make_report, evaluate SEEDS = tuple(range(8)); SWEEP_SEEDS = tuple(range(4)); WIN = 32 EPOCHS = 10; BATCH = 128; WEIGHT_DECAY = 0.0 LRS = [0.0015, 0.003, 0.006]; M = 8; TAU_MIN = 1e-3 CAPTURE = {} class BiasAttention(nn.Module): def __init__(self, d, heads, win, kind, gamma=1.0): super().__init__(); assert d % heads == 0 self.d, self.heads, self.dk, self.win, self.kind = d, heads, d//heads, win, kind self.gamma = float(gamma); self.qkv = nn.Linear(d, 3*d); self.out = nn.Linear(d, d) if kind == "table": self.bias = nn.Parameter(torch.zeros(heads, win)) else: self.alpha = nn.Parameter(torch.zeros(heads, M)) init = torch.logspace(math.log10(.03), math.log10(1.0), M) self.beta = nn.Parameter(torch.log(torch.expm1(init-TAU_MIN)).repeat(heads, 1)) def lag_bias(self, device): d = torch.arange(self.win, device=device, dtype=torch.float32) if self.kind == "table": return self.bias w = F.softmax(self.alpha, dim=-1); tau = F.softplus(self.beta) + TAU_MIN k = torch.exp(-d[None,:,None] * tau[:,None,:]).matmul(w[...,None]).squeeze(-1) return torch.log(k + 1e-6) * self.gamma def forward(self, x): B,L,_ = x.shape; q,k,v = self.qkv(x).chunk(3, dim=-1) def split(z): return z.view(B,L,self.heads,self.dk).transpose(1,2) q,k,v = map(split,(q,k,v)); logits = (q @ k.transpose(-2,-1))/math.sqrt(self.dk) lb = self.lag_bias(x.device) lags = (torch.arange(L,device=x.device)[None,:]-torch.arange(L,device=x.device)[:,None]).clamp(min=0) logits = logits + lb[:,lags][None] logits = logits.masked_fill(torch.triu(torch.ones(L,L,device=x.device,dtype=torch.bool),1),-1e4) a = torch.softmax(logits,dim=-1); y=(a@v).transpose(1,2).contiguous().view(B,L,self.d) return self.out(y),a class Block(nn.Module): def __init__(self,d,heads,win,kind,gamma): super().__init__(); self.n1=nn.LayerNorm(d); self.attn=BiasAttention(d,heads,win,kind,gamma) self.n2=nn.LayerNorm(d); self.ff=nn.Sequential(nn.Linear(d,128),nn.GELU(),nn.Linear(128,d)) def forward(self,x): z,a=self.attn(self.n1(x)); x=x+z; return x+self.ff(self.n2(x)),a class LagTransformer(nn.Module): def __init__(self,kind,gamma=1.0,depth=2,d=64,heads=2,win=WIN): super().__init__(); self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.zeros(1,win,d)); nn.init.normal_(self.pos,std=.02) self.blocks=nn.ModuleList([Block(d,heads,win,kind,gamma) for _ in range(depth)]); self.head=nn.Linear(win*d,1); self.last_attn=None def forward(self,x): h=self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]; aa=[] for block in self.blocks: h,a=block(h); aa.append(a) self.last_attn=aa[-1].detach(); return self.head(h.reshape(x.shape[0],-1)) def make_train(kind,cfg,capture=False): def run(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) ds=get_dataset("sequence",seed,n_train=800,n_test=300); model=LagTransformer(kind,gamma=cfg["gamma"]) net,metric,_=train_model(model,ds,epochs=EPOCHS,lr=cfg["lr"],batch=BATCH,weight_decay=WEIGHT_DECAY,log=lambda *_:None) if net is None: return float("inf") if capture: with torch.no_grad(): dev=next(net.parameters()).device; b=net.blocks[-1].attn.lag_bias(dev).cpu().numpy() net.eval(); net(ds["xte"][:300].to(dev)); a=net.last_attn.cpu().numpy(); L=a.shape[-1] av=np.zeros(L); ct=np.zeros(L) for q in range(L): for k in range(q+1): av[q-k]+=a[:,:,q,k].mean(); ct[q-k]+=1 CAPTURE[int(seed)]={"bias":b.tolist(),"attention":(av/np.maximum(ct,1)).tolist()} return float(metric) return run def signature(): first=[]; second=[]; att=[] for rec in CAPTURE.values(): b=np.asarray(rec["bias"]); first.extend(np.diff(b,axis=1).ravel()); second.extend(np.diff(b,n=2,axis=1).ravel()); att.append(np.diff(rec["attention"]).max()) f=float(np.max(first)); s=float(np.min(second)) return {"prediction":"trained mixture log-kernel bias decreases with lag; retained content logits can make total attention nonmonotone","observed_bias_max_first_difference":f,"observed_bias_min_second_difference":s,"observed_attention_max_first_difference":float(max(att)),"n_models":len(CAPTURE),"confirmed":bool(f<=1e-7)} def main(): grid=[{"lr":lr,"gamma":1.0} for lr in LRS] base=sweep_baseline(lambda cfg:make_train("table",cfg),grid) # Same union of learning rates; select idea on the same four sweep seeds. tried=[] for cfg in grid: r=evaluate(make_train("mixture",cfg),seeds=SWEEP_SEEDS); tried.append({"cfg":cfg,"mean":r["mean"]}) best_cfg=min(grid,key=lambda c: next(x["mean"] for x in tried if x["cfg"]==c)); CAPTURE.clear() idea=evaluate(make_train("mixture",best_cfg,capture=True),seeds=SEEDS) report=make_report("sequence","transformer_tiny",base,idea,{"mechanism_signature":signature(),"protocol_notes":"Matched sequence track; shared 2-block d=64 causal transformer, differing only in relative lag bias: unconstrained table versus 8-positive-exponential mixture."}) report["idea_sweep"]=tried; report["idea_config"]=best_cfg; report["search_space"]={"baseline_grid":grid,"idea_grid":grid} Path("bench_report.json").write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2)) if __name__ == "__main__": main()