import os, sys, json, math, time from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, sweep_baseline, make_report SEEDS = tuple(range(8)) SWEEP_SEEDS = (0, 1, 2, 3) EPOCHS = 12 BATCH = 128 LR_GRID = [1e-3, 3e-3, 6e-3] def kernel_attention_math_check(seed=1180): rng = np.random.default_rng(seed); rows = [] for k in (1, 2, 3): x = rng.random((1800, k)); q = x[:40] d = np.sqrt(((q[:, None, :] - x[None, :, :]) ** 2).sum(-1)) radii = np.geomspace(.04, .16, 7); mass, wr = [], [] for r in radii: w = np.exp(-(d / r) ** 2); mass.append(w.mean()) a = w / w.sum(1, keepdims=True); wr.append(np.sqrt((a*d*d).sum(1).mean())) rows.append({'dimension': k, 'predicted_mass_slope': k, 'observed_mass_slope': float(np.polyfit(np.log(radii), np.log(mass), 1)[0]), 'predicted_radius_slope': 1, 'observed_radius_slope': float(np.polyfit(np.log(radii), np.log(wr), 1)[0])}) return rows class LocalKernelSelfAttention(nn.Module): def __init__(self, d, nhead=2, neighbors=16): super().__init__(); assert d % nhead == 0 self.d, self.nhead, self.hd, self.neighbors = d, nhead, d//nhead, neighbors self.qkv, self.out = nn.Linear(d, 3*d), nn.Linear(d, d) def forward(self, x): b, t, d = x.shape; z = torch.nn.functional.layer_norm(x, (d,)) dist = torch.cdist(z, z).clamp_min(1e-7) eye = torch.eye(t, device=x.device, dtype=x.dtype) masked = dist + eye[None] * 1e6; m = min(self.neighbors, max(1, t-1)) knn = masked.topk(m, dim=-1, largest=False).values; r = knn[..., -1:].detach() khat = (1.0/(torch.log((r+1e-7)/(knn[..., :-1]+1e-7)).mean(-1, keepdim=True)+1e-7)).clamp(1., float(d)) g = torch.exp(-((dist/r)**2))/(r**khat); g = g.masked_fill(eye.bool()[None], 0.) a = g/(g.sum(-1, keepdim=True)+1e-7) v = self.qkv(x)[..., 2*d:].reshape(b,t,self.nhead,self.hd).transpose(1,2) y = torch.matmul(a[:,None],v).transpose(1,2).reshape(b,t,d) return self.out(y), a class KernelBlock(nn.Module): def __init__(self, d=64, nhead=2, neighbors=16): super().__init__(); self.attn=LocalKernelSelfAttention(d,nhead,neighbors) self.norm1=nn.LayerNorm(d); self.ff=nn.Sequential(nn.Linear(d,128),nn.ReLU(),nn.Linear(128,d)); self.norm2=nn.LayerNorm(d) def forward(self,x): y,a=self.attn(x); x=self.norm1(x+y); return self.norm2(x+self.ff(x)),a class KernelTransformer(nn.Module): def __init__(self, win, out_dim=1, d=64, depth=2, neighbors=16): 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([KernelBlock(d,2,neighbors) for _ in range(depth)]); self.head=nn.Linear(win*d,out_dim); self.last_attention=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_attention=aa[-1].detach(); return self.head(h.reshape(h.shape[0],-1)) def make_ds(seed): return get_dataset('sequence',seed,n_train=400,n_test=200) def trained_signature(net, x): device=next(net.parameters()).device; x=x[:64].to(device) with torch.no_grad(): h=net.inp(x.unsqueeze(-1))+net.pos[:,:x.shape[1]]; h=torch.nn.functional.layer_norm(h,(h.shape[-1],)) d=torch.cdist(h,h); eye=torch.eye(d.shape[-1],device=device,dtype=d.dtype); md=d+eye[None]*1e6 m=min(net.blocks[0].attn.neighbors,d.shape[-1]-1); knn=md.topk(m,dim=-1,largest=False).values; r=knn[...,-1:] kh=(1/(torch.log((r+1e-7)/(knn[...,:-1]+1e-7)).mean(-1,keepdim=True)+1e-7)).clamp(1.,float(h.shape[-1])) gam=np.geomspace(.7,2.,6); eff=[] for g in gam: w=torch.exp(-((d/(r*g))**2)); w=w.masked_fill(eye.bool()[None],0); w=w/(w.sum(-1,keepdim=True)+1e-7); eff.append(float((1/(w.square().sum(-1)+1e-9)).mean())) slope=float(np.polyfit(np.log(gam),np.log(eff),1)[0]); pred=float(kh.mean()) ent=float(-(net.last_attention.clamp_min(1e-9)*net.last_attention.clamp_min(1e-9).log()).sum(-1).mean()) return {'predicted_dimension':pred,'predicted_support_slope':pred,'observed_support_slope':slope,'kernel_entropy_observed':ent,'effective_neighbors_observed':eff[2],'confirmed':bool(abs(slope-pred)