Gauge-Free Inverse OT Attention / bench_experiment.py

Failed on benchmark

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