Gain-Rigid Sparse Attention / experiment.py

✓✓ Beats tuned baseline

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