diff --git a/baseline/notebooks/MNIST_MLP3_MuonClip_Polar_Jacobian_Baseline.ipynb b/baseline/notebooks/MNIST_MLP3_MuonClip_Polar_Jacobian_Baseline.ipynb new file mode 100644 index 00000000..df24a971 --- /dev/null +++ b/baseline/notebooks/MNIST_MLP3_MuonClip_Polar_Jacobian_Baseline.ipynb @@ -0,0 +1 @@ +{"nbformat":4,"nbformat_minor":5,"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11"}},"cells":[{"cell_type":"markdown","metadata":{},"source":["# MNIST MLP3 — MuonClip polar Jacobian: raw W ESD vs analytic and numerical response spectra\n","\n","For each matrix checkpoint \\(W=U\\Sigma V^\\top\\), this notebook compares three **raw, unnormalized** spectra:\n","\n","1. the ordinary weight ESD \\(\\lambda(W^\\top W)=\\{\\sigma_i^2\\}\\);\n","2. the exact analytic spectrum of \\(G_W=(D\\Pi(W))^*D\\Pi(W)\\), where \\(\\Pi(W)=UV^\\top\\);\n","3. a numerical single-checkpoint finite-difference response Gram spectrum from isotropic random perturbations.\n","\n","All three are fit by the same in-notebook continuous power-law MLE with a grid search over \\(x_{\\min}\\) minimizing the Kolmogorov--Smirnov distance. No spectral normalization, log binning, KDE, or separate fitting rule is used.\n","\n","For the analytic polar derivative, with \\(A=U^\\top E V\\),\n","\\[\n","\\Omega_{ij}=\\frac{A_{ij}-A_{ji}}{\\sigma_i+\\sigma_j},\\qquad \\Omega_{ii}=0.\n","\\]\n","The nonzero Gram eigenvalues are\n","\\[\n","\\lambda_{ij}^{\\rm rot}=\\frac{4}{(\\sigma_i+\\sigma_j)^2},\\quad ig.shape[1]; x=g.T if t else g; x=x.float()/torch.linalg.vector_norm(x.float()).clamp_min(eps); a,b,c=3.4445,-4.7750,2.0315\n"," for _ in range(steps): q=x@x.T; x=a*x+(b*q+c*(q@q))@x\n"," return x.T if t else x\n","class MuonClip(torch.optim.Optimizer):\n"," def __init__(self,params): super().__init__(params,dict(lr=MATRIX_LR,momentum=MOMENTUM,weight_decay=WEIGHT_DECAY))\n"," @torch.no_grad()\n"," def step(self):\n"," for group in self.param_groups:\n"," for p in group['params']:\n"," if p.grad is None: continue\n"," b=self.state[p].setdefault('momentum_buffer',torch.zeros_like(p.grad)); b.mul_(group['momentum']).add_(p.grad)\n"," u=zeropower(b).to(p.dtype); u.mul_(RMS_SCALE*math.sqrt(max(p.shape))); p.mul_(1-group['lr']*group['weight_decay']); p.add_(u,alpha=-group['lr'])\n","\n","rng=np.random.default_rng(123)\n","for shape in [(3,3),(4,3),(3,5)]:\n"," W=rng.normal(size=shape); E=rng.normal(size=shape); err=finite_difference_error(W,E)\n"," print(shape,err); assert err<1e-7\n"]},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["transform=transforms.Compose([transforms.ToTensor(),transforms.Normalize((0.1307,),(0.3081,))])\n","full=datasets.MNIST(str(DATA_DIR),train=True,download=True,transform=transform); test=datasets.MNIST(str(DATA_DIR),train=False,download=True,transform=transform)\n","g=torch.Generator().manual_seed(20_260_807); perm=torch.randperm(len(full),generator=g).tolist(); val_idx,train_idx=perm[:5000],perm[5000:]; train,val=Subset(full,train_idx),Subset(full,val_idx)\n","def loaders(seed):\n"," tg=torch.Generator().manual_seed(seed+101); kw=dict(batch_size=BATCH,num_workers=0,pin_memory=DEVICE.type=='cuda')\n"," return DataLoader(train,shuffle=True,generator=tg,**kw),DataLoader(train,shuffle=False,**kw),DataLoader(val,shuffle=False,**kw),DataLoader(test,shuffle=False,**kw)\n","@torch.inference_mode()\n","def evaluate(model,loader):\n"," model.eval(); loss=correct=n=0\n"," for x,y in loader:\n"," x,y=x.to(DEVICE),y.to(DEVICE); z=model(x); loss+=float(F.cross_entropy(z,y,reduction='sum').cpu()); correct+=int((z.argmax(1)==y).sum().cpu()); n+=y.numel()\n"," return loss/n,correct/n\n"]},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)\n","def numerical_due(epoch): return epoch in {0,EPOCHS} or (NUMERIC_EVERY>0 and epoch%NUMERIC_EVERY==0)\n","def fit_row(seed,epoch,layer,method,spectrum,**extra):\n"," return dict(seed=seed,epoch=epoch,layer=layer,method=method,spectrum_size=len(spectrum),**extra,**powerlaw_mle_grid_ks(spectrum,min_tail=MIN_TAIL))\n","def analyze(model,seed,epoch):\n"," rows=[]; spectra={}\n"," for j,name in enumerate(('fc1','fc2','fc3')):\n"," W=getattr(model,name).weight.detach().cpu().numpy()\n"," raw=raw_weight_esd(W); analytic,zero_modes=analytic_gram_spectrum(W)\n"," rows.append(fit_row(seed,epoch,name,'raw_weight_WtW',raw,zero_modes=np.nan,probes=np.nan,epsilon=np.nan)); spectra[f'raw__seed_{seed}__epoch_{epoch:03d}__{name}']=raw\n"," rows.append(fit_row(seed,epoch,name,'analytic_polar_JtJ',analytic,zero_modes=zero_modes,probes=np.nan,epsilon=np.nan)); spectra[f'analytic__seed_{seed}__epoch_{epoch:03d}__{name}']=analytic\n"," if numerical_due(epoch):\n"," numeric,meta=numerical_probe_gram_spectrum(W,probes=PROBES,eps_rel=EPS,seed=918273+100000*seed+100*epoch+j)\n"," rows.append(fit_row(seed,epoch,name,'numerical_polar_JtJ',numeric,zero_modes=np.nan,probes=PROBES,epsilon=meta['epsilon'])); spectra[f'numeric__seed_{seed}__epoch_{epoch:03d}__{name}']=numeric\n"," return rows,spectra\n","\n","perf=[]; fits=[]; spectra={}\n","for seed in SEEDS:\n"," seed_all(seed); tr,tre,va,te=loaders(seed); model=MLP3().to(DEVICE); matrices=[p for p in model.parameters() if p.ndim==2]; biases=[p for p in model.parameters() if p.ndim!=2]; muon=MuonClip(matrices); aux=torch.optim.AdamW(biases,lr=AUX_LR,weight_decay=WEIGHT_DECAY)\n"," r,s=analyze(model,seed,0); fits+=r; spectra.update(s)\n"," for epoch in range(1,EPOCHS+1):\n"," model.train()\n"," for x,y in tr:\n"," x,y=x.to(DEVICE),y.to(DEVICE); muon.zero_grad(set_to_none=True); aux.zero_grad(set_to_none=True); loss=F.cross_entropy(model(x),y); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),GRAD_CLIP); muon.step(); aux.step()\n"," tl,ta=evaluate(model,tre); vl,vaa=evaluate(model,va); ql,qa=evaluate(model,te); perf.append(dict(seed=seed,epoch=epoch,train_loss=tl,validation_loss=vl,test_loss=ql,train_accuracy=ta,validation_accuracy=vaa,test_accuracy=qa)); r,s=analyze(model,seed,epoch); fits+=r; spectra.update(s); print(seed,epoch,vl)\n","performance=pd.DataFrame(perf); three_way_fits=pd.DataFrame(fits); performance.to_csv(RUN_ROOT/'performance_by_epoch_and_seed.csv',index=False); three_way_fits.to_csv(RUN_ROOT/'raw_vs_polar_mle_ks_by_epoch_layer_seed.csv',index=False); np.savez_compressed(RUN_ROOT/'raw_vs_polar_spectra.npz',**spectra)\n"]},{"cell_type":"markdown","metadata":{},"source":["## Three-way exponent comparison\n","\n","The table below reports \\(\\alpha\\), KS distance \\(D\\), \\(x_{\\min}\\), and fitted-tail size for every available method. The numerical response spectrum is a finite-probe random-subspace estimate; the analytic spectrum is the full exact positive spectrum.\n"]},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["cols=['seed','epoch','layer','method','spectrum_size','alpha','D','xmin','tail_evals','probes']\n","display(three_way_fits[cols].sort_values(['seed','epoch','layer','method']))\n","alpha_table=three_way_fits.pivot_table(index=['seed','epoch','layer'],columns='method',values='alpha',aggfunc='first').reset_index(); alpha_table.to_csv(RUN_ROOT/'alpha_three_way_comparison.csv',index=False); display(alpha_table)\n","final=three_way_fits[three_way_fits.epoch.eq(EPOCHS)].copy(); display(final[cols].sort_values(['seed','layer','method']))\n"]},{"cell_type":"markdown","metadata":{},"source":["## Single-checkpoint numerical convergence\n","\n","For a selected checkpoint matrix `W`, `probe_convergence` repeats the finite-difference calculation at increasing probe counts and compares its MLE/KS exponent with the analytic polar-Gram exponent.\n"]},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["# Example after training:\n","# W=model.fc2.weight.detach().cpu().numpy()\n","# display(pd.DataFrame(probe_convergence(W,counts=(16,32,64,128,256),eps_rel=EPS,min_tail=MIN_TAIL)))\n"]}]} \ No newline at end of file diff --git a/baseline/rg_baselines/polar_jacobian.py b/baseline/rg_baselines/polar_jacobian.py new file mode 100644 index 00000000..c308e6e9 --- /dev/null +++ b/baseline/rg_baselines/polar_jacobian.py @@ -0,0 +1,155 @@ +"""Polar-projection Jacobian spectra for single matrix checkpoints.""" +from __future__ import annotations +import numpy as np + + +def polar_factor(weight: np.ndarray) -> np.ndarray: + w = np.asarray(weight, dtype=np.float64) + u, _, vh = np.linalg.svd(w, full_matrices=False) + return u @ vh + + +def frechet_action(weight: np.ndarray, perturbation: np.ndarray) -> np.ndarray: + """Closed-form D Pi(W)[E] for Pi(W)=U V^T at full-rank W.""" + w = np.asarray(weight, dtype=np.float64) + e = np.asarray(perturbation, dtype=np.float64) + u, s, vh = np.linalg.svd(w, full_matrices=False) + v = vh.T + a = u.T @ e @ v + omega = (a - a.T) / (s[:, None] + s[None, :]) + np.fill_diagonal(omega, 0.0) + out = u @ omega @ v.T + m, n = w.shape + if m > n: + out += (e - u @ (u.T @ e)) @ v @ np.diag(1.0 / s) @ v.T + elif n > m: + out += u @ np.diag(1.0 / s) @ u.T @ e @ (np.eye(n) - v @ v.T) + return out + + +def raw_weight_esd(weight: np.ndarray) -> np.ndarray: + """Raw nonzero eigenvalues of W^T W, i.e. squared singular values.""" + s = np.linalg.svd(np.asarray(weight, dtype=np.float64), compute_uv=False) + return np.sort(s * s) + + +def analytic_gram_spectrum(weight: np.ndarray) -> tuple[np.ndarray, int]: + """Exact positive spectrum of (D Pi)^* D Pi and zero-mode count.""" + w = np.asarray(weight, dtype=np.float64) + m, n = w.shape + s = np.linalg.svd(w, compute_uv=False) + r = min(m, n) + rot = np.fromiter( + (4.0 / (s[i] + s[j]) ** 2 for i in range(r) for j in range(i + 1, r)), + dtype=np.float64, + count=r * (r - 1) // 2, + ) + trans = np.repeat(1.0 / (s * s), abs(m - n)) if m != n else np.empty(0) + positive = np.sort(np.concatenate([rot, trans])) + return positive, int(m * n - positive.size) + + +def numerical_probe_gram_spectrum( + weight: np.ndarray, + *, + probes: int = 128, + eps_rel: float = 1e-5, + seed: int = 918273, +) -> tuple[np.ndarray, dict[str, float]]: + """Finite-difference single-checkpoint response Gram X^T X spectrum.""" + w = np.asarray(weight, dtype=np.float64) + rng = np.random.default_rng(int(seed)) + step = float(eps_rel) * max(1.0, float(np.linalg.norm(w, "fro"))) + x = np.empty((w.size, int(probes)), dtype=np.float64) + for k in range(int(probes)): + e = rng.normal(size=w.shape) + e /= np.linalg.norm(e, "fro") + response = (polar_factor(w + step * e) - polar_factor(w - step * e)) / (2.0 * step) + x[:, k] = response.reshape(-1) + gram = x.T @ x + evals = np.linalg.eigvalsh(gram) + tol = np.finfo(float).eps * max(gram.shape) * max(float(evals[-1]), 1.0) + return np.sort(evals[evals > tol]), { + "probes": int(probes), + "epsilon": float(step), + "epsilon_relative": float(eps_rel), + } + + +def powerlaw_mle_grid_ks(evals: np.ndarray, *, min_tail: int = 5) -> dict[str, float]: + """Continuous Pareto MLE with an x_min grid search minimizing KS distance. + + For every candidate x_min leaving at least ``min_tail`` observations, + alpha = 1 + n / sum(log(x/x_min)). The selected candidate minimizes the + Kolmogorov-Smirnov distance between the empirical tail CDF and the fitted + continuous Pareto CDF. + """ + values = np.sort(np.asarray(evals, dtype=np.float64)) + values = values[np.isfinite(values) & (values > 0)] + if values.size < max(2, int(min_tail)): + return {"alpha": np.nan, "D": np.nan, "xmin": np.nan, "xmax": np.nan, "tail_evals": 0} + + best = None + for start in range(0, values.size - int(min_tail) + 1): + xmin = float(values[start]) + tail = values[start:] + n = int(tail.size) + denom = float(np.sum(np.log(tail / xmin))) + if not np.isfinite(denom) or denom <= 0: + continue + alpha = 1.0 + n / denom + empirical = np.arange(1, n + 1, dtype=np.float64) / n + theoretical = 1.0 - np.power(tail / xmin, 1.0 - alpha) + D = float(np.max(np.abs(empirical - theoretical))) + candidate = (D, start, alpha, xmin, n) + if best is None or candidate[0] < best[0]: + best = candidate + + if best is None: + return {"alpha": np.nan, "D": np.nan, "xmin": np.nan, "xmax": float(np.max(values)), "tail_evals": 0} + D, _, alpha, xmin, n = best + return { + "alpha": float(alpha), + "D": float(D), + "xmin": float(xmin), + "xmax": float(np.max(values)), + "tail_evals": int(n), + } + + +def weightwatcher_pl_fit(evals: np.ndarray) -> dict[str, float]: + """Use WeightWatcher's own PL fitter on raw positive eigenvalues.""" + from weightwatcher.WW_powerlaw import pl_fit + values = np.asarray(evals, dtype=np.float64) + values = values[np.isfinite(values) & (values > 0)] + fit = pl_fit(data=values, xmin=None, xmax=None, verbose=False) + return { + "alpha": float(fit.alpha), + "D": float(fit.D), + "xmin": float(fit.xmin), + "xmax": float(np.max(values)), + "tail_evals": int(np.count_nonzero(values >= fit.xmin)), + } + + +def finite_difference_error(weight: np.ndarray, perturbation: np.ndarray, eps_rel: float = 1e-6) -> float: + w = np.asarray(weight, dtype=np.float64) + e = np.asarray(perturbation, dtype=np.float64) + step = float(eps_rel) * max(1.0, float(np.linalg.norm(w, "fro"))) + fd = (polar_factor(w + step * e) - polar_factor(w - step * e)) / (2.0 * step) + exact = frechet_action(w, e) + return float(np.linalg.norm(fd - exact) / max(np.linalg.norm(fd), np.finfo(float).tiny)) + + +def probe_convergence(weight: np.ndarray, counts=(16, 32, 64, 128, 256), *, eps_rel=1e-5, seed=918273, min_tail=5): + """Return analytic and finite-probe MLE/KS fits for one checkpoint.""" + analytic, _ = analytic_gram_spectrum(weight) + reference = powerlaw_mle_grid_ks(analytic, min_tail=min_tail) + rows = [{"method": "analytic", "probes": np.nan, **reference}] + for count in counts: + sampled, _ = numerical_probe_gram_spectrum(weight, probes=int(count), eps_rel=eps_rel, seed=seed) + rows.append({"method": "numerical_finite_difference", "probes": int(count), **powerlaw_mle_grid_ks(sampled, min_tail=min_tail)}) + for row in rows: + row["alpha_analytic"] = reference["alpha"] + row["alpha_difference"] = row["alpha"] - reference["alpha"] + return rows