Moment-Sharp Spectral-Norm Control / moment_sharp.py

Mechanism confirmed, baseline not beaten

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