Cut-Aware Augmentation Filtering / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, random
2from pathlib import Path
3import numpy as np
4from sklearn.neighbors import NearestNeighbors
5from sklearn.model_selection import train_test_split
6import torch
7import torch.nn as nn
8
9SEED = 523
10np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
11
12def graph_from_x(x, k=10, sigma=1.0):
13 n = len(x)
14 d = ((x[:, None, :] - x[None, :, :]) ** 2).sum(-1)
15 nnidx = np.argsort(d, axis=1)[:, 1:k+1]
16 W = np.zeros((n, n), dtype=np.float64)
17 for i in range(n):
18 for j in nnidx[i]: W[i, j] = math.exp(-d[i, j] / (2*sigma*sigma))
19 W = np.maximum(W, W.T); np.fill_diagonal(W, 0)
20 return W
21
22def cut_ratio(W, y):
23 return float((W * (y[:,None] != y[None,:])).sum() / (W.sum() + 1e-12))
24
25def weights_from_cuts(cuts, beta=12.0, qmin=0.0):
26 z = -beta*np.asarray(cuts); z -= z.max(); q = np.exp(z); q /= q.sum()
27 if qmin:
28 q = np.maximum(q, qmin); q /= q.sum()
29 return q
30
31def verify_math():
32 # Two clearly separated groups and a graph with only within-group edges.
33 y = np.array([0,0,0,1,1,1])
34 Wgood = np.zeros((6,6)); Wbad = np.ones((6,6))-np.eye(6)
35 for a in range(3):
36 for b in range(3): Wgood[a,b+3] = Wgood[b+3,a] = 0
37 for i in range(3):
38 for j in range(3):
39 if i != j: Wgood[i,j] = Wgood[i+3,j+3] = 1
40 rg, rb = cut_ratio(Wgood,y), cut_ratio(Wbad,y)
41 q = weights_from_cuts([rg, rb], beta=10)
42 assert rg < rb and q[0] > q[1] and abs(q.sum()-1) < 1e-12
43 # Directly check the stated ratio is invariant to global graph scaling.
44 assert abs(cut_ratio(7.3*Wbad,y)-rb) < 1e-12
45 return {"good_cut":rg, "bad_cut":rb, "q_beta10":q.tolist(), "scale_invariance_error":abs(cut_ratio(7.3*Wbad,y)-rb)}
46
47class Net(nn.Module):
48 def __init__(self):
49 super().__init__(); self.f = nn.Sequential(nn.Linear(2,32),nn.ReLU(),nn.Linear(32,32),nn.ReLU(),nn.Linear(32,3))
50 def forward(self,x): return self.f(x)
51
52def make_data(seed=SEED):
53 rng=np.random.RandomState(seed)
54 centers=np.array([[-2.,0.],[2.,0.],[0.,2.8]])
55 xs=[]; ys=[]
56 for c in range(3):
57 xs.append(centers[c]+rng.randn(90,2)*0.62); ys += [c]*90
58 return np.vstack(xs).astype('float32'), np.array(ys,dtype=np.int64)
59
60def train_once(mode, X, y, graphs, cuts, seed):
61 torch.manual_seed(seed)
62 # fixed 12 labeled points/class, remaining nodes participate only in graph term
63 labeled=np.concatenate([np.where(y==c)[0][:12] for c in range(3)])
64 xt=torch.tensor(X); yt=torch.tensor(y); lt=torch.tensor(labeled)
65 model=Net(); opt=torch.optim.Adam(model.parameters(),lr=0.025)
66 if mode=='uniform': q=np.ones(len(graphs))/len(graphs)
67 elif mode=='cut': q=weights_from_cuts(cuts,beta=12.)
68 else: raise ValueError(mode)
69 W=sum(float(qi)*Wi for qi,Wi in zip(q,graphs))
70 # normalize graph penalty to comparable scale across policies
71 Wt=torch.tensor(W,dtype=torch.float32); denom=Wt.sum().clamp_min(1e-8)
72 for _ in range(260):
73 opt.zero_grad(); logits=model(xt); p=torch.softmax(logits,1)
74 sup=nn.functional.cross_entropy(logits[lt],yt[lt])
75 dif=(p[:,None,:]-p[None,:,:]).pow(2).sum(-1)
76 reg=(Wt*dif).sum()/denom
77 (sup+0.8*reg).backward(); opt.step()
78 with torch.no_grad():
79 pred=model(xt).argmax(1).numpy(); p=torch.softmax(model(xt),1).numpy()
80 return {"accuracy":float((pred==y).mean()),"labeled_accuracy":float((pred[labeled]==y[labeled]).mean()),"graph_reg":float(reg.item()),"q":q.tolist(),"pred_cut_uniform_graph":cut_ratio(W,pred)}
81
82def main():
83 mathcheck=verify_math(); X,y=make_data()
84 rng=np.random.RandomState(SEED)
85 # Identity and mild noise preserve neighborhoods; shuffled coordinates intentionally mix labels.
86 policies=[X, X+rng.randn(*X.shape).astype('float32')*0.18, X[rng.permutation(len(X))]]
87 graphs=[graph_from_x(z,k=10,sigma=0.8) for z in policies]
88 cuts=[cut_ratio(W,y) for W in graphs]
89 results={m:train_once(m,X,y,graphs,cuts,SEED+17) for m in ['uniform','cut']}
90 out={"math_check":mathcheck,"policy_cuts":cuts,"cut_weights":weights_from_cuts(cuts,12).tolist(),"results":results}
91 Path('results.json').write_text(json.dumps(out,indent=2))
92 print(json.dumps(out,indent=2))
93
94if __name__=='__main__': main()