Smith-normal-form Cayley positional encoding / experiment.py
Beats tuned baseline
1"""Reproducible MVP for metric-parallel SNF Cayley positional coordinates."""
2from collections import deque
3from dataclasses import dataclass, asdict
4import json
5import numpy as np
6import sympy as sp
7from sympy.matrices.normalforms import smith_normal_form
8from sympy.polys.domains import ZZ
9
10@dataclass
11class Result:
12 name: str; vertices: int; edges: int; classes: int
13 cycle_rank: int; snf_torsion: list; free_rank: int
14 path_checks: int; path_failures: int; edge_checks: int; edge_failures: int
15 raw_dim: int; quotient_dim: int; compression_ratio: float
16
17def cycle_graph(n): return list(range(n)), [(i,(i+1)%n) for i in range(n)]
18def torus_graph(m):
19 v=[(i,j) for i in range(m) for j in range(m)]; e=[]
20 for i,j in v: e += [((i,j),((i+1)%m,j)),((i,j),(i,(j+1)%m))]
21 return v,e
22def petersen_graph():
23 v=list(range(10)); e=[]
24 for i in range(5): e += [(i,(i+1)%5),(i,5+i),(5+i,5+(i+2)%5)]
25 return v,e
26
27def distances(v,e):
28 a={x:[] for x in v}
29 for x,y in e: a[x].append(y); a[y].append(x)
30 out={}
31 for s in v:
32 d={s:0}; q=deque([s])
33 while q:
34 x=q.popleft()
35 for y in a[x]:
36 if y not in d: d[y]=d[x]+1; q.append(y)
37 out[s]=d
38 return out
39
40def phi_classes(v,e):
41 """Connected components of the paper's undirected metric-parallel relation."""
42 d=distances(v,e); p=list(range(len(e)))
43 def f(x):
44 while p[x]!=x: p[x]=p[p[x]]; x=p[x]
45 return x
46 def u(x,y):
47 x,y=f(x),f(y)
48 if x!=y: p[y]=x
49 for i,(a,b) in enumerate(e):
50 for j in range(i):
51 x,y=e[j]
52 if (d[a][x]==d[b][y] and d[a][y]==d[b][x]) or (d[a][y]==d[b][x] and d[a][x]==d[b][y]): u(i,j)
53 roots={}; cls=[]
54 for i in range(len(e)):
55 r=f(i); roots.setdefault(r,len(roots)); cls.append(roots[r])
56 return cls
57
58def spanning_tree(v,e):
59 idx={x:i for i,x in enumerate(v)}; a=[[] for _ in v]
60 for ei,(x,y) in enumerate(e):
61 a[idx[x]].append((idx[y],ei,1)); a[idx[y]].append((idx[x],ei,-1))
62 seen={0}; parent={}; tree=set(); q=deque([0])
63 while q:
64 x=q.popleft()
65 for y,ei,s in a[x]:
66 if y not in seen: seen.add(y); parent[y]=(x,ei,s); tree.add(ei); q.append(y)
67 return idx,parent,tree
68
69def path_edges(a, b, parent):
70 """Edges on the directed tree path from vertex a to vertex b."""
71 up_a=[]; x=a
72 while x != 0:
73 par,ei,s = parent[x]
74 up_a.append((x,par,ei,s)); x=par
75 up_b=[]; x=b
76 while x != 0:
77 par,ei,s = parent[x]
78 up_b.append((x,par,ei,s)); x=par
79 anc_a={x:i for i,(x,_,_,_) in enumerate(up_a)}; anc_a[0]=len(up_a)
80 lca=next((x for x,_,_,_ in up_b if x in anc_a), 0)
81 if lca not in anc_a: lca=0
82 out=[]; x=a
83 while x != lca:
84 par,ei,s=parent[x]; out.append((ei,-s)); x=par
85 down=[]; x=b
86 while x != lca:
87 par,ei,s=parent[x]; down.append((ei,s)); x=par
88 out.extend(reversed(down))
89 return out
90
91def cycle_matrix(v,e,cls):
92 idx,parent,tree=spanning_tree(v,e); k=max(cls)+1; cols=[]
93 for ei,(x,y) in enumerate(e):
94 if ei in tree: continue
95 c=[0]*k; c[cls[ei]]+=1
96 for te,s in path_edges(idx[y],idx[x],parent): c[cls[te]]+=s
97 cols.append(c)
98 return sp.Matrix(k,len(cols),lambda i,j: cols[j][i]) if cols else sp.zeros(k,0),idx,parent,tree
99
100def run(name, v, e):
101 cls=phi_classes(v,e); B,idx,parent,tree=cycle_matrix(v,e,cls); k=max(cls)+1
102 D=smith_normal_form(B,domain=ZZ); tors=[]
103 for i in range(min(D.rows,D.cols)):
104 q=abs(int(D[i,i]))
105 if q>1: tors.append(q)
106 rank=sum(1 for i in range(min(D.rows,D.cols)) if D[i,i]!=0)
107 free=k-rank; quotient=len(tors)+free
108 # Exact certificate: every fundamental cycle is a lattice generator.
109 # Also verify every edge increment against the tree-integrated labels.
110 checks=0; fail=0; edgefail=0
111 tree_cols={}
112 non_tree=list(ei for ei in range(len(e)) if ei not in tree)
113 for j,ei in enumerate(non_tree):
114 checks += 1
115 fail += int(any(not z.is_Integer for z in B[:,j]))
116 tree_cols[ei]=j
117 # Tree integration gives z; non-tree edge discrepancies must equal its
118 # fundamental cycle (up to the fixed orientation convention).
119 labels={v[0]:np.zeros(k,dtype=int)}; q=deque([v[0]]) ; adj={x:[] for x in v}
120 for ei,(a,b) in enumerate(e): adj[a].append((b,ei,1)); adj[b].append((a,ei,-1))
121 while q:
122 x=q.popleft()
123 for y,ei,sgn in adj[x]:
124 if y not in labels:
125 labels[y]=labels[x].copy(); labels[y][cls[ei]]+=sgn; q.append(y)
126 for ei,(a,b) in enumerate(e):
127 delta=labels[b]-labels[a]; sigma=np.zeros(k,dtype=int); sigma[cls[ei]]=1
128 if ei in tree:
129 edgefail += int(np.any(delta-sigma))
130 else:
131 # The discrepancy is a signed fundamental cycle; test either sign.
132 col=np.array(B[:,tree_cols[ei]],dtype=int).ravel()
133 edgefail += int(not (np.array_equal(delta-sigma,col) or np.array_equal(delta-sigma,-col)))
134 return Result(name,len(v),len(e),k,B.cols,tors,free,checks,fail,len(e),edgefail,k,quotient,round(k/quotient,3) if quotient else None)
135
136def main():
137 cases=[('C5',*cycle_graph(5)),('C6',*cycle_graph(6)),('Torus3',*torus_graph(3)),('Petersen',*petersen_graph())]
138 out=[asdict(run(n,v,e)) for n,v,e in cases]
139 print(json.dumps(out,indent=2))
140 assert all(x['path_failures']==0 and x['edge_failures']==0 for x in out)
141
142if __name__=='__main__': main()