Positive-Measure Span Regularizer / run_experiment.py
Mechanism confirmed, baseline not beaten
1import json, random, time
2import numpy as np
3import torch
4from torch import nn
5from torch.utils.data import DataLoader, TensorDataset
6from sklearn.datasets import load_digits
7from sklearn.model_selection import train_test_split
8
9SEED = 460
10random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
11try:
12 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
13 if device.type == 'cuda':
14 torch.cuda.manual_seed_all(SEED)
15except Exception:
16 device = torch.device('cpu')
17
18def rank_proxy(z, tau=0.15):
19 # z rows are samples; Frobenius normalization makes this insensitive to scale.
20 z = z / (torch.linalg.norm(z, ord='fro') + 1e-8)
21 s = torch.linalg.svdvals(z)
22 return (s*s / (s*s + tau*tau)).sum()
23
24def span_loss(emb, labels, tau=0.15, beta=5.0):
25 vals=[]
26 for c in torch.unique(labels):
27 g=emb[labels==c]
28 if g.shape[0] >= 3:
29 vals.append(rank_proxy(g, tau))
30 if not vals:
31 return emb.sum()*0
32 r=torch.stack(vals)
33 # stable soft minimum: -logsumexp(-beta*r)/beta; the stated loss
34 return torch.logsumexp(-beta*r, dim=0)/beta
35
36def direct_rank(A, tau=.15):
37 with torch.no_grad(): return float(rank_proxy(torch.tensor(A, dtype=torch.float32),tau))
38
39def toy_check():
40 rng=np.random.default_rng(SEED)
41 # Same scale, one group exactly one-dimensional and one full-rank.
42 u=rng.normal(size=(24,1)); collapsed=np.repeat(u, 8, axis=1)
43 diverse=rng.normal(size=(24,8))
44 # A scale check is important because the implementation normalizes first.
45 scale=direct_rank(diverse); scaled=direct_rank(100*diverse)
46 # Optimize a collapsed group with the differentiable proxy.
47 z=torch.tensor(collapsed, dtype=torch.float32, requires_grad=True)
48 opt=torch.optim.Adam([z], lr=.08)
49 before=float(rank_proxy(z))
50 for _ in range(120):
51 opt.zero_grad(); loss=-rank_proxy(z); loss.backward(); opt.step()
52 after=float(rank_proxy(z))
53 return {'collapsed_rank':direct_rank(collapsed), 'diverse_rank':direct_rank(diverse),
54 'scale_rank':scale, 'scaled_rank':scaled, 'optimized_collapsed_before':before,
55 'optimized_collapsed_after':after, 'optimization_increased_rank':after>before+0.1}
56
57class Net(nn.Module):
58 def __init__(self):
59 super().__init__()
60 self.enc=nn.Sequential(nn.Linear(64,64),nn.ReLU(),nn.Linear(64,16))
61 self.head=nn.Linear(16,10)
62 def forward(self,x):
63 z=self.enc(x); return self.head(z),z
64
65def eval_model(model, loader):
66 model.eval(); correct=0; total=0; ranks=[]
67 with torch.no_grad():
68 for x,y in loader:
69 x=x.to(device); y=y.to(device); logits,z=model(x)
70 correct += (logits.argmax(1)==y).sum().item(); total += y.numel()
71 for c in torch.unique(y):
72 g=z[y==c]
73 if len(g)>=3: ranks.append(float(rank_proxy(g)))
74 return correct/total, float(np.percentile(ranks,10)), float(np.mean(ranks))
75
76def train(use_span):
77 X,y=load_digits(return_X_y=True)
78 X=X.astype('float32')/16.0
79 xtr,xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
80 tr=DataLoader(TensorDataset(torch.tensor(xtr),torch.tensor(ytr,dtype=torch.long)),batch_size=128,shuffle=True)
81 te=DataLoader(TensorDataset(torch.tensor(xte),torch.tensor(yte,dtype=torch.long)),batch_size=128,shuffle=False)
82 torch.manual_seed(SEED)
83 model=Net().to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3)
84 ce=nn.CrossEntropyLoss(); span_history=[]; t=time.time()
85 for epoch in range(25):
86 model.train()
87 for x,yb in tr:
88 x=x.to(device); yb=yb.to(device); logits,z=model(x)
89 loss=ce(logits,yb)
90 if use_span:
91 # label-defined groups are a reproducible local-region approximation.
92 loss=loss + .01*span_loss(z,yb,tau=.15,beta=5.)
93 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.); opt.step()
94 if epoch in (0,4,9,24): span_history.append(float(loss.detach().cpu()))
95 acc,p10,mean=eval_model(model,te)
96 return {'accuracy':acc,'test_p10_rank':p10,'test_mean_rank':mean,
97 'last_losses':span_history,'seconds':time.time()-t}
98
99def main():
100 toy=toy_check()
101 baseline=train(False); idea=train(True)
102 out={'device':str(device),'toy':toy,'baseline':baseline,'idea':idea,
103 'delta_accuracy':idea['accuracy']-baseline['accuracy'],
104 'delta_p10_rank':idea['test_p10_rank']-baseline['test_p10_rank']}
105 print(json.dumps(out,indent=2))
106 with open('results.json','w') as f: json.dump(out,f,indent=2)
107if __name__=='__main__': main()