Moment-Sharp Spectral-Norm Control / moment_sharp.py
Mechanism confirmed, baseline not beaten
1"""MVP for moment-sharp spectral-norm control (K=2).
2
3For nonnegative squared singular values x_i, fixed m1=sum x_i and
4m2=sum x_i^2 imply the sharp maximum
5 U2 = m1/d + sqrt((d-1)*(d*m2-m1**2))/d.
6Equality is attained by (U2, (m1-U2)/(d-1), ...), when feasible.
7"""
8import json, math, time
9from pathlib import Path
10import numpy as np
11import torch
12from torch import nn
13from sklearn.datasets import load_digits
14from sklearn.model_selection import train_test_split
15from sklearn.preprocessing import StandardScaler
16
17SEED = 17
18
19def u2_from_moments(m1, m2, d, eps=0.0):
20 disc = max(0.0, d*m2 - m1*m1)
21 u = m1/d + math.sqrt((d-1)*disc)/d
22 return max(float(u), eps)
23
24def exact_moments(W):
25 # W is a 2-D array; d is the number of squared singular values.
26 s2 = np.linalg.svd(W, compute_uv=False)**2
27 # W^T W has W.shape[1] eigenvalues; rectangular W has trailing zeros.
28 d = W.shape[1]
29 return float(s2.sum()), float((s2*s2).sum()), d, float(s2.max())
30
31def hutchinson_moments(W, probes=8, seed=0):
32 """Unbiased trace estimates for W^T W and (W^T W)^2."""
33 rng = np.random.default_rng(seed)
34 A = W.T @ W
35 vals1, vals2 = [], []
36 for _ in range(probes):
37 v = rng.choice([-1., 1.], size=A.shape[0])
38 Av = A @ v
39 vals1.append(v @ Av)
40 vals2.append(Av @ Av)
41 return float(np.mean(vals1)), float(np.mean(vals2))
42
43def sanity_check():
44 rng = np.random.default_rng(SEED)
45 rows = []
46 gaps = []
47 for d in [3, 5, 10, 24]:
48 for _ in range(100):
49 x = np.exp(rng.normal(size=d))
50 m1, m2 = x.sum(), (x*x).sum()
51 bound = u2_from_moments(m1, m2, d)
52 gaps.append(bound-x.max())
53 # Construct the equality spectrum implied by the K=2 solution.
54 rest = (m1-bound)/(d-1)
55 recon = np.r_[bound, np.full(d-1, rest)]
56 assert rest >= -1e-10
57 assert abs(recon.sum()-m1) < 1e-8*max(1,m1)
58 assert abs((recon**2).sum()-m2) < 1e-7*max(1,m2)
59 assert bound >= x.max()-1e-9
60 rows.append((d, float(np.mean(gaps)), float(np.max(gaps))))
61 # A clustered spectrum is exactly recovered, demonstrating sharpness.
62 x = np.array([9., 2., 2., 2., 2.])
63 exact = u2_from_moments(x.sum(), (x*x).sum(), len(x))
64 assert abs(exact-x.max()) < 1e-10
65 # Hutchinson is deliberately tested as an estimator, not treated as exact.
66 W = np.diag(np.sqrt(np.array([9., 4., 1., .25])))
67 h1,h2 = hutchinson_moments(W, probes=2000, seed=3)
68 e1,e2,_,_ = exact_moments(W)
69 return {"random_bound_minus_true_max": rows,
70 "max_violation": float(-min(g[1] for g in rows)),
71 "clustered_exact_bound": float(exact),
72 "hutchinson_2000_probe_abs_error": [abs(h1-e1),abs(h2-e2)]}
73
74class MLP(nn.Module):
75 def __init__(self):
76 super().__init__()
77 self.net = nn.Sequential(nn.Linear(64,64),nn.ReLU(),nn.Linear(64,32),nn.ReLU(),nn.Linear(32,10))
78 def forward(self,x): return self.net(x)
79
80def spectral_stats(model):
81 out=[]
82 for layer in model.modules():
83 if isinstance(layer, nn.Linear):
84 W=layer.weight.detach().cpu().numpy()
85 m1,m2,d,true=exact_moments(W)
86 out.append((math.sqrt(true), math.sqrt(u2_from_moments(m1,m2,d))))
87 return out
88
89def train(control, Xtr, ytr, Xva, yva, steps=180):
90 torch.manual_seed(SEED)
91 model=MLP()
92 opt=torch.optim.AdamW(model.parameters(),lr=2e-3,weight_decay=1e-4)
93 lossfn=nn.CrossEntropyLoss()
94 rng=np.random.default_rng(SEED)
95 t0=time.perf_counter(); losses=[]; spikes=[]
96 model.train()
97 for step in range(steps):
98 ix=rng.integers(0,len(Xtr),size=128)
99 xb=torch.tensor(Xtr[ix],dtype=torch.float32); yb=torch.tensor(ytr[ix],dtype=torch.long)
100 opt.zero_grad(); loss=lossfn(model(xb),yb); loss.backward()
101 gn=float(torch.nn.utils.clip_grad_norm_(model.parameters(),1e9))
102 spikes.append(gn); opt.step()
103 # Moment control: exact K=2 moments are used here to isolate bound quality.
104 # q=4 means squared spectral radius target; rescale only when exceeded.
105 if control:
106 with torch.no_grad():
107 for layer in model.modules():
108 if isinstance(layer,nn.Linear):
109 W=layer.weight.detach().cpu().numpy()
110 m1,m2,d,_=exact_moments(W); U=u2_from_moments(m1,m2,d)
111 if U>4.0:
112 layer.weight.mul_(math.sqrt(4.0/(U+1e-12)))
113 losses.append(float(loss))
114 model.eval()
115 with torch.no_grad():
116 acc=float((model(torch.tensor(Xva,dtype=torch.float32)).argmax(1).numpy()==yva).mean())
117 stats=spectral_stats(model)
118 assert all(bound >= true - 1e-5 for true,bound in stats), stats
119 return {"final_loss":losses[-1],"mean_last20_loss":float(np.mean(losses[-20:])),"val_accuracy":acc,
120 "max_grad_norm":max(spikes),"seconds":time.perf_counter()-t0,
121 "true_sigma_max":max(x[0] for x in stats),"moment_bound_sigma_max":max(x[1] for x in stats),
122 "per_layer_true_sigma": [x[0] for x in stats],
123 "per_layer_bound_sigma": [x[1] for x in stats]}
124
125def mini_experiment():
126 z=load_digits(); X=StandardScaler().fit_transform(z.data).astype(np.float32); y=z.target
127 Xtr,Xva,ytr,yva=train_test_split(X,y,test_size=.25,random_state=SEED,stratify=y)
128 return {"weight_decay_baseline":train(False,Xtr,ytr,Xva,yva),
129 "K2_moment_rescale":train(True,Xtr,ytr,Xva,yva)}
130
131if __name__ == '__main__':
132 result={"sanity":sanity_check(),"experiment":mini_experiment()}
133 Path('results.json').write_text(json.dumps(result,indent=2))
134 print(json.dumps(result,indent=2))