"""MVP for moment-sharp spectral-norm control (K=2). For nonnegative squared singular values x_i, fixed m1=sum x_i and m2=sum x_i^2 imply the sharp maximum U2 = m1/d + sqrt((d-1)*(d*m2-m1**2))/d. Equality is attained by (U2, (m1-U2)/(d-1), ...), when feasible. """ import json, math, time from pathlib import Path import numpy as np import torch from torch import nn from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler SEED = 17 def u2_from_moments(m1, m2, d, eps=0.0): disc = max(0.0, d*m2 - m1*m1) u = m1/d + math.sqrt((d-1)*disc)/d return max(float(u), eps) def exact_moments(W): # W is a 2-D array; d is the number of squared singular values. s2 = np.linalg.svd(W, compute_uv=False)**2 # W^T W has W.shape[1] eigenvalues; rectangular W has trailing zeros. d = W.shape[1] return float(s2.sum()), float((s2*s2).sum()), d, float(s2.max()) def hutchinson_moments(W, probes=8, seed=0): """Unbiased trace estimates for W^T W and (W^T W)^2.""" rng = np.random.default_rng(seed) A = W.T @ W vals1, vals2 = [], [] for _ in range(probes): v = rng.choice([-1., 1.], size=A.shape[0]) Av = A @ v vals1.append(v @ Av) vals2.append(Av @ Av) return float(np.mean(vals1)), float(np.mean(vals2)) def sanity_check(): rng = np.random.default_rng(SEED) rows = [] gaps = [] for d in [3, 5, 10, 24]: for _ in range(100): x = np.exp(rng.normal(size=d)) m1, m2 = x.sum(), (x*x).sum() bound = u2_from_moments(m1, m2, d) gaps.append(bound-x.max()) # Construct the equality spectrum implied by the K=2 solution. rest = (m1-bound)/(d-1) recon = np.r_[bound, np.full(d-1, rest)] assert rest >= -1e-10 assert abs(recon.sum()-m1) < 1e-8*max(1,m1) assert abs((recon**2).sum()-m2) < 1e-7*max(1,m2) assert bound >= x.max()-1e-9 rows.append((d, float(np.mean(gaps)), float(np.max(gaps)))) # A clustered spectrum is exactly recovered, demonstrating sharpness. x = np.array([9., 2., 2., 2., 2.]) exact = u2_from_moments(x.sum(), (x*x).sum(), len(x)) assert abs(exact-x.max()) < 1e-10 # Hutchinson is deliberately tested as an estimator, not treated as exact. W = np.diag(np.sqrt(np.array([9., 4., 1., .25]))) h1,h2 = hutchinson_moments(W, probes=2000, seed=3) e1,e2,_,_ = exact_moments(W) return {"random_bound_minus_true_max": rows, "max_violation": float(-min(g[1] for g in rows)), "clustered_exact_bound": float(exact), "hutchinson_2000_probe_abs_error": [abs(h1-e1),abs(h2-e2)]} class MLP(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential(nn.Linear(64,64),nn.ReLU(),nn.Linear(64,32),nn.ReLU(),nn.Linear(32,10)) def forward(self,x): return self.net(x) def spectral_stats(model): out=[] for layer in model.modules(): if isinstance(layer, nn.Linear): W=layer.weight.detach().cpu().numpy() m1,m2,d,true=exact_moments(W) out.append((math.sqrt(true), math.sqrt(u2_from_moments(m1,m2,d)))) return out def train(control, Xtr, ytr, Xva, yva, steps=180): torch.manual_seed(SEED) model=MLP() opt=torch.optim.AdamW(model.parameters(),lr=2e-3,weight_decay=1e-4) lossfn=nn.CrossEntropyLoss() rng=np.random.default_rng(SEED) t0=time.perf_counter(); losses=[]; spikes=[] model.train() for step in range(steps): ix=rng.integers(0,len(Xtr),size=128) xb=torch.tensor(Xtr[ix],dtype=torch.float32); yb=torch.tensor(ytr[ix],dtype=torch.long) opt.zero_grad(); loss=lossfn(model(xb),yb); loss.backward() gn=float(torch.nn.utils.clip_grad_norm_(model.parameters(),1e9)) spikes.append(gn); opt.step() # Moment control: exact K=2 moments are used here to isolate bound quality. # q=4 means squared spectral radius target; rescale only when exceeded. if control: with torch.no_grad(): for layer in model.modules(): if isinstance(layer,nn.Linear): W=layer.weight.detach().cpu().numpy() m1,m2,d,_=exact_moments(W); U=u2_from_moments(m1,m2,d) if U>4.0: layer.weight.mul_(math.sqrt(4.0/(U+1e-12))) losses.append(float(loss)) model.eval() with torch.no_grad(): acc=float((model(torch.tensor(Xva,dtype=torch.float32)).argmax(1).numpy()==yva).mean()) stats=spectral_stats(model) assert all(bound >= true - 1e-5 for true,bound in stats), stats return {"final_loss":losses[-1],"mean_last20_loss":float(np.mean(losses[-20:])),"val_accuracy":acc, "max_grad_norm":max(spikes),"seconds":time.perf_counter()-t0, "true_sigma_max":max(x[0] for x in stats),"moment_bound_sigma_max":max(x[1] for x in stats), "per_layer_true_sigma": [x[0] for x in stats], "per_layer_bound_sigma": [x[1] for x in stats]} def mini_experiment(): z=load_digits(); X=StandardScaler().fit_transform(z.data).astype(np.float32); y=z.target Xtr,Xva,ytr,yva=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y) return {"weight_decay_baseline":train(False,Xtr,ytr,Xva,yva), "K2_moment_rescale":train(True,Xtr,ytr,Xva,yva)} if __name__ == '__main__': result={"sanity":sanity_check(),"experiment":mini_experiment()} Path('results.json').write_text(json.dumps(result,indent=2)) print(json.dumps(result,indent=2))