Mean-field rainbow relation router / rainbow_router_experiment.py
Unverified
1import json, time
2import numpy as np
3
4
5def softmax(x):
6 z = x - np.max(x, axis=-1, keepdims=True)
7 e = np.exp(z)
8 return e / e.sum(axis=-1, keepdims=True)
9
10
11def init_probs(n, beta, seed=0, noise=0.15):
12 rng = np.random.default_rng(seed)
13 logits = 2*np.asarray(beta, float)[None, None, :] + noise*rng.normal(size=(n,n,3))
14 logits = (logits + logits.transpose(1,0,2))/2
15 p = softmax(logits)
16 for i in range(n): p[i,i] = 1/3
17 return p
18
19
20def refine(p, beta, beta4, tau, iterations=8, damping=1.0):
21 n = p.shape[0]
22 beta = np.asarray(beta, float)
23 q = p.copy()
24 for _ in range(iterations):
25 for i in range(n):
26 for j in range(i+1, n):
27 r = np.zeros(3)
28 for k in range(n):
29 if k == i or k == j:
30 continue
31 for a in range(3):
32 o = [x for x in range(3) if x != a]
33 r[a] += q[i,k,o[0]]*q[j,k,o[1]] + q[i,k,o[1]]*q[j,k,o[0]]
34 s = 2*beta + (beta4/n)*r
35 new = softmax((s/tau)[None, :])[0]
36 q[i,j] = q[j,i] = (1-damping)*q[i,j] + damping*new
37 return q
38
39
40def rainbow_density(p):
41 n = p.shape[0]; total = 0.0; count = 0
42 for i in range(n):
43 for j in range(i+1,n):
44 for k in range(j+1,n):
45 x = 0.0
46 for a in range(3):
47 o = [b for b in range(3) if b != a]
48 x += p[i,j,a]*(p[i,k,o[0]]*p[j,k,o[1]] + p[i,k,o[1]]*p[j,k,o[0]])
49 total += x; count += 1
50 return total/count
51
52
53def entropy(p):
54 return float((-p*np.log(np.maximum(p, 1e-12))).sum(axis=-1).mean())
55
56
57def run():
58 beta = np.array([0.10, -0.03, -0.07])
59 out = {'math_checks': {}, 'mechanism': [], 'scaling': [], 'proxy': {}}
60 # beta4=0 must make refinement identical to the unary router.
61 p0 = init_probs(7, beta, seed=4)
62 pzero = refine(p0, beta, 0.0, 0.7, iterations=5)
63 unary = softmax((2*beta/0.7)[None, :])[0]
64 out['math_checks']['beta4_zero_max_change_from_unary'] = float(np.max(np.abs(pzero[0,1]-unary)))
65 out['math_checks']['simplex_max_error'] = float(np.max(np.abs(pzero.sum(-1)-1)))
66 # Positive coupling response, using a symmetric unbiased initialization.
67 for b4 in [0., 0.5, 1., 2., 4., 8.]:
68 p = init_probs(8, np.zeros(3), seed=11, noise=0.0)
69 q = refine(p, np.zeros(3), b4, 0.35, iterations=12, damping=0.7)
70 out['mechanism'].append({'beta4': b4, 'rainbow_density': rainbow_density(q), 'entropy': entropy(q), 'max_p': float(q.max())})
71 # Claimed 1/n scaling: compare motif-logit contribution at two graph sizes.
72 for n in [6, 12, 20]:
73 p = init_probs(n, np.zeros(3), seed=3, noise=0.0)
74 # one explicit update's motif score magnitude
75 rvals=[]
76 for i in range(n):
77 for j in range(i+1,n):
78 r=np.zeros(3)
79 for k in range(n):
80 if k in (i,j): continue
81 for a in range(3):
82 o=[x for x in range(3) if x!=a]
83 r[a]+=p[i,k,o[0]]*p[j,k,o[1]]+p[i,k,o[1]]*p[j,k,o[0]]
84 rvals.append(np.max(np.abs(1.0*r/n)))
85 out['scaling'].append({'n': n, 'median_normalized_score': float(np.median(rvals))})
86 # Small fixed classification proxy: target is a noisy unary color; compare unary vs refined routing.
87 rng=np.random.default_rng(21); n=10; d=5; samples=160
88 X=rng.normal(size=(samples,n,d)); true_w=rng.normal(size=(d,3)); y=np.argmax(X@true_w,axis=2)
89 # A router predicts the edge color from endpoint features; evaluate held-out edges.
90 split=120; losses={'baseline':[],'idea':[]}; t0=time.perf_counter()
91 for s in range(samples):
92 logits=np.zeros((n,n,3))
93 for i in range(n):
94 for j in range(i+1,n):
95 logits[i,j]=logits[j,i]=(X[s,i]+X[s,j])@true_w/2
96 p=softmax(logits)
97 if s>=split: losses['baseline'].append(-np.log(np.maximum(p[np.arange(n), (np.arange(n)+1)%n, y[s,np.arange(n)]],1e-12)).mean())
98 q=refine(p, np.zeros(3), 2.0, 0.5, iterations=2, damping=0.5)
99 if s>=split: losses['idea'].append(-np.log(np.maximum(q[np.arange(n), (np.arange(n)+1)%n, y[s,np.arange(n)]],1e-12)).mean())
100 out['proxy']={'baseline_test_nll':float(np.mean(losses['baseline'])), 'idea_test_nll':float(np.mean(losses['idea'])), 'seconds':time.perf_counter()-t0}
101 return out
102
103if __name__ == '__main__':
104 print(json.dumps(run(), indent=2))