Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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 i<j,\n","\\]\n","and for rectangular matrices\n","\\[\n","\\lambda_i^\\perp=\\frac{1}{\\sigma_i^2}\n","\\]\n","with multiplicity \\(|m-n|\\).\n","\n","For the numerical method, draw isotropic unit-Frobenius directions \\(E_a\\), compute\n","\\[\n","X_a=\\frac{\\Pi(W+\\epsilon E_a)-\\Pi(W-\\epsilon E_a)}{2\\epsilon},\n","\\]\n","stack \\(\\operatorname{vec}(X_a)\\) into \\(X\\), and diagonalize \\(X^\\top X\\).\n"]},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["from pathlib import Path\n","import os,sys,math,random\n","import numpy as np,pandas as pd,torch\n","import torch.nn.functional as F\n","from torch.utils.data import DataLoader,Subset\n","from torchvision import datasets,transforms\n","from IPython.display import display\n","import matplotlib.pyplot as plt\n","\n","ROOT=None\n","for q in [Path.cwd(),*Path.cwd().parents]:\n"," if (q/'baseline'/'rg_baselines').is_dir(): ROOT=(q/'baseline').resolve(); break\n"," if (q/'rg_baselines').is_dir(): ROOT=q.resolve(); break\n","if ROOT is None: raise RuntimeError('Run from CalculatedContent/rg_optimizers')\n","sys.path.insert(0,str(ROOT))\n","from rg_baselines import MLP3,DEFAULT_BASELINE_SEEDS,MNIST_REFERENCE_SUITE_SLUG\n","from rg_baselines.polar_jacobian import raw_weight_esd,analytic_gram_spectrum,numerical_probe_gram_spectrum,powerlaw_mle_grid_ks,finite_difference_error,probe_convergence\n","\n","DEVICE=torch.device('cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu')\n","DATA_DIR=Path(os.environ.get('RG_BASELINE_DATA_DIR',ROOT/'data')).expanduser().resolve()\n","RUN_ROOT=Path(os.environ.get('RG_BASELINE_RUN_ROOT',ROOT/'runs')).expanduser().resolve()/MNIST_REFERENCE_SUITE_SLUG/'muonclip_polar_jacobian_three_way_mle_ks'; RUN_ROOT.mkdir(parents=True,exist_ok=True)\n","EPOCHS=int(os.environ.get('MUONCLIP_POLAR_EPOCHS','30')); BATCH=int(os.environ.get('MUONCLIP_POLAR_BATCH','256'))\n","SEEDS=tuple(int(x) for x in os.environ.get('MUONCLIP_POLAR_SEEDS',','.join(map(str,DEFAULT_BASELINE_SEEDS))).split(','))\n","PROBES=int(os.environ.get('POLAR_NUMERIC_PROBES','32')); EPS=float(os.environ.get('POLAR_NUMERIC_EPS_REL','1e-5')); NUMERIC_EVERY=int(os.environ.get('POLAR_NUMERIC_EVERY','0'))\n","MIN_TAIL=int(os.environ.get('POLAR_MLE_KS_MIN_TAIL','5'))\n","print(DEVICE,RUN_ROOT)\n"]},{"cell_type":"markdown","metadata":{},"source":["## Power-law estimator\n","\n","For each candidate tail start \\(x_{\\min}\\) leaving at least `MIN_TAIL` eigenvalues, fit the continuous Pareto exponent\n","\\[\n","\\widehat\\alpha=1+\\frac{n}{\\sum_i\\log(x_i/x_{\\min})}.\n","\\]\n","Select the candidate minimizing\n","\\[\n","D=\\sup_x |F_{\\rm empirical}(x)-F_{\\rm Pareto}(x)|.\n","\\]\n","The same `powerlaw_mle_grid_ks` function is applied to the raw weight ESD, analytic polar-Gram ESD, and numerical probe-Gram ESD.\n"]},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":["MATRIX_LR=2e-3; AUX_LR=2e-3; MOMENTUM=.95; WEIGHT_DECAY=1e-2; RMS_SCALE=.20; GRAD_CLIP=1.0\n","@torch.no_grad()\n","def zeropower(g,steps=5,eps=1e-7):\n"," t=g.shape[0]>g.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"]}]}
155 changes: 155 additions & 0 deletions baseline/rg_baselines/polar_jacobian.py
Original file line number Diff line number Diff line change
@@ -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
Loading