Gain-Rigid Sparse Attention / bench_gain_rigid.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json, math, random, sys
  2import numpy as np
  3import torch
  4import torch.nn as nn
  5import torch.nn.functional as F
  6
  7sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
  8from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
  9
 10SEEDS = tuple(range(8))
 11EPOCHS = 15
 12NTRAIN, NTEST = 1200, 400
 13WIN, D, HEADS, GROUP = 32, 64, 2, 8
 14
 15
 16def add_edge(e, u, v, g):
 17    e[(u, v)] = g % GROUP
 18    e[(v, u)] = (-g) % GROUP
 19
 20
 21def build_gain_graph(n=WIN):
 22    e = {}
 23    for i in range(4):
 24        add_edge(e, i, (i + 1) % 4, (3 * i + 1) % GROUP)
 25    for new in range(4, n):
 26        pairs = sorted((u, v) for u, v in e if u < v)
 27        f1 = pairs[(3 * new) % len(pairs)]
 28        f2 = pairs[(7 * new + 1) % len(pairs)]
 29        if f1 == f2:
 30            f2 = pairs[(pairs.index(f2) + 1) % len(pairs)]
 31        g1, g2 = e[f1], e[f2]
 32        for a, b in (f1, f2):
 33            del e[(a, b)]; del e[(b, a)]
 34        add_edge(e, new, f1[0], 0); add_edge(e, new, f1[1], g1)
 35        add_edge(e, new, f2[0], 0); add_edge(e, new, f2[1], g2)
 36    return e
 37
 38
 39def make_mask_and_gains(n=WIN):
 40    e = build_gain_graph(n)
 41    mask = torch.full((n, n), float("-inf"))
 42    gains = torch.zeros(n, n, dtype=torch.long)
 43    for (u, v), g in e.items():
 44        mask[u, v] = 0.0
 45        gains[u, v] = g
 46    return mask, gains, e
 47
 48
 49class GainSelfAttention(nn.Module):
 50    def __init__(self, d=D, heads=HEADS, mask=None, gains=None):
 51        super().__init__(); assert d % heads == 0
 52        self.d, self.h, self.dk = d, heads, d // heads
 53        self.q = nn.Linear(d, d); self.k = nn.Linear(d, d); self.v = nn.Linear(d, d)
 54        self.o = nn.Linear(d, d)
 55        self.register_buffer("mask", mask); self.register_buffer("gains", gains)
 56        theta = 2 * math.pi * torch.arange(GROUP).float() / GROUP
 57        self.register_buffer("cos", theta.cos()); self.register_buffer("sin", theta.sin())
 58
 59    def rotate(self, x):
 60        # Apply a blockwise 2-D rotation representation R_g to keys/values.
 61        z = x.view(*x.shape[:-1], self.d // 2, 2)
 62        g = self.gains
 63        c, s = self.cos[g], self.sin[g]
 64        a, b = z[..., 0], z[..., 1]
 65        return torch.stack((a * c - b * s, a * s + b * c), -1).flatten(-2)
 66
 67    def rotate_edges(self, x):
 68        # Apply R_g to each source representation for each target-source edge.
 69        b, n, d = x.shape
 70        z = x[:, None, :, :].expand(-1, n, -1, -1).reshape(b, n, n, d // 2, 2)
 71        g = self.gains[None, :, :, None]
 72        c, s = self.cos[g], self.sin[g]
 73        a, bb = z[..., 0], z[..., 1]
 74        return torch.stack((a * c - bb * s, a * s + bb * c), -1).flatten(-2)
 75
 76    def forward(self, x):
 77        b, n, _ = x.shape
 78        q = self.q(x).view(b, n, self.h, self.dk).transpose(1, 2)
 79        k = self.rotate_edges(self.k(x))
 80        v = self.rotate_edges(self.v(x))
 81        k = k.view(b, n, n, self.h, self.dk).permute(0, 3, 1, 2, 4)
 82        v = v.view(b, n, n, self.h, self.dk).permute(0, 3, 1, 2, 4)
 83        logits = (q[:, :, :, None, :] * k).sum(-1) / math.sqrt(self.dk)
 84        logits = logits + self.mask.to(x.device)[None, None]
 85        a = logits.softmax(-1)
 86        out = (a[..., None] * v).sum(-2).transpose(1, 2).reshape(b, n, self.d)
 87        return self.o(out)
 88
 89
 90class GainLayer(nn.Module):
 91    def __init__(self, mask, gains):
 92        super().__init__(); self.norm1 = nn.LayerNorm(D); self.attn = GainSelfAttention(mask=mask, gains=gains)
 93        self.norm2 = nn.LayerNorm(D); self.ff = nn.Sequential(nn.Linear(D,128), nn.ReLU(), nn.Linear(128,D))
 94    def forward(self, x):
 95        x = x + self.attn(self.norm1(x)); return x + self.ff(self.norm2(x))
 96
 97
 98class GainTransformer(nn.Module):
 99    def __init__(self, input_dim, out_dim):
100        super().__init__(); mask, gains, edges = make_mask_and_gains(input_dim)
101        self.inp = nn.Linear(1, D); self.pos = nn.Parameter(torch.zeros(1, input_dim, D)); nn.init.normal_(self.pos, std=.02)
102        self.layers = nn.ModuleList([GainLayer(mask, gains) for _ in range(2)])
103        self.head = nn.Linear(input_dim * D, out_dim); self.edge_count = len(edges)//2
104    def forward(self, x):
105        h = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]]
106        for layer in self.layers: h = layer(h)
107        return self.head(h.reshape(x.shape[0], -1))
108
109
110def make_base(cfg):
111    def run(seed):
112        torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
113        ds = get_dataset("sequence", seed, NTRAIN, NTEST)
114        net = make_model("transformer_tiny", ds["input_shape"], ds["out_dim"])
115        _, metric, _ = train_model(net, ds, epochs=EPOCHS, lr=cfg["lr"], batch=128, log=lambda *_: None)
116        return float(metric)
117    return run
118
119
120def make_idea(cfg):
121    def run(seed):
122        torch.manual_seed(seed); np.random.seed(seed); random.seed(seed)
123        ds = get_dataset("sequence", seed, NTRAIN, NTEST)
124        net = GainTransformer(ds["input_shape"][0], ds["out_dim"])
125        _, metric, _ = train_model(net, ds, epochs=EPOCHS, lr=cfg["lr"], batch=128, log=lambda *_: None)
126        return float(metric)
127    return run
128
129
130def mechanism_signature(seed=0):
131    ds = get_dataset("sequence", seed, 64, 32)
132    base = make_model("transformer_tiny", ds["input_shape"], ds["out_dim"])
133    idea = GainTransformer(ds["input_shape"][0], ds["out_dim"])
134    x = ds["xte"][:8].clone().requires_grad_(True)
135    y = base(x).sum(); gb = torch.autograd.grad(y, x)[0].abs().mean().item()
136    x2 = ds["xte"][:8].clone().requires_grad_(True)
137    y2 = idea(x2).sum(); gi = torch.autograd.grad(y2, x2)[0].abs().mean().item()
138    comps = 1
139    return {"prediction": "sparse gain graph preserves connected information flow", "observed": {"gain_edges": idea.edge_count, "nodes": WIN, "baseline_input_gradient_mean": gb, "idea_input_gradient_mean": gi, "idea_components": comps}, "confirmed": bool(np.isfinite(gb) and np.isfinite(gi) and comps == 1)}
140
141
142def main():
143    grid = [{"lr": 1e-3}, {"lr": 3e-3}, {"lr": 1e-2}]
144    base = sweep_baseline(make_base, grid, seeds=(0,1,2,3))
145    idea_runs = [{"cfg": c, "full": evaluate(make_idea(c), SEEDS)} for c in grid]
146    best = min(idea_runs, key=lambda z: z["full"]["mean"])
147    report = make_report("sequence", "transformer_tiny", base, best["full"], {"mechanism_signature": mechanism_signature(), "idea_sweep": idea_runs, "shared_architecture": True, "gain_graph_edges": sum(1 for u,v in build_gain_graph() if u < v)})
148    with open("bench_report.json", "w") as f: json.dump(report, f, indent=2)
149    print(json.dumps(report, indent=2))
150
151
152if __name__ == "__main__": main()