import json, sys, math, random import numpy as np import torch import torch.nn as nn sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import get_dataset, sweep_baseline, evaluate, make_report SEEDS = tuple(range(8)) SWEEP_SEEDS = (0, 1, 2, 3) LRS = [1e-3, 3e-3, 6e-3] EPOCHS = 12 BATCH = 128 class AttentionBlock(nn.Module): def __init__(self, d=64, sinkhorn=False, eps=0.7, iters=12): super().__init__() self.d, self.sinkhorn, self.eps, self.iters = d, sinkhorn, eps, iters self.q, self.k, self.v = nn.Linear(d,d), nn.Linear(d,d), nn.Linear(d,d) self.o, self.norm = nn.Linear(d,d), nn.LayerNorm(d) self.last_signature = None def forward(self, x): q, k, v = self.q(x), self.k(x), self.v(x) scores = q @ k.transpose(-1, -2) / math.sqrt(self.d) if not self.sinkhorn: w = torch.softmax(scores, dim=-1) else: # Positive kernel exp(scores/epsilon), uniform marginals, log-domain scaling. logk = scores / self.eps lm = -math.log(x.shape[1]) lu = torch.zeros_like(logk[..., 0]) lv = torch.zeros_like(logk[..., 0, :]) for _ in range(self.iters): lu = lm - torch.logsumexp(logk + lv.unsqueeze(-2), dim=-1) lv = lm - torch.logsumexp(logk + lu.unsqueeze(-1), dim=-2) w = torch.exp(lu.unsqueeze(-1) + logk + lv.unsqueeze(-2)) with torch.no_grad(): rowerr = (w.sum(-1) - 1.0/x.shape[1]).abs().mean().item() colerr = (w.sum(-2) - 1.0/x.shape[1]).abs().mean().item() self.last_signature = { "row_marginal_error": rowerr, "column_marginal_error": colerr, "min_attention": float(w.min().item()), "mean_row_sum": float(w.sum(-1).mean().item()) } return self.norm(x + self.o(w @ v)) class TinyTransformer(nn.Module): def __init__(self, sinkhorn=False, eps=0.7, iters=12): super().__init__() self.embed = nn.Linear(1, 64) self.pos = nn.Parameter(torch.zeros(1, 32, 64)) self.attn1 = AttentionBlock(64, sinkhorn, eps, iters) self.ff1 = nn.Sequential(nn.Linear(64,128), nn.GELU(), nn.Linear(128,64), nn.LayerNorm(64)) self.attn2 = AttentionBlock(64, sinkhorn, eps, iters) self.ff2 = nn.Sequential(nn.Linear(64,128), nn.GELU(), nn.Linear(128,64), nn.LayerNorm(64)) self.head = nn.Linear(64, 1) def forward(self, x): x = self.embed(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]] x = self.ff1(self.attn1(x) + x) x = self.ff2(self.attn2(x) + x) return self.head(x[:, -1]) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def train_one(seed, lr, sinkhorn, eps=0.7, iters=12, collect=False): seed_all(seed) ds = get_dataset("sequence", seed, n_train=800, n_test=240) dev = "cuda" if torch.cuda.is_available() else "cpu" try: model = TinyTransformer(sinkhorn, eps, iters).to(dev) xtr,ytr,xte,yte = [ds[k].to(dev) for k in ("xtr","ytr","xte","yte")] opt = torch.optim.Adam(model.parameters(), lr=lr) lossf = nn.MSELoss() for _ in range(EPOCHS): model.train(); perm=torch.randperm(len(xtr), device=dev) for i in range(0,len(xtr),BATCH): z=perm[i:i+BATCH]; loss=lossf(model(xtr[z]),ytr[z]) opt.zero_grad(); loss.backward(); opt.step() model.eval() with torch.no_grad(): metric=float(lossf(model(xte),yte).item()) sig = model.attn2.last_signature if collect else None return metric, sig except RuntimeError: if dev == "cpu": raise torch.cuda.empty_cache(); seed_all(seed) model = TinyTransformer(sinkhorn, eps, iters) xtr,ytr,xte,yte = [ds[k] for k in ("xtr","ytr","xte","yte")] opt=torch.optim.Adam(model.parameters(),lr=lr); lossf=nn.MSELoss() for _ in range(EPOCHS): perm=torch.randperm(len(xtr)) for i in range(0,len(xtr),BATCH): z=perm[i:i+BATCH]; loss=lossf(model(xtr[z]),ytr[z]) opt.zero_grad(); loss.backward(); opt.step() with torch.no_grad(): metric=float(lossf(model(xte),yte).item()) return metric, model.attn2.last_signature if collect else None def train_metric(sinkhorn, cfg, seed): return train_one(seed, cfg["lr"], sinkhorn, cfg.get("eps",0.7), cfg.get("iters",12))[0] def main(): baseline_grid=[{"lr":lr,"eps":eps,"iters":iters} for lr in LRS for eps in ([0.5,0.7,1.0] if False else [0.7]) for iters in [12]] base=sweep_baseline(lambda cfg: lambda seed: train_metric(False,cfg,seed), baseline_grid, seeds=SWEEP_SEEDS) best_lr=base["best_cfg"]["lr"] idea_grid=[{"lr":best_lr,"eps":e,"iters":12} for e in [0.5,0.7,1.0]] # Equal union of learning rates: baseline already evaluated every idea-side lr. idea_cfg=min(idea_grid, key=lambda c: evaluate(lambda s: train_metric(True,c,s), SWEEP_SEEDS)["mean"]) idea=evaluate(lambda s: train_metric(True,idea_cfg,s), seeds=SEEDS) sigs=[] for s in SEEDS: _,sig=train_one(s,idea_cfg["lr"],True,idea_cfg["eps"],idea_cfg["iters"],True); sigs.append(sig) signature={"predicted_max_marginal_error":1e-4,"observed_mean_row_marginal_error":float(np.mean([x["row_marginal_error"] for x in sigs])),"observed_mean_column_marginal_error":float(np.mean([x["column_marginal_error"] for x in sigs])),"observed_min_attention_mean":float(np.mean([x["min_attention"] for x in sigs])),"confirmed":bool(max(x["row_marginal_error"] for x in sigs)<1e-4 and max(x["column_marginal_error"] for x in sigs)<1e-4)} report=make_report("sequence","transformer_tiny",base,idea,{"mechanism_signature":signature,"idea_sweep":idea_grid,"track_justification":"Sequence forecast contains multi-token correlations and directly exercises attention."}) with open("bench_report.json","w") as f: json.dump(report,f,indent=2) print(json.dumps(report,indent=2)) if __name__ == "__main__": main()