Directed distance-curvature positional encoding / experiment.py

Mechanism failed

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