Directed distance-curvature positional encoding / experiment.py
Mechanism failed
1import numpy as np
2import torch
3from scipy.linalg import solve
4from scipy.sparse.csgraph import shortest_path
5import time, json
6
7SEED=417
8
9def directed_dist(A):
10 D=shortest_path(A, directed=True, unweighted=True)
11 if not np.isfinite(D).all(): raise ValueError('graph is not strongly connected')
12 return D
13
14def curvature(D, eps=1e-6, clip=5.):
15 n=len(D)
16 # np.linalg.pinv implements the specified Moore-Penrose SVD cutoff convention
17 mo=np.linalg.pinv(D, rcond=eps) @ (n*np.ones(n))
18 mi=np.linalg.pinv(D.T, rcond=eps) @ (n*np.ones(n))
19 scale=(np.abs(mo).sum()+np.abs(mi).sum())/(2*n)+1e-6
20 z=np.stack([np.clip(mo/scale,-clip,clip), np.clip(mi/scale,-clip,clip)],1)
21 return mo,mi,z,scale
22
23def graph(n=96, p=0.10, seed=0):
24 rng=np.random.default_rng(seed)
25 y=np.repeat(np.arange(2), n//2)
26 rng.shuffle(y)
27 # Group 0 has relatively strong outgoing connections; group 1 receives them.
28 P=np.full((2,2),p)
29 P[0,1]=.22; P[1,0]=.035; P[0,0]=.10; P[1,1]=.10
30 A=np.zeros((n,n),dtype=np.float32)
31 for i in range(n):
32 for j in range(n):
33 if i!=j: A[i,j]=rng.random()<P[y[i],y[j]]
34 # Guaranteed strongly-connected directed backbone, negligible as a local feature.
35 for i in range(n): A[i,(i+1)%n]=1
36 D=directed_dist(A)
37 return A,y,D
38
39def math_check():
40 signed_sum_errors=[]; l1_errors=[]; residuals=[]; negative_cases=0
41 for seed in range(20):
42 A,_,D=graph(24,seed=100+seed)
43 mo,mi,_,_=curvature(D)
44 residuals.append(max(np.max(np.abs(D@mo-len(D))),
45 np.max(np.abs(D.T@mi-len(D)))))
46 signed_sum_errors.append(abs(mo.sum()-mi.sum()))
47 l1_errors.append(abs(np.abs(mo).sum()-np.abs(mi).sum()))
48 if (mo < 0).any() or (mi < 0).any(): negative_cases += 1
49 return {
50 'instances': 20,
51 'max_out_in_system_residual': float(max(residuals)),
52 'max_signed_sum_difference': float(max(signed_sum_errors)),
53 'mean_absolute_l1_difference': float(np.mean(l1_errors)),
54 'max_absolute_l1_difference': float(max(l1_errors)),
55 'negative_solution_cases': negative_cases
56 }
57
58class DirectedNet(torch.nn.Module):
59 def __init__(self, in_dim):
60 super().__init__()
61 self.lin=torch.nn.Linear(in_dim*3,24)
62 self.out=torch.nn.Linear(24,2)
63 def forward(self,x,A):
64 # A[i,j]=edge i->j; incoming and outgoing directed aggregations.
65 out=A@x/(A.sum(1,keepdims=True)+1e-6)
66 inc=A.T@x/(A.T.sum(1,keepdims=True)+1e-6)
67 h=torch.relu(self.lin(torch.cat([x,out,inc],1)))
68 return self.out(h)
69
70def run_one(seed, use_curv, labeled_per_class=5):
71 torch.manual_seed(seed); np.random.seed(seed)
72 A,y,D=graph(seed=seed)
73 n=len(y); rng=np.random.default_rng(seed+88)
74 # Features deliberately carry no class information; degrees are the standard local control.
75 X=rng.normal(0,1,(n,4)).astype(np.float32)
76 if use_curv:
77 _,_,z,_=curvature(D)
78 X=np.concatenate([X,z.astype(np.float32)],1)
79 At=torch.tensor(A); Xt=torch.tensor(X); yt=torch.tensor(y,dtype=torch.long)
80 train=[]
81 for c in [0,1]: train.extend(rng.choice(np.where(y==c)[0],labeled_per_class,replace=False))
82 train=torch.tensor(train); test=torch.tensor([i for i in range(n) if i not in set(train.tolist())])
83 model=DirectedNet(X.shape[1]); opt=torch.optim.Adam(model.parameters(),lr=.02,weight_decay=1e-4)
84 for _ in range(250):
85 opt.zero_grad(); logits=model(Xt,At); loss=torch.nn.functional.cross_entropy(logits[train],yt[train]); loss.backward(); opt.step()
86 with torch.no_grad():
87 logits=model(Xt,At); pred=logits.argmax(1); prob=logits.softmax(1)[:,1]
88 acc=(pred[test]==yt[test]).float().mean().item()
89 # Brier score is a simple calibration metric (lower is better).
90 brier=((prob[test]-(yt[test]==1).float())**2).mean().item()
91 return acc,brier
92
93def main():
94 t=time.perf_counter(); check=math_check(); prep=time.perf_counter()-t
95 rows=[]
96 for k in [5,10]:
97 for mode in [False,True]:
98 r=[run_one(100+s,mode,k) for s in range(5)]
99 rows.append({'labels_per_class':k,'curvature':mode,'accuracy_mean':float(np.mean([x[0] for x in r])),'accuracy_std':float(np.std([x[0] for x in r])),'brier_mean':float(np.mean([x[1] for x in r]))})
100 print(json.dumps({'math_check':check,'math_check_seconds':prep,'results':rows},indent=2))
101if __name__=='__main__': main()