summaryrefslogtreecommitdiff
path: root/scripts/minimax_initialization.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-05-28 22:54:53 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-05-28 22:54:53 -0500
commit48f43ffa66c8af1b5aec123ebe9877bc3852ecd3 (patch)
tree605f9cbc0fe5e7e147dd81770f9a3282c7a0af19 /scripts/minimax_initialization.py
parent590baa92c66caeb939bae3fad0ed7cc9b2c87f75 (diff)
Add prior-free minimax initialization simulation
Diffstat (limited to 'scripts/minimax_initialization.py')
-rwxr-xr-xscripts/minimax_initialization.py364
1 files changed, 364 insertions, 0 deletions
diff --git a/scripts/minimax_initialization.py b/scripts/minimax_initialization.py
new file mode 100755
index 0000000..ba73b2e
--- /dev/null
+++ b/scripts/minimax_initialization.py
@@ -0,0 +1,364 @@
+#!/usr/bin/env python3
+"""Empirically illustrate the prior-free minimax feedback bound.
+
+For normalized feedback directions b_hat in S^{D-1}, any initialization
+distribution mu induces M_mu = E[b_hat b_hat^T] with trace 1. For any target
+direction a, the expected squared alignment is a^T M_mu a, so the worst-case
+target receives lambda_min(M_mu) <= 1/D. Isotropic feedback attains 1/D.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from dataclasses import asdict, dataclass
+from pathlib import Path
+
+import matplotlib.pyplot as plt
+import numpy as np
+
+
+@dataclass(frozen=True)
+class RunConfig:
+ dimension: int
+ feedback_samples: int
+ target_samples: int
+ batch_size: int
+ seed: int
+ distributions: list[str]
+ subspace_dim: int
+ anisotropy: float
+ outdir: str
+ plot: bool
+
+
+@dataclass(frozen=True)
+class DistributionSummary:
+ name: str
+ trace: float
+ lambda_min: float
+ lambda_max: float
+ minimax_bound: float
+ isotropic_gap: float
+ random_target_mean: float
+ random_target_std: float
+ random_target_p01: float
+ random_target_p50: float
+ random_target_p99: float
+ min_eigen_target_alignment: float
+ max_eigen_target_alignment: float
+ frobenius_distance_to_isotropic: float
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(
+ description="Simulate the prior-free minimax bound for feedback initialization."
+ )
+ parser.add_argument(
+ "--dimension",
+ type=int,
+ default=64,
+ help="Flattened feedback dimension D.",
+ )
+ parser.add_argument(
+ "--feedback-samples",
+ type=int,
+ default=50_000,
+ help="Number of feedback directions used to estimate M_mu.",
+ )
+ parser.add_argument(
+ "--target-samples",
+ type=int,
+ default=20_000,
+ help="Number of random target directions for evaluating a^T M_mu a.",
+ )
+ parser.add_argument(
+ "--batch-size",
+ type=int,
+ default=10_000,
+ help="Sampling batch size.",
+ )
+ parser.add_argument("--seed", type=int, default=0, help="Random seed.")
+ parser.add_argument(
+ "--distribution",
+ nargs="+",
+ default=["isotropic", "rademacher", "anisotropic", "subspace", "axis"],
+ choices=["isotropic", "rademacher", "anisotropic", "subspace", "axis"],
+ help="Feedback initialization distributions to compare.",
+ )
+ parser.add_argument(
+ "--subspace-dim",
+ type=int,
+ default=8,
+ help="Active dimension for the subspace distribution.",
+ )
+ parser.add_argument(
+ "--anisotropy",
+ type=float,
+ default=16.0,
+ help="Variance ratio for the anisotropic Gaussian distribution.",
+ )
+ parser.add_argument(
+ "--outdir",
+ type=Path,
+ default=Path("outputs/minimax_initialization"),
+ help="Directory for summaries and plots.",
+ )
+ parser.add_argument("--plot", action="store_true", help="Save diagnostic plots.")
+ return parser.parse_args()
+
+
+def normalize(values: np.ndarray) -> np.ndarray:
+ norms = np.linalg.norm(values, axis=1, keepdims=True)
+ return values / np.maximum(norms, np.finfo(values.dtype).tiny)
+
+
+def sample_feedback(
+ rng: np.random.Generator,
+ distribution: str,
+ count: int,
+ dimension: int,
+ subspace_dim: int,
+ anisotropy: float,
+) -> np.ndarray:
+ if distribution == "isotropic":
+ return normalize(rng.standard_normal((count, dimension)))
+
+ if distribution == "rademacher":
+ values = rng.choice(np.array([-1.0, 1.0]), size=(count, dimension))
+ return values / np.sqrt(dimension)
+
+ if distribution == "axis":
+ values = np.zeros((count, dimension), dtype=np.float64)
+ values[:, 0] = rng.choice(np.array([-1.0, 1.0]), size=count)
+ return values
+
+ if distribution == "subspace":
+ active = min(subspace_dim, dimension)
+ values = np.zeros((count, dimension), dtype=np.float64)
+ values[:, :active] = rng.standard_normal((count, active))
+ return normalize(values)
+
+ if distribution == "anisotropic":
+ variances = np.geomspace(anisotropy, 1.0, num=dimension)
+ values = rng.standard_normal((count, dimension)) * np.sqrt(variances)
+ return normalize(values)
+
+ raise ValueError(f"Unknown distribution: {distribution}")
+
+
+def estimate_second_moment(
+ rng: np.random.Generator, config: RunConfig, distribution: str
+) -> np.ndarray:
+ moment = np.zeros((config.dimension, config.dimension), dtype=np.float64)
+ remaining = config.feedback_samples
+ while remaining > 0:
+ count = min(config.batch_size, remaining)
+ feedback = sample_feedback(
+ rng,
+ distribution,
+ count,
+ config.dimension,
+ config.subspace_dim,
+ config.anisotropy,
+ )
+ moment += feedback.T @ feedback
+ remaining -= count
+ return moment / config.feedback_samples
+
+
+def random_targets(
+ rng: np.random.Generator, count: int, dimension: int
+) -> np.ndarray:
+ return normalize(rng.standard_normal((count, dimension)))
+
+
+def summarize_distribution(
+ name: str,
+ moment: np.ndarray,
+ rng: np.random.Generator,
+ target_samples: int,
+) -> tuple[DistributionSummary, np.ndarray, np.ndarray]:
+ dimension = moment.shape[0]
+ eigenvalues = np.linalg.eigvalsh(moment)
+ targets = random_targets(rng, target_samples, dimension)
+ target_alignments = np.einsum("ij,jk,ik->i", targets, moment, targets)
+ minimax_bound = 1.0 / dimension
+ isotropic = np.eye(dimension) / dimension
+
+ summary = DistributionSummary(
+ name=name,
+ trace=float(np.trace(moment)),
+ lambda_min=float(eigenvalues[0]),
+ lambda_max=float(eigenvalues[-1]),
+ minimax_bound=minimax_bound,
+ isotropic_gap=float(minimax_bound - eigenvalues[0]),
+ random_target_mean=float(np.mean(target_alignments)),
+ random_target_std=float(np.std(target_alignments, ddof=1)),
+ random_target_p01=float(np.quantile(target_alignments, 0.01)),
+ random_target_p50=float(np.quantile(target_alignments, 0.50)),
+ random_target_p99=float(np.quantile(target_alignments, 0.99)),
+ min_eigen_target_alignment=float(eigenvalues[0]),
+ max_eigen_target_alignment=float(eigenvalues[-1]),
+ frobenius_distance_to_isotropic=float(np.linalg.norm(moment - isotropic)),
+ )
+ return summary, eigenvalues, target_alignments
+
+
+def write_outputs(
+ config: RunConfig,
+ summaries: list[DistributionSummary],
+ spectra: dict[str, np.ndarray],
+ target_alignments: dict[str, np.ndarray],
+ outdir: Path,
+) -> None:
+ outdir.mkdir(parents=True, exist_ok=True)
+ payload = {
+ "config": asdict(config),
+ "summaries": [asdict(summary) for summary in summaries],
+ }
+ (outdir / "summary.json").write_text(
+ json.dumps(payload, indent=2, sort_keys=True) + "\n"
+ )
+
+ spectrum_payload = {
+ name: values.tolist() for name, values in sorted(spectra.items())
+ }
+ (outdir / "eigenvalues.json").write_text(
+ json.dumps(spectrum_payload, indent=2, sort_keys=True) + "\n"
+ )
+
+ target_payload = {
+ name: values.tolist() for name, values in sorted(target_alignments.items())
+ }
+ (outdir / "random_target_alignments.json").write_text(
+ json.dumps(target_payload, indent=2, sort_keys=True) + "\n"
+ )
+
+
+def save_plots(
+ summaries: list[DistributionSummary],
+ spectra: dict[str, np.ndarray],
+ target_alignments: dict[str, np.ndarray],
+ outdir: Path,
+) -> list[Path]:
+ outdir.mkdir(parents=True, exist_ok=True)
+ paths: list[Path] = []
+ dimension = len(next(iter(spectra.values())))
+ minimax_bound = 1.0 / dimension
+
+ spectrum_path = outdir / "eigenvalue_spectra.png"
+ plt.figure(figsize=(7, 4.5))
+ for name, values in sorted(spectra.items()):
+ plt.plot(np.arange(1, dimension + 1), np.sort(values), label=name)
+ plt.axhline(minimax_bound, color="black", linestyle="--", linewidth=1, label="1/D")
+ plt.xlabel("eigenvalue index")
+ plt.ylabel("eigenvalue of M_mu")
+ plt.title("Feedback second-moment spectra")
+ plt.legend()
+ plt.tight_layout()
+ plt.savefig(spectrum_path, dpi=180)
+ plt.close()
+ paths.append(spectrum_path)
+
+ worst_case_path = outdir / "worst_case_alignment.png"
+ names = [summary.name for summary in summaries]
+ worst = [summary.lambda_min for summary in summaries]
+ best = [summary.lambda_max for summary in summaries]
+ x = np.arange(len(names))
+ plt.figure(figsize=(7, 4.5))
+ plt.bar(x - 0.18, worst, width=0.36, label="worst target")
+ plt.bar(x + 0.18, best, width=0.36, label="best target")
+ plt.axhline(minimax_bound, color="black", linestyle="--", linewidth=1, label="1/D")
+ plt.xticks(x, names, rotation=25, ha="right")
+ plt.ylabel("expected squared alignment")
+ plt.title("Worst-case target penalty")
+ plt.legend()
+ plt.tight_layout()
+ plt.savefig(worst_case_path, dpi=180)
+ plt.close()
+ paths.append(worst_case_path)
+
+ target_path = outdir / "random_target_alignment_hist.png"
+ plt.figure(figsize=(7, 4.5))
+ for name, values in sorted(target_alignments.items()):
+ plt.hist(values, bins=70, density=True, histtype="step", linewidth=1.5, label=name)
+ plt.axvline(minimax_bound, color="black", linestyle="--", linewidth=1, label="1/D")
+ plt.xlabel("a^T M_mu a for random target a")
+ plt.ylabel("density")
+ plt.title("Random target alignment")
+ plt.legend()
+ plt.tight_layout()
+ plt.savefig(target_path, dpi=180)
+ plt.close()
+ paths.append(target_path)
+
+ return paths
+
+
+def main() -> None:
+ args = parse_args()
+ if args.dimension < 2:
+ raise ValueError("--dimension must be at least 2.")
+ if args.feedback_samples < 1:
+ raise ValueError("--feedback-samples must be positive.")
+ if args.target_samples < 1:
+ raise ValueError("--target-samples must be positive.")
+ if args.batch_size < 1:
+ raise ValueError("--batch-size must be positive.")
+ if args.subspace_dim < 1:
+ raise ValueError("--subspace-dim must be positive.")
+ if args.anisotropy < 1:
+ raise ValueError("--anisotropy must be at least 1.")
+
+ config = RunConfig(
+ dimension=args.dimension,
+ feedback_samples=args.feedback_samples,
+ target_samples=args.target_samples,
+ batch_size=args.batch_size,
+ seed=args.seed,
+ distributions=args.distribution,
+ subspace_dim=args.subspace_dim,
+ anisotropy=args.anisotropy,
+ outdir=str(args.outdir),
+ plot=args.plot,
+ )
+
+ rng = np.random.default_rng(args.seed)
+ summaries: list[DistributionSummary] = []
+ spectra: dict[str, np.ndarray] = {}
+ target_alignments: dict[str, np.ndarray] = {}
+
+ for distribution in args.distribution:
+ moment = estimate_second_moment(rng, config, distribution)
+ summary, eigenvalues, alignments = summarize_distribution(
+ distribution, moment, rng, args.target_samples
+ )
+ summaries.append(summary)
+ spectra[distribution] = eigenvalues
+ target_alignments[distribution] = alignments
+
+ write_outputs(config, summaries, spectra, target_alignments, args.outdir)
+ plot_paths = (
+ save_plots(summaries, spectra, target_alignments, args.outdir)
+ if args.plot
+ else []
+ )
+
+ print(f"dimension: {args.dimension}")
+ print(f"minimax bound 1/D: {1.0 / args.dimension:.8g}")
+ for summary in summaries:
+ print(
+ f"{summary.name}: "
+ f"lambda_min={summary.lambda_min:.8g}, "
+ f"lambda_max={summary.lambda_max:.8g}, "
+ f"random_mean={summary.random_target_mean:.8g}, "
+ f"gap_to_1/D={summary.isotropic_gap:.8g}"
+ )
+ print(f"summary: {args.outdir / 'summary.json'}")
+ for path in plot_paths:
+ print(f"plot: {path}")
+
+
+if __name__ == "__main__":
+ main()