Positive-Measure Span Regularizer / run_experiment.py

Mechanism confirmed, baseline not beaten

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