SabaPivot's picture
download
raw
10.5 kB
#!/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.