import json, random, time import numpy as np import torch from torch import nn from torch.utils.data import DataLoader, TensorDataset from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split SEED = 460 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) try: device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') if device.type == 'cuda': torch.cuda.manual_seed_all(SEED) except Exception: device = torch.device('cpu') def rank_proxy(z, tau=0.15): # z rows are samples; Frobenius normalization makes this insensitive to scale. z = z / (torch.linalg.norm(z, ord='fro') + 1e-8) s = torch.linalg.svdvals(z) return (s*s / (s*s + tau*tau)).sum() def span_loss(emb, labels, tau=0.15, beta=5.0): vals=[] for c in torch.unique(labels): g=emb[labels==c] if g.shape[0] >= 3: vals.append(rank_proxy(g, tau)) if not vals: return emb.sum()*0 r=torch.stack(vals) # stable soft minimum: -logsumexp(-beta*r)/beta; the stated loss return torch.logsumexp(-beta*r, dim=0)/beta def direct_rank(A, tau=.15): with torch.no_grad(): return float(rank_proxy(torch.tensor(A, dtype=torch.float32),tau)) def toy_check(): rng=np.random.default_rng(SEED) # Same scale, one group exactly one-dimensional and one full-rank. u=rng.normal(size=(24,1)); collapsed=np.repeat(u, 8, axis=1) diverse=rng.normal(size=(24,8)) # A scale check is important because the implementation normalizes first. scale=direct_rank(diverse); scaled=direct_rank(100*diverse) # Optimize a collapsed group with the differentiable proxy. z=torch.tensor(collapsed, dtype=torch.float32, requires_grad=True) opt=torch.optim.Adam([z], lr=.08) before=float(rank_proxy(z)) for _ in range(120): opt.zero_grad(); loss=-rank_proxy(z); loss.backward(); opt.step() after=float(rank_proxy(z)) return {'collapsed_rank':direct_rank(collapsed), 'diverse_rank':direct_rank(diverse), 'scale_rank':scale, 'scaled_rank':scaled, 'optimized_collapsed_before':before, 'optimized_collapsed_after':after, 'optimization_increased_rank':after>before+0.1} class Net(nn.Module): def __init__(self): super().__init__() self.enc=nn.Sequential(nn.Linear(64,64),nn.ReLU(),nn.Linear(64,16)) self.head=nn.Linear(16,10) def forward(self,x): z=self.enc(x); return self.head(z),z def eval_model(model, loader): model.eval(); correct=0; total=0; ranks=[] with torch.no_grad(): for x,y in loader: x=x.to(device); y=y.to(device); logits,z=model(x) correct += (logits.argmax(1)==y).sum().item(); total += y.numel() for c in torch.unique(y): g=z[y==c] if len(g)>=3: ranks.append(float(rank_proxy(g))) return correct/total, float(np.percentile(ranks,10)), float(np.mean(ranks)) def train(use_span): X,y=load_digits(return_X_y=True) X=X.astype('float32')/16.0 xtr,xte,ytr,yte=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y) tr=DataLoader(TensorDataset(torch.tensor(xtr),torch.tensor(ytr,dtype=torch.long)),batch_size=128,shuffle=True) te=DataLoader(TensorDataset(torch.tensor(xte),torch.tensor(yte,dtype=torch.long)),batch_size=128,shuffle=False) torch.manual_seed(SEED) model=Net().to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3) ce=nn.CrossEntropyLoss(); span_history=[]; t=time.time() for epoch in range(25): model.train() for x,yb in tr: x=x.to(device); yb=yb.to(device); logits,z=model(x) loss=ce(logits,yb) if use_span: # label-defined groups are a reproducible local-region approximation. loss=loss + .01*span_loss(z,yb,tau=.15,beta=5.) opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),5.); opt.step() if epoch in (0,4,9,24): span_history.append(float(loss.detach().cpu())) acc,p10,mean=eval_model(model,te) return {'accuracy':acc,'test_p10_rank':p10,'test_mean_rank':mean, 'last_losses':span_history,'seconds':time.time()-t} def main(): toy=toy_check() baseline=train(False); idea=train(True) out={'device':str(device),'toy':toy,'baseline':baseline,'idea':idea, 'delta_accuracy':idea['accuracy']-baseline['accuracy'], 'delta_p10_rank':idea['test_p10_rank']-baseline['test_p10_rank']} print(json.dumps(out,indent=2)) with open('results.json','w') as f: json.dump(out,f,indent=2) if __name__=='__main__': main()