Cut-Aware Augmentation Filtering / experiment.py

Mechanism confirmed, baseline not beaten

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