import json, math, random from itertools import combinations, product import numpy as np import torch from torch import nn SEED = 7 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) try: device = 'cuda' if torch.cuda.is_available() else 'cpu' except Exception: device = 'cpu' def mec_radius(points): """Exact brute-force smallest enclosing-circle radius for <=4 points.""" p = np.asarray(points, dtype=float) if len(p) == 0: return 0.0 best = float('inf') candidates = [p[i] for i in range(len(p))] for i,j in combinations(range(len(p)), 2): candidates.append((p[i]+p[j])/2) for i,j,k in combinations(range(len(p)), 3): a,b,c=p[i],p[j],p[k] d=2*np.cross(b-a,c-a) if abs(d)>1e-12: aa=np.dot(a,a); bb=np.dot(b,b); cc=np.dot(c,c) center=np.array([(aa*(b[1]-c[1])+bb*(c[1]-a[1])+cc*(a[1]-b[1]))/d, (aa*(c[0]-b[0])+bb*(a[0]-c[0])+cc*(b[0]-a[0]))/d]) candidates.append(center) for center in candidates: best=min(best, float(np.max(np.linalg.norm(p-center,axis=1)))) return best def lunar_cost(points, colors): """Two-color combinatorial lunar EMST cost (paper Eq. 2.1).""" points=np.asarray(points,float); colors=np.asarray(colors) groups=[np.where(colors==c)[0].tolist() for c in sorted(set(colors))] assert len(groups)==2 and all(groups) nodes=[] for i,j in product(groups[0],groups[1]): nodes.append((i,j,0.5*np.linalg.norm(points[i]-points[j]))) # Kruskal on colorful pair nodes; arc diameter is 2*MEC radius of all four vertices. edges=[] for u,v in combinations(range(len(nodes)),2): ids=list(nodes[u][:2])+list(nodes[v][:2]) edges.append((2*mec_radius(points[ids]),u,v)) edges.sort(); parent=list(range(len(nodes))) def find(x): while parent[x]!=x: parent[x]=parent[parent[x]]; x=parent[x] return x total=0.; chosen=0 for w,u,v in edges: a,b=find(u),find(v) if a!=b: parent[a]=b; total += w; chosen += 1 if chosen==len(nodes)-1: break births=sorted(2*x[2] for x in nodes) return float(total-sum(births)+births[0]), {'births':births, 'arc_sum':total, 'nodes':len(nodes)} def math_check(): # Verify the paper's nonnegative birth/death cost on fixed and random sets. cases=[] for name,p in [ ('separated', [[-2,0],[-2,1],[-1.5,.5],[2,0],[2,1],[1.5,.5]]), ('mixed', [[-1,0],[0,1],[1,0],[-1,1],[0,0],[1,1]]), ('collapsed', [[0,0]]*6)]: c=np.array([0,0,0,1,1,1]); cost,info=lunar_cost(p,c) cases.append((name,cost,info)) for seed in range(20): rng=np.random.default_rng(seed) pp=rng.normal(size=(6,2)); cc=np.array([0,0,0,1,1,1]) val,_=lunar_cost(pp,cc) assert val >= -1e-8, (seed,val) assert all(x[1] >= -1e-8 for x in cases) return cases + [('random_nonnegative_cases',20, {})] class Net(nn.Module): def __init__(self): super().__init__(); self.f=nn.Sequential(nn.Linear(2,24),nn.Tanh(),nn.Linear(24,2)); self.head=nn.Linear(2,3) def forward(self,x): z=self.f(x); return z,self.head(z) def differentiable_pair_penalty(z, colors): # engineering approximation: mean cross-color pair radius, normalized by batch diameter. a=z[colors==0]; b=z[colors==1] d=torch.cdist(a,b) diam=torch.pdist(z).max().clamp_min(1e-5) if len(z)>1 else z.new_tensor(1.) return (d.min(dim=1).values.mean()+d.min(dim=0).values.mean())/(2*diam) def metric(model,X,Y,C): with torch.no_grad(): z,log=model(X); pred=log.argmax(1); acc=(pred==Y).float().mean().item() a=z[C==0]; b=z[C==1]; cross=torch.cdist(a,b).min(1).values.mean().item() within=torch.pdist(z).mean().item(); spread=z.std().item() D=torch.cdist(z,z); D.fill_diagonal_(float('inf')) nnidx=D.topk(3,largest=False).indices knn=(Y[nnidx].mode(dim=1).values==Y).float().mean().item() # exact cost on a small deterministic subset n=min(12,len(X)); cost,_=lunar_cost(z[:n].cpu().numpy(),C[:n].cpu().numpy()) diameter=torch.pdist(z[:n]).max().item() if n>1 else 1.0 return {'ce_acc':acc,'cross_color_nearest':cross,'mean_pair_distance':within, 'embedding_std':spread,'knn3_acc':knn,'lunar_cost_subset':cost, 'lunar_cost_normalized':cost/max(diameter,1e-8)} def train(lam): torch.manual_seed(SEED) n=180; centers=torch.tensor([[-2.,0.],[0.,2.],[2.,0.]]) y=torch.arange(3).repeat_interleave(n//3) x=centers[y]+.55*torch.randn(n,2) # two augmentation identities/colors for every class, balanced in each minibatch c=torch.arange(n)%2 model=Net().to(device); opt=torch.optim.Adam(model.parameters(),lr=.025) X,Y,C=x.to(device),y.to(device),c.to(device) for step in range(250): idx=torch.randperm(n,device=device)[:96]; z,log=model(X[idx]) loss=nn.functional.cross_entropy(log,Y[idx]) if lam: loss=loss+lam*differentiable_pair_penalty(z,C[idx]) opt.zero_grad(); loss.backward(); opt.step() return metric(model,X,Y,C) def main(): global device try: out={'device':device,'math_check':math_check()} out['baseline']=train(0.0); out['lunar_lambda_0.5']=train(.5); out['lunar_lambda_2']=train(2.0) except Exception as exc: if device != 'cuda': raise device='cpu' out={'device':device,'cuda_error':repr(exc),'math_check':math_check()} out['baseline']=train(0.0); out['lunar_lambda_0.5']=train(.5); out['lunar_lambda_2']=train(2.0) print(json.dumps(out,indent=2)) if __name__=='__main__': main()