Gauge-Free Inverse OT Attention / bench_experiment.py
Failed on benchmark
1import json, sys, 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, sweep_baseline, evaluate, make_report
8
9SEEDS = tuple(range(8))
10SWEEP_SEEDS = (0, 1, 2, 3)
11LRS = [1e-3, 3e-3, 6e-3]
12EPOCHS = 12
13BATCH = 128
14
15class AttentionBlock(nn.Module):
16 def __init__(self, d=64, sinkhorn=False, eps=0.7, iters=12):
17 super().__init__()
18 self.d, self.sinkhorn, self.eps, self.iters = d, sinkhorn, eps, iters
19 self.q, self.k, self.v = nn.Linear(d,d), nn.Linear(d,d), nn.Linear(d,d)
20 self.o, self.norm = nn.Linear(d,d), nn.LayerNorm(d)
21 self.last_signature = None
22
23 def forward(self, x):
24 q, k, v = self.q(x), self.k(x), self.v(x)
25 scores = q @ k.transpose(-1, -2) / math.sqrt(self.d)
26 if not self.sinkhorn:
27 w = torch.softmax(scores, dim=-1)
28 else:
29 # Positive kernel exp(scores/epsilon), uniform marginals, log-domain scaling.
30 logk = scores / self.eps
31 lm = -math.log(x.shape[1])
32 lu = torch.zeros_like(logk[..., 0])
33 lv = torch.zeros_like(logk[..., 0, :])
34 for _ in range(self.iters):
35 lu = lm - torch.logsumexp(logk + lv.unsqueeze(-2), dim=-1)
36 lv = lm - torch.logsumexp(logk + lu.unsqueeze(-1), dim=-2)
37 w = torch.exp(lu.unsqueeze(-1) + logk + lv.unsqueeze(-2))
38 with torch.no_grad():
39 rowerr = (w.sum(-1) - 1.0/x.shape[1]).abs().mean().item()
40 colerr = (w.sum(-2) - 1.0/x.shape[1]).abs().mean().item()
41 self.last_signature = {
42 "row_marginal_error": rowerr,
43 "column_marginal_error": colerr,
44 "min_attention": float(w.min().item()),
45 "mean_row_sum": float(w.sum(-1).mean().item())
46 }
47 return self.norm(x + self.o(w @ v))
48
49class TinyTransformer(nn.Module):
50 def __init__(self, sinkhorn=False, eps=0.7, iters=12):
51 super().__init__()
52 self.embed = nn.Linear(1, 64)
53 self.pos = nn.Parameter(torch.zeros(1, 32, 64))
54 self.attn1 = AttentionBlock(64, sinkhorn, eps, iters)
55 self.ff1 = nn.Sequential(nn.Linear(64,128), nn.GELU(), nn.Linear(128,64), nn.LayerNorm(64))
56 self.attn2 = AttentionBlock(64, sinkhorn, eps, iters)
57 self.ff2 = nn.Sequential(nn.Linear(64,128), nn.GELU(), nn.Linear(128,64), nn.LayerNorm(64))
58 self.head = nn.Linear(64, 1)
59
60 def forward(self, x):
61 x = self.embed(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]
62 x = self.ff1(self.attn1(x) + x)
63 x = self.ff2(self.attn2(x) + x)
64 return self.head(x[:, -1])
65
66
67def seed_all(seed):
68 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
69 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
70
71
72def train_one(seed, lr, sinkhorn, eps=0.7, iters=12, collect=False):
73 seed_all(seed)
74 ds = get_dataset("sequence", seed, n_train=800, n_test=240)
75 dev = "cuda" if torch.cuda.is_available() else "cpu"
76 try:
77 model = TinyTransformer(sinkhorn, eps, iters).to(dev)
78 xtr,ytr,xte,yte = [ds[k].to(dev) for k in ("xtr","ytr","xte","yte")]
79 opt = torch.optim.Adam(model.parameters(), lr=lr)
80 lossf = nn.MSELoss()
81 for _ in range(EPOCHS):
82 model.train(); perm=torch.randperm(len(xtr), device=dev)
83 for i in range(0,len(xtr),BATCH):
84 z=perm[i:i+BATCH]; loss=lossf(model(xtr[z]),ytr[z])
85 opt.zero_grad(); loss.backward(); opt.step()
86 model.eval()
87 with torch.no_grad(): metric=float(lossf(model(xte),yte).item())
88 sig = model.attn2.last_signature if collect else None
89 return metric, sig
90 except RuntimeError:
91 if dev == "cpu": raise
92 torch.cuda.empty_cache(); seed_all(seed)
93 model = TinyTransformer(sinkhorn, eps, iters)
94 xtr,ytr,xte,yte = [ds[k] for k in ("xtr","ytr","xte","yte")]
95 opt=torch.optim.Adam(model.parameters(),lr=lr); lossf=nn.MSELoss()
96 for _ in range(EPOCHS):
97 perm=torch.randperm(len(xtr))
98 for i in range(0,len(xtr),BATCH):
99 z=perm[i:i+BATCH]; loss=lossf(model(xtr[z]),ytr[z])
100 opt.zero_grad(); loss.backward(); opt.step()
101 with torch.no_grad(): metric=float(lossf(model(xte),yte).item())
102 return metric, model.attn2.last_signature if collect else None
103
104
105def train_metric(sinkhorn, cfg, seed):
106 return train_one(seed, cfg["lr"], sinkhorn, cfg.get("eps",0.7), cfg.get("iters",12))[0]
107
108
109def main():
110 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]]
111 base=sweep_baseline(lambda cfg: lambda seed: train_metric(False,cfg,seed), baseline_grid, seeds=SWEEP_SEEDS)
112 best_lr=base["best_cfg"]["lr"]
113 idea_grid=[{"lr":best_lr,"eps":e,"iters":12} for e in [0.5,0.7,1.0]]
114 # Equal union of learning rates: baseline already evaluated every idea-side lr.
115 idea_cfg=min(idea_grid, key=lambda c: evaluate(lambda s: train_metric(True,c,s), SWEEP_SEEDS)["mean"])
116 idea=evaluate(lambda s: train_metric(True,idea_cfg,s), seeds=SEEDS)
117 sigs=[]
118 for s in SEEDS:
119 _,sig=train_one(s,idea_cfg["lr"],True,idea_cfg["eps"],idea_cfg["iters"],True); sigs.append(sig)
120 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)}
121 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."})
122 with open("bench_report.json","w") as f: json.dump(report,f,indent=2)
123 print(json.dumps(report,indent=2))
124
125if __name__ == "__main__": main()