import json, math, random, sys import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, "/home/maxwelhelp/all/math2nn") from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) EPOCHS = 15 NTRAIN, NTEST = 1200, 400 WIN, D, HEADS, GROUP = 32, 64, 2, 8 def add_edge(e, u, v, g): e[(u, v)] = g % GROUP e[(v, u)] = (-g) % GROUP def build_gain_graph(n=WIN): e = {} for i in range(4): add_edge(e, i, (i + 1) % 4, (3 * i + 1) % GROUP) for new in range(4, n): pairs = sorted((u, v) for u, v in e if u < v) f1 = pairs[(3 * new) % len(pairs)] f2 = pairs[(7 * new + 1) % len(pairs)] if f1 == f2: f2 = pairs[(pairs.index(f2) + 1) % len(pairs)] g1, g2 = e[f1], e[f2] for a, b in (f1, f2): del e[(a, b)]; del e[(b, a)] add_edge(e, new, f1[0], 0); add_edge(e, new, f1[1], g1) add_edge(e, new, f2[0], 0); add_edge(e, new, f2[1], g2) return e def make_mask_and_gains(n=WIN): e = build_gain_graph(n) mask = torch.full((n, n), float("-inf")) gains = torch.zeros(n, n, dtype=torch.long) for (u, v), g in e.items(): mask[u, v] = 0.0 gains[u, v] = g return mask, gains, e class GainSelfAttention(nn.Module): def __init__(self, d=D, heads=HEADS, mask=None, gains=None): super().__init__(); assert d % heads == 0 self.d, self.h, self.dk = d, heads, d // heads self.q = nn.Linear(d, d); self.k = nn.Linear(d, d); self.v = nn.Linear(d, d) self.o = nn.Linear(d, d) self.register_buffer("mask", mask); self.register_buffer("gains", gains) theta = 2 * math.pi * torch.arange(GROUP).float() / GROUP self.register_buffer("cos", theta.cos()); self.register_buffer("sin", theta.sin()) def rotate(self, x): # Apply a blockwise 2-D rotation representation R_g to keys/values. z = x.view(*x.shape[:-1], self.d // 2, 2) g = self.gains c, s = self.cos[g], self.sin[g] a, b = z[..., 0], z[..., 1] return torch.stack((a * c - b * s, a * s + b * c), -1).flatten(-2) def rotate_edges(self, x): # Apply R_g to each source representation for each target-source edge. b, n, d = x.shape z = x[:, None, :, :].expand(-1, n, -1, -1).reshape(b, n, n, d // 2, 2) g = self.gains[None, :, :, None] c, s = self.cos[g], self.sin[g] a, bb = z[..., 0], z[..., 1] return torch.stack((a * c - bb * s, a * s + bb * c), -1).flatten(-2) def forward(self, x): b, n, _ = x.shape q = self.q(x).view(b, n, self.h, self.dk).transpose(1, 2) k = self.rotate_edges(self.k(x)) v = self.rotate_edges(self.v(x)) k = k.view(b, n, n, self.h, self.dk).permute(0, 3, 1, 2, 4) v = v.view(b, n, n, self.h, self.dk).permute(0, 3, 1, 2, 4) logits = (q[:, :, :, None, :] * k).sum(-1) / math.sqrt(self.dk) logits = logits + self.mask.to(x.device)[None, None] a = logits.softmax(-1) out = (a[..., None] * v).sum(-2).transpose(1, 2).reshape(b, n, self.d) return self.o(out) class GainLayer(nn.Module): def __init__(self, mask, gains): super().__init__(); self.norm1 = nn.LayerNorm(D); self.attn = GainSelfAttention(mask=mask, gains=gains) self.norm2 = nn.LayerNorm(D); self.ff = nn.Sequential(nn.Linear(D,128), nn.ReLU(), nn.Linear(128,D)) def forward(self, x): x = x + self.attn(self.norm1(x)); return x + self.ff(self.norm2(x)) class GainTransformer(nn.Module): def __init__(self, input_dim, out_dim): super().__init__(); mask, gains, edges = make_mask_and_gains(input_dim) self.inp = nn.Linear(1, D); self.pos = nn.Parameter(torch.zeros(1, input_dim, D)); nn.init.normal_(self.pos, std=.02) self.layers = nn.ModuleList([GainLayer(mask, gains) for _ in range(2)]) self.head = nn.Linear(input_dim * D, out_dim); self.edge_count = len(edges)//2 def forward(self, x): h = self.inp(x.unsqueeze(-1)) + self.pos[:, :x.shape[1]] for layer in self.layers: h = layer(h) return self.head(h.reshape(x.shape[0], -1)) def make_base(cfg): def run(seed): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) ds = get_dataset("sequence", seed, NTRAIN, NTEST) net = make_model("transformer_tiny", ds["input_shape"], ds["out_dim"]) _, metric, _ = train_model(net, ds, epochs=EPOCHS, lr=cfg["lr"], batch=128, log=lambda *_: None) return float(metric) return run def make_idea(cfg): def run(seed): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) ds = get_dataset("sequence", seed, NTRAIN, NTEST) net = GainTransformer(ds["input_shape"][0], ds["out_dim"]) _, metric, _ = train_model(net, ds, epochs=EPOCHS, lr=cfg["lr"], batch=128, log=lambda *_: None) return float(metric) return run def mechanism_signature(seed=0): ds = get_dataset("sequence", seed, 64, 32) base = make_model("transformer_tiny", ds["input_shape"], ds["out_dim"]) idea = GainTransformer(ds["input_shape"][0], ds["out_dim"]) x = ds["xte"][:8].clone().requires_grad_(True) y = base(x).sum(); gb = torch.autograd.grad(y, x)[0].abs().mean().item() x2 = ds["xte"][:8].clone().requires_grad_(True) y2 = idea(x2).sum(); gi = torch.autograd.grad(y2, x2)[0].abs().mean().item() comps = 1 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)} def main(): grid = [{"lr": 1e-3}, {"lr": 3e-3}, {"lr": 1e-2}] base = sweep_baseline(make_base, grid, seeds=(0,1,2,3)) idea_runs = [{"cfg": c, "full": evaluate(make_idea(c), SEEDS)} for c in grid] best = min(idea_runs, key=lambda z: z["full"]["mean"]) 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)}) with open("bench_report.json", "w") as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == "__main__": main()