Buckets:
| """Claim 1: STE dynamics converge to a deterministic ODE in the high-d limit. | |
| Reproduces Theorem V.3 of Ichikawa et al. (arXiv:2510.10693) which states: | |
| E ||Psi_t - Psi(t/d)||^2 <= C / d. | |
| We run STE simulations at increasing dimensions d in {100, 300, 900, 3000} under | |
| the Figure 2 setting (joint weight+input quantization, b=2, omega=1, eta=0.04, | |
| lambda=1, T=0, rho=1, sigma2=0), and compare the macroscopic trajectory | |
| (m(tau), q(tau), eps_g(tau)) to the ODE prediction. The L2 distance between the | |
| STE mean (over 5 seeds) and the ODE solution is expected to shrink like | |
| O(d^{-1/2}) (the theorem's L1 bound is C/sqrt(d)). | |
| """ | |
| import json | |
| import math | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| from scipy.integrate import solve_ivp | |
| sys.path.insert(0, str(Path(__file__).parent)) | |
| from ste_repro import Quantizer, ode_rhs, eps_g, macro_m_psi, macro_q_psi | |
| from ste_sim_torch import run_ste, STEConfig | |
| OUT = Path(__file__).resolve().parents[1] / "outputs" / "claim1" | |
| OUT.mkdir(parents=True, exist_ok=True) | |
| def run_claim1(): | |
| # Figure 2 setting: weight-only quantization (input unquantized: | |
| # kappa_x = sigma2_x = 1), weight quantizer b=2, omega=1. | |
| b = 2 | |
| omega = 1.0 | |
| eta = 0.04 | |
| lam = 1.0 | |
| rho = 1.0 | |
| sigma2 = 0.0 | |
| qw = Quantizer(b=b, omega=omega) | |
| # input is unquantized -> identity quantizer moments | |
| kappa_x, sigma2_x = 1.0, 1.0 | |
| print(f"[Claim 1] weight quant: {qw.describe()}") | |
| print(f"[Claim 1] input unquantized: kappa_x=1.0 sigma2_x=1.0") | |
| # Initial condition: w ~ N(0,1), so m0 = 0, q0 = 1 | |
| m0, q0 = 0.0, 1.0 | |
| # Integrate ODE for tau up to tau_max (matches Figure 2 x-axis) | |
| tau_max = 300.0 | |
| n_eval = 600 | |
| t_eval = np.linspace(0.0, tau_max, n_eval) | |
| sol = solve_ivp( | |
| ode_rhs, [0.0, tau_max], [m0, q0], | |
| args=(qw, kappa_x, sigma2_x, eta, lam, rho, sigma2), | |
| t_eval=t_eval, rtol=1e-9, atol=1e-12, method="DOP853", | |
| ) | |
| m_ode = sol.y[0] | |
| q_ode = sol.y[1] | |
| # eps_g from ODE | |
| eps_ode = np.empty_like(m_ode) | |
| for i, (m, q) in enumerate(zip(m_ode, q_ode)): | |
| s = math.sqrt(max(q - m * m / rho, 1e-12)) | |
| eps_ode[i] = eps_g(m, s, qw, kappa_x, sigma2_x, rho, sigma2) | |
| print(f"[Claim 1] ODE integrated. Final eps_g(ODE, tau={tau_max}) = {eps_ode[-1]:.6f}") | |
| # STE simulations at increasing d. We use weight-only quantization (matches | |
| # Figure 2 setting: input unquantized -> kappa_x = sigma2_x = 1, weight | |
| # quantizer b=2, omega=1). tau_max = 300 keeps wall-clock modest while | |
| # spanning many relaxation times (eta * tau_max = 12). | |
| dims = [100, 300, 900, 3000] | |
| n_seeds = 10 | |
| tau_max = 300.0 | |
| log_period = 50 | |
| results = {} | |
| for d in dims: | |
| n_steps = int(d * tau_max) | |
| # cap log buffer to ~ n_eval points | |
| lp = max(log_period, n_steps // (n_eval * 2)) | |
| cfg = STEConfig( | |
| d=d, eta=eta, lam=lam, rho=rho, sigma2=sigma2, | |
| n_steps=n_steps, log_period=lp, n_seeds=n_seeds, | |
| w_init_std=1.0, w_star_value=1.0, | |
| sigma2_x=sigma2_x, kappa_x=kappa_x, device="cuda", | |
| ) | |
| print(f"[Claim 1] STE: d={d}, n_steps={n_steps}, log_period={lp}") | |
| t0 = time.time() | |
| taus, metrics, wall = run_ste( | |
| (qw.levels, qw.theta, qw.omega, qw.Delta), | |
| None, # identity input (weight-only quantization) | |
| cfg, | |
| ) | |
| print(f" wall={wall:.1f}s (incl. {time.time()-t0:.1f}s total)") | |
| # mean over seeds | |
| m_ste = metrics["m"].mean(axis=1) | |
| q_ste = metrics["q"].mean(axis=1) | |
| eps_ste = metrics["eps_g"].mean(axis=1) | |
| # interpolate ODE at STE taus | |
| m_ode_at = np.interp(taus, t_eval, m_ode) | |
| q_ode_at = np.interp(taus, t_eval, q_ode) | |
| eps_ode_at = np.interp(taus, t_eval, eps_ode) | |
| # L2 (RMS) error -- full trajectory | |
| rms_m = float(np.sqrt(np.mean((m_ste - m_ode_at) ** 2))) | |
| rms_q = float(np.sqrt(np.mean((q_ste - q_ode_at) ** 2))) | |
| rms_eps = float(np.sqrt(np.mean((eps_ste - eps_ode_at) ** 2))) | |
| # L2 (RMS) error -- late trajectory (tau >= 20, after early transient) | |
| late_mask = taus >= 20.0 | |
| rms_m_late = float(np.sqrt(np.mean((m_ste[late_mask] - m_ode_at[late_mask]) ** 2))) | |
| rms_q_late = float(np.sqrt(np.mean((q_ste[late_mask] - q_ode_at[late_mask]) ** 2))) | |
| rms_eps_late = float(np.sqrt(np.mean((eps_ste[late_mask] - eps_ode_at[late_mask]) ** 2))) | |
| results[d] = { | |
| "wall_s": wall, | |
| "rms_m": rms_m, "rms_q": rms_q, "rms_eps": rms_eps, | |
| "rms_m_late": rms_m_late, "rms_q_late": rms_q_late, "rms_eps_late": rms_eps_late, | |
| "m_final_ste": float(m_ste[-1]), "m_final_ode": float(m_ode_at[-1]), | |
| "q_final_ste": float(q_ste[-1]), "q_final_ode": float(q_ode_at[-1]), | |
| "eps_final_ste": float(eps_ste[-1]), "eps_final_ode": float(eps_ode_at[-1]), | |
| "taus": taus.tolist(), | |
| "m_ste_mean": m_ste.tolist(), "q_ste_mean": q_ste.tolist(), | |
| "eps_ste_mean": eps_ste.tolist(), | |
| "m_ste_std": metrics["m"].std(axis=1).tolist(), | |
| "q_ste_std": metrics["q"].std(axis=1).tolist(), | |
| "eps_ste_std": metrics["eps_g"].std(axis=1).tolist(), | |
| } | |
| print(f" d={d}: RMS(full) m={rms_m:.5f} q={rms_q:.5f} eps={rms_eps:.5f}") | |
| print(f" d={d}: RMS(late) m={rms_m_late:.5f} q={rms_q_late:.5f} eps={rms_eps_late:.5f}") | |
| print(f" final eps: STE={eps_ste[-1]:.6f} ODE={eps_ode_at[-1]:.6f}") | |
| # save ODE trajectory too | |
| results["_ode"] = { | |
| "t_eval": t_eval.tolist(), | |
| "m_ode": m_ode.tolist(), | |
| "q_ode": q_ode.tolist(), | |
| "eps_ode": eps_ode.tolist(), | |
| } | |
| results["_meta"] = { | |
| "b": b, "omega": omega, "eta": eta, "lam": lam, "rho": rho, "sigma2": sigma2, | |
| "kappa_x": kappa_x, "sigma2_x": sigma2_x, | |
| "n_seeds": n_seeds, "tau_max": tau_max, | |
| "dims": dims, | |
| } | |
| # ---- Fit scaling: rms(d) ~ A / sqrt(d) -> log-log slope ~ -1/2 ---- | |
| ds = np.array(dims, dtype=float) | |
| rms_eps_arr = np.array([results[d]["rms_eps"] for d in dims]) | |
| rms_eps_late_arr = np.array([results[d]["rms_eps_late"] for d in dims]) | |
| # linear fit log(rms) = log(A) + slope * log(d) | |
| A_full = np.polyfit(np.log(ds), np.log(rms_eps_arr), 1) | |
| A_late = np.polyfit(np.log(ds), np.log(rms_eps_late_arr), 1) | |
| slope_full = float(A_full[0]) | |
| slope_late = float(A_late[0]) | |
| results["_scaling"] = { | |
| "dims": ds.tolist(), | |
| "rms_eps": rms_eps_arr.tolist(), | |
| "rms_eps_late": rms_eps_late_arr.tolist(), | |
| "loglog_slope_full": slope_full, | |
| "loglog_slope_late": slope_late, | |
| "expected_slope": -0.5, | |
| } | |
| print(f"[Claim 1] log-log slope (full): {slope_full:.4f} (expected ~ -0.5)") | |
| print(f"[Claim 1] log-log slope (late): {slope_late:.4f} (expected ~ -0.5)") | |
| with open(OUT / "claim1_results.json", "w") as f: | |
| json.dump(results, f, indent=2) | |
| # ---- Plot: convergence (Figure like Fig 1 + scaling inset) ---- | |
| fig, axes = plt.subplots(1, 3, figsize=(13, 4)) | |
| ax = axes[0] | |
| ax.plot(t_eval, m_ode, "b-", label="ODE") | |
| for d in dims: | |
| r = results[d] | |
| ax.plot(r["taus"], r["m_ste_mean"], label=f"STE d={d}", alpha=0.7) | |
| ax.set_xlabel(r"$\tau$") | |
| ax.set_ylabel(r"$m(\tau)$") | |
| ax.set_title("m(τ): STE → ODE as d grows") | |
| ax.legend(fontsize=8) | |
| ax.grid(True, alpha=0.3) | |
| ax = axes[1] | |
| ax.plot(t_eval, q_ode, "b-", label="ODE") | |
| for d in dims: | |
| r = results[d] | |
| ax.plot(r["taus"], r["q_ste_mean"], label=f"STE d={d}", alpha=0.7) | |
| ax.set_xlabel(r"$\tau$") | |
| ax.set_ylabel(r"$q(\tau)$") | |
| ax.set_title("q(τ): STE → ODE as d grows") | |
| ax.legend(fontsize=8) | |
| ax.grid(True, alpha=0.3) | |
| ax = axes[2] | |
| ax.plot(t_eval, eps_ode, "b-", label="ODE", lw=2) | |
| for d in dims: | |
| r = results[d] | |
| ax.plot(r["taus"], r["eps_ste_mean"], label=f"STE d={d}", alpha=0.7) | |
| ax.set_xlabel(r"$\tau$") | |
| ax.set_ylabel(r"$\varepsilon_g(\tau)$") | |
| ax.set_title("ε_g(τ): STE → ODE as d grows") | |
| ax.legend(fontsize=8) | |
| ax.grid(True, alpha=0.3) | |
| fig.tight_layout() | |
| fig.savefig(OUT / "claim1_trajectories.png", dpi=120) | |
| plt.close(fig) | |
| # ---- Plot: scaling inset (RMS vs d, log-log) ---- | |
| fig, ax = plt.subplots(figsize=(5, 4)) | |
| ax.loglog(ds, rms_eps_arr, "bo-", label="RMS(ε_g) full trajectory") | |
| ax.loglog(ds, rms_eps_late_arr, "go-", label="RMS(ε_g) late (τ≥20)") | |
| # reference 1/sqrt(d) line scaled | |
| ref = rms_eps_late_arr[0] * np.sqrt(ds[0]) / np.sqrt(ds) | |
| ax.loglog(ds, ref, "k--", label=r"$\propto d^{-1/2}$ (Thm V.3 upper bound)") | |
| ax.set_xlabel("dimension d") | |
| ax.set_ylabel(r"RMS$(\varepsilon_g^{STE} - \varepsilon_g^{ODE})$") | |
| ax.set_title(f"Claim 1: ODE convergence rate\nlate-traj log-log slope = {slope_late:.3f} (expected -0.5)") | |
| ax.legend(fontsize=9) | |
| ax.grid(True, alpha=0.3, which="both") | |
| fig.tight_layout() | |
| fig.savefig(OUT / "claim1_scaling.png", dpi=120) | |
| plt.close(fig) | |
| print(f"[Claim 1] Done. Outputs in {OUT}") | |
| return results | |
| if __name__ == "__main__": | |
| run_claim1() | |
Xet Storage Details
- Size:
- 9.31 kB
- Xet hash:
- 8f280423a675f9801f8c157b98fbd6cbbf859a91747858e94c165f4a1df49bac
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.