#!/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()