Gain-Rigid Sparse Attention / bench_gain_rigid.py
Beats tuned baseline
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()