Lunar Color-Connectivity Regularizer / lunar_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, random
  2from itertools import combinations, product
  3import numpy as np
  4import torch
  5from torch import nn
  6
  7SEED = 7
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9try:
 10    device = 'cuda' if torch.cuda.is_available() else 'cpu'
 11except Exception:
 12    device = 'cpu'
 13
 14
 15def mec_radius(points):
 16    """Exact brute-force smallest enclosing-circle radius for <=4 points."""
 17    p = np.asarray(points, dtype=float)
 18    if len(p) == 0: return 0.0
 19    best = float('inf')
 20    candidates = [p[i] for i in range(len(p))]
 21    for i,j in combinations(range(len(p)), 2):
 22        candidates.append((p[i]+p[j])/2)
 23    for i,j,k in combinations(range(len(p)), 3):
 24        a,b,c=p[i],p[j],p[k]
 25        d=2*np.cross(b-a,c-a)
 26        if abs(d)>1e-12:
 27            aa=np.dot(a,a); bb=np.dot(b,b); cc=np.dot(c,c)
 28            center=np.array([(aa*(b[1]-c[1])+bb*(c[1]-a[1])+cc*(a[1]-b[1]))/d,
 29                             (aa*(c[0]-b[0])+bb*(a[0]-c[0])+cc*(b[0]-a[0]))/d])
 30            candidates.append(center)
 31    for center in candidates:
 32        best=min(best, float(np.max(np.linalg.norm(p-center,axis=1))))
 33    return best
 34
 35
 36def lunar_cost(points, colors):
 37    """Two-color combinatorial lunar EMST cost (paper Eq. 2.1)."""
 38    points=np.asarray(points,float); colors=np.asarray(colors)
 39    groups=[np.where(colors==c)[0].tolist() for c in sorted(set(colors))]
 40    assert len(groups)==2 and all(groups)
 41    nodes=[]
 42    for i,j in product(groups[0],groups[1]):
 43        nodes.append((i,j,0.5*np.linalg.norm(points[i]-points[j])))
 44    # Kruskal on colorful pair nodes; arc diameter is 2*MEC radius of all four vertices.
 45    edges=[]
 46    for u,v in combinations(range(len(nodes)),2):
 47        ids=list(nodes[u][:2])+list(nodes[v][:2])
 48        edges.append((2*mec_radius(points[ids]),u,v))
 49    edges.sort(); parent=list(range(len(nodes)))
 50    def find(x):
 51        while parent[x]!=x:
 52            parent[x]=parent[parent[x]]; x=parent[x]
 53        return x
 54    total=0.; chosen=0
 55    for w,u,v in edges:
 56        a,b=find(u),find(v)
 57        if a!=b:
 58            parent[a]=b; total += w; chosen += 1
 59            if chosen==len(nodes)-1: break
 60    births=sorted(2*x[2] for x in nodes)
 61    return float(total-sum(births)+births[0]), {'births':births, 'arc_sum':total, 'nodes':len(nodes)}
 62
 63
 64def math_check():
 65    # Verify the paper's nonnegative birth/death cost on fixed and random sets.
 66    cases=[]
 67    for name,p in [
 68        ('separated', [[-2,0],[-2,1],[-1.5,.5],[2,0],[2,1],[1.5,.5]]),
 69        ('mixed', [[-1,0],[0,1],[1,0],[-1,1],[0,0],[1,1]]),
 70        ('collapsed', [[0,0]]*6)]:
 71        c=np.array([0,0,0,1,1,1]); cost,info=lunar_cost(p,c)
 72        cases.append((name,cost,info))
 73    for seed in range(20):
 74        rng=np.random.default_rng(seed)
 75        pp=rng.normal(size=(6,2)); cc=np.array([0,0,0,1,1,1])
 76        val,_=lunar_cost(pp,cc)
 77        assert val >= -1e-8, (seed,val)
 78    assert all(x[1] >= -1e-8 for x in cases)
 79    return cases + [('random_nonnegative_cases',20, {})]
 80
 81class Net(nn.Module):
 82    def __init__(self):
 83        super().__init__(); self.f=nn.Sequential(nn.Linear(2,24),nn.Tanh(),nn.Linear(24,2)); self.head=nn.Linear(2,3)
 84    def forward(self,x):
 85        z=self.f(x); return z,self.head(z)
 86
 87def differentiable_pair_penalty(z, colors):
 88    # engineering approximation: mean cross-color pair radius, normalized by batch diameter.
 89    a=z[colors==0]; b=z[colors==1]
 90    d=torch.cdist(a,b)
 91    diam=torch.pdist(z).max().clamp_min(1e-5) if len(z)>1 else z.new_tensor(1.)
 92    return (d.min(dim=1).values.mean()+d.min(dim=0).values.mean())/(2*diam)
 93
 94def metric(model,X,Y,C):
 95    with torch.no_grad():
 96        z,log=model(X); pred=log.argmax(1); acc=(pred==Y).float().mean().item()
 97        a=z[C==0]; b=z[C==1]; cross=torch.cdist(a,b).min(1).values.mean().item()
 98        within=torch.pdist(z).mean().item(); spread=z.std().item()
 99        D=torch.cdist(z,z); D.fill_diagonal_(float('inf'))
100        nnidx=D.topk(3,largest=False).indices
101        knn=(Y[nnidx].mode(dim=1).values==Y).float().mean().item()
102    # exact cost on a small deterministic subset
103    n=min(12,len(X)); cost,_=lunar_cost(z[:n].cpu().numpy(),C[:n].cpu().numpy())
104    diameter=torch.pdist(z[:n]).max().item() if n>1 else 1.0
105    return {'ce_acc':acc,'cross_color_nearest':cross,'mean_pair_distance':within,
106            'embedding_std':spread,'knn3_acc':knn,'lunar_cost_subset':cost,
107            'lunar_cost_normalized':cost/max(diameter,1e-8)}
108
109def train(lam):
110    torch.manual_seed(SEED)
111    n=180; centers=torch.tensor([[-2.,0.],[0.,2.],[2.,0.]])
112    y=torch.arange(3).repeat_interleave(n//3)
113    x=centers[y]+.55*torch.randn(n,2)
114    # two augmentation identities/colors for every class, balanced in each minibatch
115    c=torch.arange(n)%2
116    model=Net().to(device); opt=torch.optim.Adam(model.parameters(),lr=.025)
117    X,Y,C=x.to(device),y.to(device),c.to(device)
118    for step in range(250):
119        idx=torch.randperm(n,device=device)[:96]; z,log=model(X[idx])
120        loss=nn.functional.cross_entropy(log,Y[idx])
121        if lam: loss=loss+lam*differentiable_pair_penalty(z,C[idx])
122        opt.zero_grad(); loss.backward(); opt.step()
123    return metric(model,X,Y,C)
124
125def main():
126    global device
127    try:
128        out={'device':device,'math_check':math_check()}
129        out['baseline']=train(0.0); out['lunar_lambda_0.5']=train(.5); out['lunar_lambda_2']=train(2.0)
130    except Exception as exc:
131        if device != 'cuda': raise
132        device='cpu'
133        out={'device':device,'cuda_error':repr(exc),'math_check':math_check()}
134        out['baseline']=train(0.0); out['lunar_lambda_0.5']=train(.5); out['lunar_lambda_2']=train(2.0)
135    print(json.dumps(out,indent=2))
136
137if __name__=='__main__': main()