amkkk's picture
download
raw
9.31 kB
"""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.