Mean-field rainbow relation router / rainbow_router_experiment.py

Unverified

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