summaryrefslogtreecommitdiff
path: root/scripts
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-05-29 08:43:19 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-05-29 08:43:19 -0500
commitdd14582dcfd3b6e3e7b0e28f66bb0f2994f106c4 (patch)
tree99c06ef0e6bb26a040409408040e1e88bc14f69f /scripts
parentc47a74792e6ee78181b087ead595465140b77825 (diff)
Add multilayer capacity distribution matching
Diffstat (limited to 'scripts')
-rw-r--r--scripts/README.md27
-rwxr-xr-xscripts/multilayer_capacity_distribution.py326
2 files changed, 353 insertions, 0 deletions
diff --git a/scripts/README.md b/scripts/README.md
index e30b309..efd4208 100644
--- a/scripts/README.md
+++ b/scripts/README.md
@@ -85,6 +85,33 @@ Outputs are written under `outputs/capacity_empirical_validation/`:
- `multilayer_capacity.csv`
- diagnostic plots when `--plot` is set.
+## Multilayer Observed Capacity Distribution
+
+Run:
+
+```bash
+python scripts/multilayer_capacity_distribution.py --dimensions 64 256 1024 4096 --layers 1 2 4 8 16 --samples 100000 --batch-size 8192 --seed 456 --plot
+```
+
+This validates the full null distribution of observed capacity surprisal:
+
+\[
+S_l=-\log P(Q_l'\ge Q_l).
+\]
+
+Theory predicts:
+
+\[
+S_l\sim \mathrm{Exp}(1),
+\qquad
+\sum_{l=1}^L S_l\sim \mathrm{Gamma}(L,1).
+\]
+
+The script samples random direction cosines through the exact chi-square
+representation \(Q=X/(X+Y)\), with \(X\sim\chi^2_1\) and
+\(Y\sim\chi^2_{D-1}\). Outputs are written under
+`outputs/multilayer_capacity_distribution/`.
+
## Minimax Initialization Bound
Run:
diff --git a/scripts/multilayer_capacity_distribution.py b/scripts/multilayer_capacity_distribution.py
new file mode 100755
index 0000000..d238241
--- /dev/null
+++ b/scripts/multilayer_capacity_distribution.py
@@ -0,0 +1,326 @@
+#!/usr/bin/env python3
+"""Match observed FA capacity surprisal to Exp/Gamma laws.
+
+If Q follows the beta alignment law and S = -log P(Q' >= Q), then S is Exp(1).
+For L independent layers, sum_l S_l is Gamma(L, 1). This script validates the
+full predicted distribution, not only tail probabilities. It samples random
+direction cosines using the exact chi-square representation of a Gaussian
+direction, Q = X / (X + Y), with X ~ chi^2_1 and Y ~ chi^2_{D-1}.
+"""
+
+from __future__ import annotations
+
+import argparse
+import csv
+import json
+from dataclasses import asdict, dataclass
+from pathlib import Path
+
+import matplotlib.pyplot as plt
+import numpy as np
+from scipy import stats
+
+
+@dataclass(frozen=True)
+class RunConfig:
+ dimensions: list[int]
+ layers: list[int]
+ samples: int
+ batch_size: int
+ seed: int
+ outdir: str
+ plot: bool
+
+
+@dataclass(frozen=True)
+class DistributionMatchRow:
+ dimension: int
+ layers: int
+ samples: int
+ empirical_mean: float
+ theoretical_mean: float
+ empirical_var: float
+ theoretical_var: float
+ ks_statistic: float
+ ks_pvalue: float
+ q50_empirical: float
+ q50_theoretical: float
+ q90_empirical: float
+ q90_theoretical: float
+ q99_empirical: float
+ q99_theoretical: float
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(
+ description="Validate observed capacity surprisal distributions."
+ )
+ parser.add_argument(
+ "--dimensions",
+ type=int,
+ nargs="+",
+ default=[64, 256, 1024, 4096],
+ )
+ parser.add_argument("--layers", type=int, nargs="+", default=[1, 2, 4, 8, 16])
+ parser.add_argument("--samples", type=int, default=100_000)
+ parser.add_argument("--batch-size", type=int, default=2048)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument(
+ "--outdir",
+ type=Path,
+ default=Path("outputs/multilayer_capacity_distribution"),
+ )
+ parser.add_argument("--plot", action="store_true")
+ return parser.parse_args()
+
+
+def validate_config(config: RunConfig) -> None:
+ if any(d < 2 for d in config.dimensions):
+ raise ValueError("All dimensions must be at least 2.")
+ if any(l < 1 for l in config.layers):
+ raise ValueError("All layer counts must be positive.")
+ if config.samples < 1:
+ raise ValueError("--samples must be positive.")
+ if config.batch_size < 1:
+ raise ValueError("--batch-size must be positive.")
+
+
+def sample_q_matrix(
+ rng: np.random.Generator,
+ dimension: int,
+ layers: int,
+ samples: int,
+ batch_size: int,
+) -> np.ndarray:
+ values = np.empty((samples, layers), dtype=np.float64)
+ offset = 0
+ while offset < samples:
+ count = min(batch_size, samples - offset)
+ z = rng.standard_normal((count, layers, dimension))
+ numerator = z[:, :, 0] * z[:, :, 0]
+ denominator = np.einsum("bld,bld->bl", z, z)
+ values[offset : offset + count, :] = numerator / denominator
+ offset += count
+ return values
+
+
+def sample_total_surprisal_by_layer_count(
+ rng: np.random.Generator,
+ dimension: int,
+ layer_counts: list[int],
+ samples: int,
+ batch_size: int,
+) -> dict[int, np.ndarray]:
+ max_layers = max(layer_counts)
+ outputs = {
+ layers: np.empty(samples, dtype=np.float64) for layers in sorted(layer_counts)
+ }
+ beta_dist = stats.beta(0.5, (dimension - 1) / 2)
+
+ offset = 0
+ while offset < samples:
+ count = min(batch_size, samples - offset)
+ numerator = rng.chisquare(df=1.0, size=(count, max_layers))
+ remainder = rng.chisquare(df=dimension - 1.0, size=(count, max_layers))
+ q_values = numerator / (numerator + remainder)
+ survival = beta_dist.sf(q_values)
+ survival = np.clip(survival, np.finfo(float).tiny, 1.0)
+ cumulative = np.cumsum(-np.log(survival), axis=1)
+ for layers, values in outputs.items():
+ values[offset : offset + count] = cumulative[:, layers - 1]
+ offset += count
+ return outputs
+
+
+def observed_surprisal(q_values: np.ndarray, dimension: int) -> np.ndarray:
+ beta_dist = stats.beta(0.5, (dimension - 1) / 2)
+ survival = beta_dist.sf(q_values)
+ survival = np.clip(survival, np.finfo(float).tiny, 1.0)
+ return -np.log(survival)
+
+
+def summarize(values: np.ndarray, dimension: int, layers: int) -> DistributionMatchRow:
+ gamma_dist = stats.gamma(a=layers, scale=1.0)
+ ks = stats.kstest(values, gamma_dist.cdf)
+ return DistributionMatchRow(
+ dimension=dimension,
+ layers=layers,
+ samples=len(values),
+ empirical_mean=float(np.mean(values)),
+ theoretical_mean=float(gamma_dist.mean()),
+ empirical_var=float(np.var(values, ddof=1)),
+ theoretical_var=float(gamma_dist.var()),
+ ks_statistic=float(ks.statistic),
+ ks_pvalue=float(ks.pvalue),
+ q50_empirical=float(np.quantile(values, 0.50)),
+ q50_theoretical=float(gamma_dist.ppf(0.50)),
+ q90_empirical=float(np.quantile(values, 0.90)),
+ q90_theoretical=float(gamma_dist.ppf(0.90)),
+ q99_empirical=float(np.quantile(values, 0.99)),
+ q99_theoretical=float(gamma_dist.ppf(0.99)),
+ )
+
+
+def write_csv(path: Path, rows: list[object]) -> None:
+ if not rows:
+ return
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("w", newline="") as handle:
+ first = asdict(rows[0]) # type: ignore[arg-type]
+ writer = csv.DictWriter(handle, fieldnames=list(first.keys()))
+ writer.writeheader()
+ for row in rows:
+ writer.writerow(asdict(row)) # type: ignore[arg-type]
+
+
+def write_outputs(
+ config: RunConfig, rows: list[DistributionMatchRow], outdir: Path
+) -> None:
+ outdir.mkdir(parents=True, exist_ok=True)
+ write_csv(outdir / "distribution_match.csv", rows)
+ payload = {
+ "config": asdict(config),
+ "distribution_match": [asdict(row) for row in rows],
+ }
+ (outdir / "summary.json").write_text(
+ json.dumps(payload, indent=2, sort_keys=True) + "\n"
+ )
+
+
+def save_plots(
+ row_values: dict[tuple[int, int], np.ndarray],
+ rows: list[DistributionMatchRow],
+ outdir: Path,
+) -> list[Path]:
+ outdir.mkdir(parents=True, exist_ok=True)
+ paths: list[Path] = []
+
+ ks_path = outdir / "ks_heatmap.png"
+ dimensions = sorted({row.dimension for row in rows})
+ layer_counts = sorted({row.layers for row in rows})
+ heatmap = np.empty((len(layer_counts), len(dimensions)), dtype=np.float64)
+ for i, layers in enumerate(layer_counts):
+ for j, dimension in enumerate(dimensions):
+ match = next(
+ row
+ for row in rows
+ if row.dimension == dimension and row.layers == layers
+ )
+ heatmap[i, j] = match.ks_statistic
+
+ plt.figure(figsize=(7, 4.5))
+ plt.imshow(heatmap, aspect="auto", origin="lower", cmap="viridis")
+ plt.colorbar(label="KS statistic")
+ plt.xticks(range(len(dimensions)), dimensions)
+ plt.yticks(range(len(layer_counts)), layer_counts)
+ plt.xlabel("dimension D")
+ plt.ylabel("layers L")
+ plt.title("Gamma distribution calibration")
+ plt.tight_layout()
+ plt.savefig(ks_path, dpi=180)
+ plt.close()
+ paths.append(ks_path)
+
+ selected = [
+ (dimensions[0], layer_counts[0]),
+ (dimensions[0], layer_counts[-1]),
+ (dimensions[-1], layer_counts[0]),
+ (dimensions[-1], layer_counts[-1]),
+ ]
+ for dimension, layers in selected:
+ values = row_values[(dimension, layers)]
+ gamma_dist = stats.gamma(a=layers, scale=1.0)
+
+ hist_path = outdir / f"hist_D{dimension}_L{layers}.png"
+ x_max = float(max(np.quantile(values, 0.999), gamma_dist.ppf(0.999)))
+ xs = np.linspace(0.0, x_max, 600)
+ plt.figure(figsize=(7, 4.5))
+ plt.hist(values, bins=90, density=True, alpha=0.45, label="empirical")
+ plt.plot(xs, gamma_dist.pdf(xs), color="black", linewidth=2, label="Gamma")
+ plt.xlabel("observed total capacity surprisal")
+ plt.ylabel("density")
+ plt.title(f"D={dimension}, L={layers}: empirical vs Gamma({layers},1)")
+ plt.legend()
+ plt.tight_layout()
+ plt.savefig(hist_path, dpi=180)
+ plt.close()
+ paths.append(hist_path)
+
+ qq_path = outdir / f"qq_D{dimension}_L{layers}.png"
+ probs = (np.arange(1, len(values) + 1) - 0.5) / len(values)
+ empirical = np.sort(values)
+ theoretical = gamma_dist.ppf(probs)
+ max_value = float(max(empirical[-1], theoretical[-1]))
+ plt.figure(figsize=(5, 5))
+ plt.scatter(theoretical, empirical, s=4, alpha=0.25)
+ plt.plot([0, max_value], [0, max_value], color="black", linewidth=1)
+ plt.xlabel("theoretical Gamma quantile")
+ plt.ylabel("empirical quantile")
+ plt.title(f"Q-Q: D={dimension}, L={layers}")
+ plt.tight_layout()
+ plt.savefig(qq_path, dpi=180)
+ plt.close()
+ paths.append(qq_path)
+
+ return paths
+
+
+def parse_config(args: argparse.Namespace) -> RunConfig:
+ return RunConfig(
+ dimensions=args.dimensions,
+ layers=args.layers,
+ samples=args.samples,
+ batch_size=args.batch_size,
+ seed=args.seed,
+ outdir=str(args.outdir),
+ plot=args.plot,
+ )
+
+
+def main() -> None:
+ args = parse_args()
+ config = parse_config(args)
+ validate_config(config)
+ rng = np.random.default_rng(config.seed)
+
+ rows: list[DistributionMatchRow] = []
+ row_values: dict[tuple[int, int], np.ndarray] = {}
+
+ for dimension in config.dimensions:
+ surprisal_by_layers = sample_total_surprisal_by_layer_count(
+ rng, dimension, config.layers, config.samples, config.batch_size
+ )
+ for layers in config.layers:
+ surprisal = surprisal_by_layers[layers]
+ row = summarize(surprisal, dimension, layers)
+ rows.append(row)
+ row_values[(dimension, layers)] = surprisal
+ print(
+ f"D={dimension}, L={layers}: "
+ f"mean_emp={row.empirical_mean:.6g}, "
+ f"mean_theory={row.theoretical_mean:.6g}, "
+ f"KS={row.ks_statistic:.6g}"
+ )
+
+ outdir = Path(config.outdir)
+ write_outputs(config, rows, outdir)
+ plot_paths = save_plots(row_values, rows, outdir) if config.plot else []
+
+ max_ks = max(row.ks_statistic for row in rows)
+ mean_abs_mean_error = float(
+ np.mean([abs(row.empirical_mean - row.theoretical_mean) for row in rows])
+ )
+ mean_abs_var_error = float(
+ np.mean([abs(row.empirical_var - row.theoretical_var) for row in rows])
+ )
+ print(f"rows: {len(rows)}")
+ print(f"max_ks_statistic: {max_ks:.8g}")
+ print(f"mean_abs_mean_error: {mean_abs_mean_error:.8g}")
+ print(f"mean_abs_var_error: {mean_abs_var_error:.8g}")
+ print(f"distribution_match: {outdir / 'distribution_match.csv'}")
+ for path in plot_paths:
+ print(f"plot: {path}")
+
+
+if __name__ == "__main__":
+ main()