Gain-Rigid Sparse Attention / experiment.py
Beats tuned baseline
1import json, math, random
2from collections import deque
3import numpy as np
4
5SEED = 418
6random.seed(SEED); np.random.seed(SEED)
7
8# Cyclic group Z_m, represented additively. Its 2-D orthogonal representation is a rotation.
9def inv(g, m): return (-g) % m
10def compose(a, b, m): return (a + b) % m
11def rotation(g, m):
12 t = 2 * math.pi * g / m
13 return np.array([[math.cos(t), -math.sin(t)], [math.sin(t), math.cos(t)]])
14
15def add_edge(edges, u, v, g, m):
16 edges[(u, v)] = g % m
17 edges[(v, u)] = inv(g, m)
18
19def undirected_pairs(edges):
20 return sorted((u, v) for (u, v) in edges if u < v)
21
22def two_extension(edges, f1, f2, m, new_vertex):
23 """Delete directed representatives f1/f2 and insert v -> four endpoints.
24 Labels use a=c=0, b=psi(f1), d=psi(f2), hence a^-1 b=f1 and c^-1 d=f2.
25 Reverse arcs are inserted automatically.
26 """
27 out = dict(edges)
28 (v1, v2), (v3, v4) = f1, f2
29 g1, g2 = out[f1], out[f2]
30 assert out[(v2, v1)] == inv(g1, m) and out[(v4, v3)] == inv(g2, m)
31 del out[(v1, v2)]; del out[(v2, v1)]
32 del out[(v3, v4)]; del out[(v4, v3)]
33 add_edge(out, new_vertex, v1, 0, m)
34 add_edge(out, new_vertex, v2, g1, m)
35 add_edge(out, new_vertex, v3, 0, m)
36 add_edge(out, new_vertex, v4, g2, m)
37 return out, (g1, g2)
38
39def components(edges, n):
40 adj=[[] for _ in range(n)]
41 for u,v in undirected_pairs(edges): adj[u].append(v); adj[v].append(u)
42 seen=set(); cs=[]
43 for s in range(n):
44 if s in seen: continue
45 q=[s]; seen.add(s); c=[]
46 while q:
47 u=q.pop(); c.append(u)
48 for v in adj[u]:
49 if v not in seen: seen.add(v); q.append(v)
50 cs.append(c)
51 return cs
52
53def diameter(edges, n):
54 adj=[[] for _ in range(n)]
55 for u,v in undirected_pairs(edges): adj[u].append(v); adj[v].append(u)
56 best=0
57 for s in range(n):
58 d={s:0}; q=deque([s])
59 while q:
60 u=q.popleft()
61 for v in adj[u]:
62 if v not in d: d[v]=d[u]+1; q.append(v)
63 if len(d)!=n: return float('inf')
64 best=max(best,max(d.values()))
65 return best
66
67def build_rigid(n=32, m=8):
68 # A connected seed, then repeated 2-extensions. Edge count is n+2 here.
69 e={}
70 for i in range(4): add_edge(e, i, (i+1)%4, (i*3+1)%m, m)
71 next_v=4
72 while next_v<n:
73 pairs=undirected_pairs(e)
74 # deterministic, non-identical endpoint pairs when possible
75 f1=pairs[(next_v*3) % len(pairs)]
76 f2=pairs[(next_v*7+1) % len(pairs)]
77 if f1==f2: f2=pairs[(pairs.index(f2)+1)%len(pairs)]
78 e,_=two_extension(e, f1, f2, m, next_v); next_v+=1
79 return e
80
81def random_graph(n, edge_count, rng):
82 allp=[(u,v) for u in range(n) for v in range(u+1,n)]
83 chosen=rng.choice(len(allp), edge_count, replace=False)
84 e={}
85 for j in chosen:
86 u,v=allp[j]; add_edge(e,u,v,int(rng.integers(8)),8)
87 return e
88
89def reach_fraction(edges, n, steps):
90 adj=[set() for _ in range(n)]
91 for u,v in undirected_pairs(edges): adj[u].add(v); adj[v].add(u)
92 reachable={0}
93 for _ in range(steps):
94 reachable |= {v for u in reachable for v in adj[u]}
95 return len(reachable)/n
96
97def propagation(edges, n, steps=8, d=2):
98 # Gain-conditioned linear message passing, row-normalized over undirected neighbors.
99 A=np.zeros((n*d,n*d))
100 deg=[0]*n
101 for u,v in undirected_pairs(edges): deg[u]+=1; deg[v]+=1
102 for (u,v),g in edges.items():
103 if deg[u]: A[u*d:(u+1)*d,v*d:(v+1)*d]=rotation(g,8)/deg[u]
104 x=np.zeros(n*d); x[0]=1.; x[1]=0.
105 for _ in range(steps): x=A@x
106 return float(np.linalg.norm(x.reshape(n,d),axis=1).sum()), float(np.count_nonzero(np.linalg.norm(x.reshape(n,d),axis=1)>1e-10)) / n
107
108def main():
109 n=32; rigid=build_rigid(n); E=len(undirected_pairs(rigid)); rng=np.random.default_rng(SEED)
110 random_stats=[]
111 for _ in range(100):
112 g=random_graph(n,E,rng); random_stats.append((len(components(g,n)),diameter(g,n),reach_fraction(g,n,8),propagation(g,n)[1]))
113 de={}
114 for u in range(n):
115 for v in range(u+1,n): add_edge(de,u,v,(u*13+v*7)%8,8)
116 dense={"edges":len(undirected_pairs(de)),"components":len(components(de,n)),"diameter":diameter(de,n),"reach_8_steps":reach_fraction(de,n,8),"active_fraction":propagation(de,n)[1]}
117 # Verify every extension's defining equations and reverse-edge inverse law independently.
118 check={"reverse_inverse":True,"extension_equations":True}
119 e={}; add_edge(e,0,1,3,8); add_edge(e,1,2,5,8); add_edge(e,2,3,1,8); add_edge(e,3,0,6,8)
120 for (u,v),g in e.items(): check["reverse_inverse"] &= (e[(v,u)]==inv(g,8))
121 old={(u,v):g for (u,v),g in e.items()}
122 # Test the general construction with nonzero a,c: b=a*f1 and d=c*f2.
123 out=dict(e); del out[(0,1)]; del out[(1,0)]; del out[(2,3)]; del out[(3,2)]
124 a,c=3,6
125 add_edge(out,4,0,a,8); add_edge(out,4,1,compose(a,old[(0,1)],8),8)
126 add_edge(out,4,2,c,8); add_edge(out,4,3,compose(c,old[(2,3)],8),8)
127 check["extension_equations"] = (compose(inv(a,8),out[(4,1)],8)==old[(0,1)] and compose(inv(c,8),out[(4,3)],8)==old[(2,3)])
128 check["general_extension_reverse_edges"] = all(out[(v,u)]==inv(g,8) for (u,v),g in out.items())
129 rs=[len(components(rigid,n)),diameter(rigid,n),reach_fraction(rigid,n,8),propagation(rigid,n)[1]]
130 result={"seed":SEED,"n":n,"undirected_edges":E,"math_check":check,
131 "gain_rotation_orthogonality_error":float(np.linalg.norm(rotation(3,8).T@rotation(3,8)-np.eye(2))),
132 "rigid":{"components":rs[0],"diameter":rs[1],"reach_8_steps":rs[2],"active_fraction":rs[3]},
133 "dense_reference":dense,"random_100":{"mean_components":float(np.mean([x[0] for x in random_stats])),"disconnected_rate":float(np.mean([x[0]>1 for x in random_stats])),"mean_diameter_connected":float(np.mean([x[1] for x in random_stats if np.isfinite(x[1])] or [float('inf')])),"mean_reach_8_steps":float(np.mean([x[2] for x in random_stats])),"mean_active_fraction":float(np.mean([x[3] for x in random_stats]))}}
134 with open('results.json','w') as f: json.dump(result,f,indent=2)
135 print(json.dumps(result,indent=2))
136
137if __name__=='__main__': main()