Lunar Color-Connectivity Regularizer / lunar_experiment.py
Mechanism confirmed, baseline not beaten
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()