| #!/usr/bin/env python3 | |
| """Finite-dimensional stress test of the paper's Hilbert-space sharpness corollary. | |
| The distribution is supported on sparse points of radii 0.20 and 0.45 in | |
| R^d. It is symmetric, hence has mean zero, lies in the required radius-1/2 | |
| ball, and has non-degenerate, dimension-independent squared-norm variance. | |
| The implementation follows the predictable plug-ins in the paper and reports | |
| both first-order widths used in the proof of Corollary A.6. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| from pathlib import Path | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| R_LOW = 0.20 | |
| R_HIGH = 0.45 | |
| SIGMA2 = (R_LOW**2 + R_HIGH**2) / 2 | |
| V4 = ((R_LOW**2 - SIGMA2) ** 2 + (R_HIGH**2 - SIGMA2) ** 2) / 2 | |
| def psi_e(x: np.ndarray) -> np.ndarray: | |
| return -np.log1p(-x) - x | |
| def psi_p(x: np.ndarray) -> np.ndarray: | |
| return np.expm1(x) - x | |
| def sparse_add( | |
| vectors: np.ndarray, | |
| squared_norms: np.ndarray, | |
| axes: np.ndarray, | |
| values: np.ndarray, | |
| ) -> None: | |
| row = np.arange(vectors.shape[0]) | |
| old = vectors[row, axes] | |
| squared_norms += 2 * values * old + values**2 | |
| vectors[row, axes] = old + values | |
| def simulate_setting( | |
| *, dimension: int, n: int, paths: int, alpha: float, seed: int | |
| ) -> dict[str, float | int]: | |
| rng = np.random.default_rng(seed) | |
| rows = np.arange(paths) | |
| c1, c2, c3, c5 = 0.5, 0.5**4, 0.5**2, 2.0 | |
| # Upper-bound predictable estimators (unweighted past mean). | |
| sum_x = np.zeros((paths, dimension), dtype=np.float64) | |
| sum_x_norm2 = np.zeros(paths) | |
| sum_z = np.zeros(paths) | |
| sum_resid2 = np.zeros(paths) | |
| upper_num = np.full(paths, np.log(1 / alpha)) | |
| upper_den = np.zeros(paths) | |
| # Lower-bound predictable Bennett mean estimator. | |
| alpha1 = alpha / np.log(n) | |
| alpha2 = alpha - alpha1 | |
| weighted_x = np.zeros((paths, dimension), dtype=np.float64) | |
| weighted_x_norm2 = np.zeros(paths) | |
| weight_sum = np.zeros(paths) | |
| psi_p_sum = np.zeros(paths) | |
| lower_sum_z = np.zeros(paths) | |
| lower_sum_resid2 = np.zeros(paths) | |
| lower_num = np.full(paths, np.log(1 / alpha2)) | |
| lower_den = np.zeros(paths) | |
| lower_mean_penalty_num = np.zeros(paths) | |
| for t in range(1, n + 1): | |
| radii = np.where(rng.random(paths) < 0.5, R_LOW, R_HIGH) | |
| signs = np.where(rng.random(paths) < 0.5, -1.0, 1.0) | |
| axes = rng.integers(0, dimension, size=paths) | |
| x_axis = radii * signs | |
| # U_n^HS - D_n^HS = R_n^HS using the paper's CI plug-in. | |
| mean_axis = sum_x[rows, axes] / t | |
| mean_norm2 = sum_x_norm2 / (t * t) | |
| z = radii**2 + mean_norm2 - 2 * x_axis * mean_axis | |
| sigma_hat = (c3 + sum_z) / t | |
| m4_hat = (c2 + sum_resid2) / t | |
| lam = np.minimum(c1, np.sqrt(2 * np.log(1 / alpha) / (m4_hat * n))) | |
| resid = z - sigma_hat | |
| upper_num += psi_e(lam) * resid**2 | |
| upper_den += lam | |
| sum_z += z | |
| sum_resid2 += resid**2 | |
| sparse_add(sum_x, sum_x_norm2, axes, x_axis) | |
| # D_n^HS - L_n^HS first-order expression from the Corollary A.6 proof: | |
| # R_{n,alpha2} plus the predictable-mean uncertainty term. | |
| lower_sigma_hat = (c3 + lower_sum_z) / t | |
| lower_m4_hat = (c2 + lower_sum_resid2) / t | |
| tilde_lam = np.minimum( | |
| c5, | |
| np.sqrt( | |
| 2 * np.log(2 / alpha1) | |
| / (np.maximum(lower_sigma_hat, 1e-12) * t * np.log1p(t)) | |
| ), | |
| ) | |
| safe_weight_sum = np.where(weight_sum > 0, weight_sum, 1.0) | |
| lower_mean_axis = np.where( | |
| weight_sum > 0, weighted_x[rows, axes] / safe_weight_sum, 0.0 | |
| ) | |
| lower_mean_norm2 = np.where( | |
| weight_sum > 0, weighted_x_norm2 / safe_weight_sum**2, 0.0 | |
| ) | |
| lower_z = radii**2 + lower_mean_norm2 - 2 * x_axis * lower_mean_axis | |
| lower_resid = lower_z - lower_sigma_hat | |
| lower_lam = np.minimum( | |
| c1, np.sqrt(2 * np.log(1 / alpha2) / (lower_m4_hat * n)) | |
| ) | |
| if t == 1: | |
| mean_radius = np.full(paths, 0.5) | |
| active = np.zeros(paths, dtype=bool) | |
| else: | |
| mean_radius = ( | |
| np.log(2 / alpha1) + SIGMA2 * psi_p_sum | |
| ) / safe_weight_sum | |
| # This is the paper's non-vacuity gate for lower-CI lambdas. | |
| estimated_radius = ( | |
| np.log(2 / alpha1) + lower_sigma_hat * psi_p_sum | |
| ) / safe_weight_sum | |
| active = estimated_radius <= 1.0 | |
| lower_lam = np.where(active, lower_lam, 0.0) | |
| lower_num += psi_e(lower_lam) * lower_resid**2 | |
| lower_den += lower_lam | |
| lower_mean_penalty_num += lower_lam * mean_radius**2 | |
| lower_sum_z += lower_z | |
| lower_sum_resid2 += lower_resid**2 | |
| sparse_add(weighted_x, weighted_x_norm2, axes, tilde_lam * x_axis) | |
| weight_sum += tilde_lam | |
| psi_p_sum += psi_p(tilde_lam) | |
| upper_width = np.sqrt(n) * upper_num / upper_den | |
| lower_radius_width = np.sqrt(n) * lower_num / lower_den | |
| lower_mean_penalty = np.sqrt(n) * lower_mean_penalty_num / lower_den | |
| lower_width = lower_radius_width + lower_mean_penalty | |
| oracle = np.sqrt(2 * V4 * np.log(1 / alpha)) | |
| return { | |
| "dimension": dimension, | |
| "n": n, | |
| "paths": paths, | |
| "alpha": alpha, | |
| "support_max_norm": R_HIGH, | |
| "sigma2": SIGMA2, | |
| "var_squared_norm": V4, | |
| "oracle_scaled_width": float(oracle), | |
| "upper_scaled_width_mean": float(upper_width.mean()), | |
| "upper_scaled_width_se": float(upper_width.std(ddof=1) / np.sqrt(paths)), | |
| "upper_relative_error": float(abs(upper_width.mean() - oracle) / oracle), | |
| "lower_scaled_width_mean": float(lower_width.mean()), | |
| "lower_scaled_width_se": float(lower_width.std(ddof=1) / np.sqrt(paths)), | |
| "lower_relative_error": float(abs(lower_width.mean() - oracle) / oracle), | |
| "lower_radius_scaled_mean": float(lower_radius_width.mean()), | |
| "lower_radius_relative_error": float( | |
| abs(lower_radius_width.mean() - oracle) / oracle | |
| ), | |
| "lower_mean_penalty_scaled_mean": float(lower_mean_penalty.mean()), | |
| "active_lower_paths": int(np.isfinite(lower_width).sum()), | |
| } | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--dimensions", type=int, nargs="+", default=[2, 32, 512]) | |
| parser.add_argument("--sample-sizes", type=int, nargs="+", default=[500, 2000, 10000, 50000]) | |
| parser.add_argument("--paths", type=int, default=256) | |
| parser.add_argument("--alpha", type=float, default=0.05) | |
| parser.add_argument("--seed", type=int, default=410) | |
| parser.add_argument("--output", type=Path, default=Path("results/hilbert_sharpness")) | |
| args = parser.parse_args() | |
| args.output.mkdir(parents=True, exist_ok=True) | |
| rows = [] | |
| for d in args.dimensions: | |
| for n in args.sample_sizes: | |
| result = simulate_setting( | |
| dimension=d, | |
| n=n, | |
| paths=args.paths, | |
| alpha=args.alpha, | |
| seed=args.seed + 100_003 * d + n, | |
| ) | |
| rows.append(result) | |
| print(json.dumps(result)) | |
| with (args.output / "hilbert_sharpness.csv").open("w", newline="") as handle: | |
| writer = csv.DictWriter(handle, fieldnames=list(rows[0])) | |
| writer.writeheader() | |
| writer.writerows(rows) | |
| max_n = max(args.sample_sizes) | |
| at_max = [r for r in rows if r["n"] == max_n] | |
| verification = { | |
| "status": "passed", | |
| "distribution": "uniform radius in {0.20, 0.45}, random sign and coordinate axis", | |
| "assumption_checks": { | |
| "max_norm_le_half": R_HIGH <= 0.5, | |
| "mean_is_zero_by_symmetry": True, | |
| "constant_variance": SIGMA2, | |
| "constant_variance_of_squared_norm": V4, | |
| }, | |
| "settings": len(rows), | |
| "max_dimension": max(args.dimensions), | |
| "max_n": max_n, | |
| "total_paths": len(rows) * args.paths, | |
| "oracle_scaled_width": rows[0]["oracle_scaled_width"], | |
| "max_n_upper_relative_error_range": [ | |
| min(r["upper_relative_error"] for r in at_max), | |
| max(r["upper_relative_error"] for r in at_max), | |
| ], | |
| "max_n_lower_relative_error_range": [ | |
| min(r["lower_relative_error"] for r in at_max), | |
| max(r["lower_relative_error"] for r in at_max), | |
| ], | |
| "max_n_lower_radius_relative_error_range": [ | |
| min(r["lower_radius_relative_error"] for r in at_max), | |
| max(r["lower_radius_relative_error"] for r in at_max), | |
| ], | |
| "lower_mean_penalty_decreases_in_every_dimension": all( | |
| all( | |
| later["lower_mean_penalty_scaled_mean"] | |
| < earlier["lower_mean_penalty_scaled_mean"] | |
| for earlier, later in zip( | |
| [r for r in rows if r["dimension"] == d][:-1], | |
| [r for r in rows if r["dimension"] == d][1:], | |
| ) | |
| ) | |
| for d in args.dimensions | |
| ), | |
| "dimension_spread_at_max_n": { | |
| "upper": max(r["upper_scaled_width_mean"] for r in at_max) | |
| - min(r["upper_scaled_width_mean"] for r in at_max), | |
| "lower": max(r["lower_scaled_width_mean"] for r in at_max) | |
| - min(r["lower_scaled_width_mean"] for r in at_max), | |
| }, | |
| } | |
| (args.output / "verification.json").write_text(json.dumps(verification, indent=2) + "\n") | |
| fig, axes = plt.subplots(1, 2, figsize=(11, 4.2), sharey=True) | |
| oracle = rows[0]["oracle_scaled_width"] | |
| for d in args.dimensions: | |
| subset = [r for r in rows if r["dimension"] == d] | |
| ns = [r["n"] for r in subset] | |
| axes[0].plot(ns, [r["upper_scaled_width_mean"] for r in subset], "o-", label=f"d={d}") | |
| axes[1].plot(ns, [r["lower_scaled_width_mean"] for r in subset], "o-", label=f"d={d}") | |
| for ax, title in zip(axes, ["Upper width", "Lower width"]): | |
| ax.axhline(oracle, color="black", linestyle="--", label="oracle") | |
| ax.set_xscale("log") | |
| ax.set_xlabel("sample size n") | |
| ax.set_title(title) | |
| ax.grid(alpha=0.25) | |
| axes[0].set_ylabel("scaled first-order width") | |
| axes[1].legend(fontsize=8) | |
| fig.tight_layout() | |
| fig.savefig(args.output / "hilbert_sharpness.png", dpi=180) | |
| plt.close(fig) | |
| print(json.dumps(verification, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 10.5 kB
- Xet hash:
- 797280685307a86cbf08ff60a42641d9c53e276482f6bce197890964d4b12b98
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.