summaryrefslogtreecommitdiff
path: root/scripts/aaai_depth_experiments.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 17:47:13 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-27 17:47:13 -0500
commite65dd2e2460b1da48c83a930304cf9b269fc4447 (patch)
tree94f583b0fb5e9989bd6c68f7f0e950c6ef91756d /scripts/aaai_depth_experiments.py
parent82a49011de15287583d8cfec11ac5cca7efee747 (diff)
Add AAAI depth experiments and diagnostic figures
Diffstat (limited to 'scripts/aaai_depth_experiments.py')
-rw-r--r--scripts/aaai_depth_experiments.py633
1 files changed, 633 insertions, 0 deletions
diff --git a/scripts/aaai_depth_experiments.py b/scripts/aaai_depth_experiments.py
new file mode 100644
index 0000000..bb52f85
--- /dev/null
+++ b/scripts/aaai_depth_experiments.py
@@ -0,0 +1,633 @@
+#!/usr/bin/env python3
+"""AAAI depth experiments for fixed random feedback.
+
+Part ``init`` validates the exact expected initial first-order deficit for both
+FA and DFA at hidden depths 1, 2, 3, 4, and 6. Part ``finite`` measures the
+FA/BP finite-time loss gap and evaluates the frozen-initialization, linear
+operator-velocity, and early-retangent predictors over depth and width.
+"""
+
+from __future__ import annotations
+
+import argparse
+import csv
+import json
+import math
+import sys
+import time
+from dataclasses import asdict, dataclass
+from pathlib import Path
+
+import matplotlib
+
+matplotlib.use("Agg")
+import matplotlib.pyplot as plt
+import numpy as np
+import torch
+
+SCRIPT_DIR = Path(__file__).resolve().parent
+if str(SCRIPT_DIR) not in sys.path:
+ sys.path.insert(0, str(SCRIPT_DIR))
+
+import feedback_rules as fr # noqa: E402
+import real_data_validation as rdv # noqa: E402
+
+
+@dataclass(frozen=True)
+class InitRow:
+ rule: str
+ depth: int
+ width: int
+ init_seed: int
+ feedback_draws: int
+ bp_speed: float
+ output_speed: float
+ predicted_deficit: float
+ empirical_deficit: float
+ empirical_stderr: float
+ empirical_std: float
+ calibration_error: float
+
+
+@dataclass(frozen=True)
+class FiniteRow:
+ depth: int
+ width: int
+ train_samples: int
+ init_seed: int
+ feedback_seed: int
+ lr: float
+ horizon: int
+ early_step: int
+ initial_deficit_prediction: float
+ empirical_bp_loss: float
+ empirical_fa_loss: float
+ empirical_gap: float
+ fixed_gap: float
+ velocity_gap: float
+ retangent_gap: float
+ fixed_error: float
+ velocity_error: float
+ retangent_error: float
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser(description=__doc__)
+ parser.add_argument(
+ "--part", choices=["init", "finite", "both", "summarize"], default="both"
+ )
+ parser.add_argument("--depths", type=int, nargs="+", default=[1, 2, 3, 4, 6])
+ parser.add_argument("--width", type=int, default=64, help="fixed width for initialization")
+ parser.add_argument("--finite-widths", type=int, nargs="+", default=[32, 64, 96])
+ parser.add_argument("--input-dim", type=int, default=16)
+ parser.add_argument("--output-dim", type=int, default=4)
+ parser.add_argument("--train-samples", type=int, default=64)
+ parser.add_argument("--data-seed", type=int, default=2027)
+ parser.add_argument("--init-seeds", type=int, default=6)
+ parser.add_argument("--feedback-draws", type=int, default=256)
+ parser.add_argument("--finite-init-seeds", type=int, default=3)
+ parser.add_argument("--finite-feedback-seeds", type=int, default=4)
+ parser.add_argument("--lr", type=float, default=1e-3)
+ parser.add_argument("--horizon", type=int, default=50)
+ parser.add_argument("--early-step", type=int, default=20)
+ parser.add_argument("--feedback-scale", choices=["relu", "fan-in", "unit"], default="relu")
+ parser.add_argument("--torch-threads", type=int, default=16)
+ parser.add_argument("--outdir", type=Path, default=Path("outputs/aaai_depth_experiments"))
+ return parser.parse_args()
+
+
+def make_data(args: argparse.Namespace) -> tuple[torch.Tensor, torch.Tensor]:
+ generator = torch.Generator().manual_seed(args.data_seed)
+ x = torch.randn(args.train_samples, args.input_dim, generator=generator, dtype=torch.float64)
+ y = torch.randn(args.train_samples, args.output_dim, generator=generator, dtype=torch.float64)
+ return x, y
+
+
+def write_dataclass_csv(path: Path, rows: list[object]) -> None:
+ if not rows:
+ return
+ path.parent.mkdir(parents=True, exist_ok=True)
+ first = asdict(rows[0]) # type: ignore[arg-type]
+ with path.open("w", newline="") as handle:
+ writer = csv.DictWriter(handle, fieldnames=list(first))
+ writer.writeheader()
+ for row in rows:
+ writer.writerow(asdict(row)) # type: ignore[arg-type]
+
+
+def run_initialization(args: argparse.Namespace, x: torch.Tensor, y: torch.Tensor) -> list[InitRow]:
+ rows: list[InitRow] = []
+ for depth in args.depths:
+ dims = [args.input_dim, *([args.width] * depth), args.output_dim]
+ for init_index in range(args.init_seeds):
+ init_seed = 10_000 + init_index
+ weights = fr.initialize_mlp(dims, seed=init_seed)
+ bp_grads = fr.gradients(weights, x, y)
+ bp_speed = fr.squared_norm(bp_grads)
+ output_speed = fr.squared_norm([bp_grads[-1]])
+ predicted = 1.0 - output_speed / bp_speed
+ for rule_index, rule in enumerate(("fa", "dfa")):
+ deficits: list[float] = []
+ for draw in range(args.feedback_draws):
+ feedback_seed = 100_000_000 + 10_000 * depth + 1_000 * init_index + 2 * draw + rule_index
+ feedback = fr.init_feedback(
+ dims,
+ seed=feedback_seed,
+ rule=rule,
+ mode=args.feedback_scale,
+ )
+ rule_grads = fr.gradients(weights, x, y, rule=rule, feedback=feedback)
+ deficits.append(1.0 - fr.inner_product(bp_grads, rule_grads) / bp_speed)
+ values = np.asarray(deficits)
+ empirical = float(values.mean())
+ row = InitRow(
+ rule=rule.upper(),
+ depth=depth,
+ width=args.width,
+ init_seed=init_seed,
+ feedback_draws=args.feedback_draws,
+ bp_speed=bp_speed,
+ output_speed=output_speed,
+ predicted_deficit=predicted,
+ empirical_deficit=empirical,
+ empirical_stderr=float(values.std(ddof=1) / math.sqrt(len(values))),
+ empirical_std=float(values.std(ddof=1)),
+ calibration_error=empirical - predicted,
+ )
+ rows.append(row)
+ print(
+ f"[init] {rule.upper()} depth={depth} init={init_index}: "
+ f"prediction={predicted:.5f}, measured={empirical:.5f} "
+ f"+/- {row.empirical_stderr:.5f}",
+ flush=True,
+ )
+ return rows
+
+
+def train_with_snapshots(
+ weights0: list[torch.Tensor],
+ x: torch.Tensor,
+ y: torch.Tensor,
+ lr: float,
+ horizon: int,
+ feedback: list[torch.Tensor] | None,
+ early_step: int,
+) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
+ final, snapshots = rdv.train_with_snapshots(
+ weights0,
+ x,
+ y,
+ lr,
+ horizon,
+ feedback,
+ {early_step},
+ )
+ return final, snapshots[early_step]
+
+
+def retangent_loss(
+ kernel_s: torch.Tensor,
+ residual_s: torch.Tensor,
+ lr: float,
+ samples: int,
+ horizon: int,
+ early_step: int,
+) -> float:
+ return rdv.rollout_loss(
+ kernel_s,
+ None,
+ residual_s,
+ lr,
+ samples,
+ horizon - early_step,
+ early_step,
+ )
+
+
+def run_finite_time(args: argparse.Namespace, x: torch.Tensor, y: torch.Tensor) -> list[FiniteRow]:
+ rows: list[FiniteRow] = []
+ n = x.shape[0]
+ for depth in args.depths:
+ for width in args.finite_widths:
+ dims = [args.input_dim, *([width] * depth), args.output_dim]
+ for init_index in range(args.finite_init_seeds):
+ block_start = time.time()
+ init_seed = 20_000 + init_index
+ weights0 = fr.initialize_mlp(dims, seed=init_seed)
+ r0 = (fr.predict(weights0, x) - y).reshape(-1)
+ bp_grads = fr.gradients(weights0, x, y)
+ initial_prediction = 1.0 - fr.squared_norm([bp_grads[-1]]) / fr.squared_norm(bp_grads)
+
+ k_bp0 = rdv.tangent_kernel(weights0, x, None)
+ bp_final, bp_early = train_with_snapshots(
+ weights0, x, y, args.lr, args.horizon, None, args.early_step
+ )
+ k_bps = rdv.tangent_kernel(bp_early, x, None)
+ bp_emp = fr.mse(bp_final, x, y)
+ bp_fixed = rdv.rollout_loss(
+ k_bp0, None, r0, args.lr, n, args.horizon, args.early_step
+ )
+ bp_velocity = rdv.rollout_loss(
+ k_bp0,
+ k_bps - k_bp0,
+ r0,
+ args.lr,
+ n,
+ args.horizon,
+ args.early_step,
+ )
+ bp_residual_s = (fr.predict(bp_early, x) - y).reshape(-1)
+ bp_retangent = retangent_loss(
+ k_bps, bp_residual_s, args.lr, n, args.horizon, args.early_step
+ )
+
+ for feedback_index in range(args.finite_feedback_seeds):
+ feedback_seed = 200_000 + 1_000 * init_index + feedback_index
+ feedback = fr.init_feedback(
+ dims,
+ seed=feedback_seed,
+ rule="fa",
+ mode=args.feedback_scale,
+ )
+ k_fa0 = rdv.tangent_kernel(weights0, x, feedback)
+ fa_final, fa_early = train_with_snapshots(
+ weights0,
+ x,
+ y,
+ args.lr,
+ args.horizon,
+ feedback,
+ args.early_step,
+ )
+ k_fas = rdv.tangent_kernel(fa_early, x, feedback)
+ fa_emp = fr.mse(fa_final, x, y)
+ fa_fixed = rdv.rollout_loss(
+ k_fa0, None, r0, args.lr, n, args.horizon, args.early_step
+ )
+ fa_velocity = rdv.rollout_loss(
+ k_fa0,
+ k_fas - k_fa0,
+ r0,
+ args.lr,
+ n,
+ args.horizon,
+ args.early_step,
+ )
+ fa_residual_s = (fr.predict(fa_early, x) - y).reshape(-1)
+ fa_retangent = retangent_loss(
+ k_fas, fa_residual_s, args.lr, n, args.horizon, args.early_step
+ )
+
+ empirical_gap = fa_emp - bp_emp
+ fixed_gap = fa_fixed - bp_fixed
+ velocity_gap = fa_velocity - bp_velocity
+ retangent_gap = fa_retangent - bp_retangent
+ rows.append(
+ FiniteRow(
+ depth=depth,
+ width=width,
+ train_samples=n,
+ init_seed=init_seed,
+ feedback_seed=feedback_seed,
+ lr=args.lr,
+ horizon=args.horizon,
+ early_step=args.early_step,
+ initial_deficit_prediction=initial_prediction,
+ empirical_bp_loss=bp_emp,
+ empirical_fa_loss=fa_emp,
+ empirical_gap=empirical_gap,
+ fixed_gap=fixed_gap,
+ velocity_gap=velocity_gap,
+ retangent_gap=retangent_gap,
+ fixed_error=fixed_gap - empirical_gap,
+ velocity_error=velocity_gap - empirical_gap,
+ retangent_error=retangent_gap - empirical_gap,
+ )
+ )
+ print(
+ f"[finite] depth={depth} width={width} init={init_index}: "
+ f"{args.finite_feedback_seeds} feedbacks in {time.time() - block_start:.1f}s",
+ flush=True,
+ )
+ return rows
+
+
+def safe_corr(left: np.ndarray, right: np.ndarray) -> float:
+ if len(left) < 2 or float(np.std(left) * np.std(right)) == 0.0:
+ return math.nan
+ return float(np.corrcoef(left, right)[0, 1])
+
+
+def summarize_finite(rows: list[FiniteRow]) -> list[dict[str, float | int | str]]:
+ metrics: list[dict[str, float | int | str]] = []
+ groups = sorted({(row.depth, row.width) for row in rows})
+ for depth, width in groups:
+ subset = [row for row in rows if row.depth == depth and row.width == width]
+ empirical = np.asarray([row.empirical_gap for row in subset])
+ for predictor, field in (
+ ("frozen K(0)", "fixed_gap"),
+ ("linear velocity", "velocity_gap"),
+ ("early retangent", "retangent_gap"),
+ ):
+ predicted = np.asarray([getattr(row, field) for row in subset])
+ metrics.append(
+ {
+ "depth": depth,
+ "width": width,
+ "predictor": predictor,
+ "rows": len(subset),
+ "empirical_gap_mean": float(empirical.mean()),
+ "empirical_gap_std": float(empirical.std(ddof=1)),
+ "mae": float(np.mean(np.abs(predicted - empirical))),
+ "bias": float(np.mean(predicted - empirical)),
+ "corr": safe_corr(predicted, empirical),
+ }
+ )
+ return metrics
+
+
+def write_dict_csv(path: Path, rows: list[dict[str, object]]) -> None:
+ if not rows:
+ return
+ with path.open("w", newline="") as handle:
+ writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
+ writer.writeheader()
+ writer.writerows(rows)
+
+
+def plot_initial(rows: list[InitRow], outdir: Path) -> None:
+ fig, ax = plt.subplots(figsize=(5.7, 5.2), dpi=180)
+ colors = plt.cm.viridis(np.linspace(0.1, 0.9, len(sorted({row.depth for row in rows}))))
+ for color, depth in zip(colors, sorted({row.depth for row in rows})):
+ for rule, marker in (("FA", "o"), ("DFA", "s")):
+ subset = [row for row in rows if row.depth == depth and row.rule == rule]
+ ax.errorbar(
+ [row.predicted_deficit for row in subset],
+ [row.empirical_deficit for row in subset],
+ yerr=[2 * row.empirical_stderr for row in subset],
+ fmt=marker,
+ ms=4.5,
+ color=color,
+ alpha=0.78,
+ capsize=1.5,
+ lw=0.8,
+ label=f"depth {depth}, {rule}",
+ )
+ values = [value for row in rows for value in (row.predicted_deficit, row.empirical_deficit)]
+ lo, hi = min(values), max(values)
+ pad = 0.04 * (hi - lo + 1e-12)
+ ax.plot([lo - pad, hi + pad], [lo - pad, hi + pad], color="black", lw=1)
+ ax.set_xlabel("exact expected initial deficit")
+ ax.set_ylabel("measured mean initial deficit")
+ ax.set_title("FA and DFA share the exact initialization cost")
+ ax.grid(alpha=0.16)
+ ax.legend(fontsize=6.5, ncols=2)
+ fig.tight_layout()
+ fig.savefig(outdir / "initial_deficit_calibration.png", bbox_inches="tight")
+ plt.close(fig)
+
+
+def plot_finite(rows: list[FiniteRow], metrics: list[dict[str, object]], outdir: Path) -> None:
+ fig, ax = plt.subplots(figsize=(6.4, 4.7), dpi=180)
+ for width, marker in zip(sorted({row.width for row in rows}), ("o", "s", "^", "D")):
+ depths = sorted({row.depth for row in rows if row.width == width})
+ means = [np.mean([row.empirical_gap for row in rows if row.width == width and row.depth == d]) for d in depths]
+ sems = []
+ for depth in depths:
+ cell = [row for row in rows if row.width == width and row.depth == depth]
+ init_means = [
+ np.mean([row.empirical_gap for row in cell if row.init_seed == init_seed])
+ for init_seed in sorted({row.init_seed for row in cell})
+ ]
+ sems.append(float(np.std(init_means, ddof=1) / math.sqrt(len(init_means))))
+ ax.errorbar(depths, means, yerr=sems, marker=marker, capsize=3, label=f"width {width}")
+ ax.set_xlabel("hidden-layer depth")
+ ax.set_ylabel("FA loss - BP loss at fixed horizon")
+ ax.set_title("Finite-time optimization cost across depth and width")
+ ax.grid(alpha=0.16)
+ ax.legend()
+ fig.tight_layout()
+ fig.savefig(outdir / "finite_gap_by_depth_width.png", bbox_inches="tight")
+ plt.close(fig)
+
+ fig, ax = plt.subplots(figsize=(7.0, 4.8), dpi=180)
+ predictors = ["frozen K(0)", "linear velocity", "early retangent"]
+ for predictor, marker in zip(predictors, ("^", "o", "s")):
+ subset = [row for row in metrics if row["predictor"] == predictor and row["width"] == 64]
+ ax.plot(
+ [int(row["depth"]) for row in subset],
+ [float(row["mae"]) for row in subset],
+ marker=marker,
+ label=predictor,
+ )
+ ax.set_xlabel("hidden-layer depth")
+ ax.set_ylabel("gap prediction MAE (width 64)")
+ ax.set_title("Operator drift, not initialization mismatch, controls later error")
+ ax.set_yscale("log")
+ ax.grid(alpha=0.16, which="both")
+ ax.legend()
+ fig.tight_layout()
+ fig.savefig(outdir / "predictor_mae_by_depth.png", bbox_inches="tight")
+ plt.close(fig)
+
+
+def summarize_mismatch_vs_cost(
+ init_rows: list[dict[str, str]],
+ finite_rows: list[dict[str, str]],
+ width: int,
+ output_dim: int,
+ outdir: Path,
+) -> list[dict[str, object]]:
+ """Contrast exponentially compounding matrix mismatch with measured costs."""
+ depths = sorted({int(row["depth"]) for row in init_rows})
+ records: list[dict[str, object]] = []
+ for depth in depths:
+ init_subset = [row for row in init_rows if int(row["depth"]) == depth]
+ finite_subset = [
+ row
+ for row in finite_rows
+ if int(row["depth"]) == depth and int(row["width"]) == width
+ ]
+ # One width x output block plus depth-1 width x width FA blocks.
+ joint_alignment_proxy = (1.0 / (width * output_dim)) * (
+ 1.0 / (width * width)
+ ) ** (depth - 1)
+ records.append(
+ {
+ "depth": depth,
+ "width": width,
+ "joint_alignment_proxy": joint_alignment_proxy,
+ "negative_log10_joint_alignment": -math.log10(joint_alignment_proxy),
+ "initial_deficit_mean": float(
+ np.mean([float(row["empirical_deficit"]) for row in init_subset])
+ ),
+ "finite_gap_mean": float(
+ np.mean([float(row["empirical_gap"]) for row in finite_subset])
+ ),
+ }
+ )
+
+ write_dict_csv(outdir / "mismatch_vs_cost_by_depth.csv", records)
+ fig, ax = plt.subplots(figsize=(7.0, 4.8), dpi=180)
+ ax.plot(
+ [int(row["depth"]) for row in records],
+ [float(row["joint_alignment_proxy"]) for row in records],
+ "o--",
+ color="black",
+ label=r"joint matrix-alignment proxy $\prod_l 1/D_l$",
+ )
+ ax.set_yscale("log")
+ ax.set_xlabel("hidden-layer depth")
+ ax.set_ylabel("joint squared-alignment proxy")
+ ax.grid(alpha=0.16, which="both")
+ ax2 = ax.twinx()
+ ax2.plot(
+ [int(row["depth"]) for row in records],
+ [float(row["initial_deficit_mean"]) for row in records],
+ "s-",
+ color="#2f6f9f",
+ label="measured initial deficit",
+ )
+ ax2.plot(
+ [int(row["depth"]) for row in records],
+ [float(row["finite_gap_mean"]) for row in records],
+ "^-",
+ color="#c65f16",
+ label="measured finite-time gap",
+ )
+ ax2.set_ylabel("optimization cost")
+ handles1, labels1 = ax.get_legend_handles_labels()
+ handles2, labels2 = ax2.get_legend_handles_labels()
+ ax2.legend(handles1 + handles2, labels1 + labels2, fontsize=8, loc="center right")
+ ax.set_title("Matrix mismatch compounds exponentially; optimization cost does not")
+ fig.tight_layout()
+ fig.savefig(outdir / "mismatch_vs_cost_by_depth.png", bbox_inches="tight")
+ plt.close(fig)
+ return records
+
+
+def main() -> None:
+ args = parse_args()
+ if args.early_step <= 0 or args.early_step >= args.horizon:
+ raise ValueError("early-step must lie strictly between 0 and horizon")
+ torch.set_num_threads(args.torch_threads)
+ args.outdir.mkdir(parents=True, exist_ok=True)
+ x, y = make_data(args)
+ summary: dict[str, object] = {
+ "config": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}
+ }
+
+ if args.part == "summarize":
+ with (args.outdir / "initialization_rows.csv").open() as handle:
+ existing_init = list(csv.DictReader(handle))
+ with (args.outdir / "finite_time_rows.csv").open() as handle:
+ existing_finite = list(csv.DictReader(handle))
+ init_errors = np.asarray([float(row["calibration_error"]) for row in existing_init])
+ init_stderr = np.asarray([float(row["empirical_stderr"]) for row in existing_init])
+ summary["initialization"] = {
+ "rows": len(existing_init),
+ "max_abs_calibration_error": float(np.max(np.abs(init_errors))),
+ "max_standardized_error": float(np.max(np.abs(init_errors) / init_stderr)),
+ "rules": sorted({row["rule"] for row in existing_init}),
+ "depths": sorted({int(row["depth"]) for row in existing_init}),
+ }
+ empirical = np.asarray([float(row["empirical_gap"]) for row in existing_finite])
+ summary["finite_time"] = {
+ "rows": len(existing_finite),
+ "depths": sorted({int(row["depth"]) for row in existing_finite}),
+ "widths": sorted({int(row["width"]) for row in existing_finite}),
+ "overall": {
+ predictor: {
+ "mae": float(np.mean(np.abs(predicted - empirical))),
+ "corr": safe_corr(predicted, empirical),
+ }
+ for predictor, predicted in (
+ (
+ "frozen",
+ np.asarray([float(row["fixed_gap"]) for row in existing_finite]),
+ ),
+ (
+ "linear_velocity",
+ np.asarray([float(row["velocity_gap"]) for row in existing_finite]),
+ ),
+ (
+ "early_retangent",
+ np.asarray([float(row["retangent_gap"]) for row in existing_finite]),
+ ),
+ )
+ },
+ }
+ mismatch_records = summarize_mismatch_vs_cost(
+ existing_init,
+ existing_finite,
+ width=args.width,
+ output_dim=args.output_dim,
+ outdir=args.outdir,
+ )
+ summary["mismatch_vs_cost"] = {
+ "joint_alignment_proxy_depth_1": mismatch_records[0]["joint_alignment_proxy"],
+ "joint_alignment_proxy_depth_6": mismatch_records[-1]["joint_alignment_proxy"],
+ "orders_of_magnitude_drop": float(
+ math.log10(
+ float(mismatch_records[0]["joint_alignment_proxy"])
+ / float(mismatch_records[-1]["joint_alignment_proxy"])
+ )
+ ),
+ }
+ (args.outdir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n")
+ print(f"summary: {args.outdir / 'summary.json'}")
+ return
+
+ if args.part in {"init", "both"}:
+ init_rows = run_initialization(args, x, y)
+ write_dataclass_csv(args.outdir / "initialization_rows.csv", init_rows)
+ plot_initial(init_rows, args.outdir)
+ summary["initialization"] = {
+ "rows": len(init_rows),
+ "max_abs_calibration_error": max(abs(row.calibration_error) for row in init_rows),
+ "max_standardized_error": max(
+ abs(row.calibration_error) / row.empirical_stderr for row in init_rows
+ ),
+ "rules": sorted({row.rule for row in init_rows}),
+ "depths": sorted({row.depth for row in init_rows}),
+ }
+
+ if args.part in {"finite", "both"}:
+ finite_rows = run_finite_time(args, x, y)
+ write_dataclass_csv(args.outdir / "finite_time_rows.csv", finite_rows)
+ metrics = summarize_finite(finite_rows)
+ write_dict_csv(args.outdir / "finite_time_metrics.csv", metrics) # type: ignore[arg-type]
+ plot_finite(finite_rows, metrics, args.outdir)
+ summary["finite_time"] = {
+ "rows": len(finite_rows),
+ "depths": sorted({row.depth for row in finite_rows}),
+ "widths": sorted({row.width for row in finite_rows}),
+ "overall": {
+ predictor: {
+ "mae": float(
+ np.mean(
+ [
+ abs(getattr(row, field) - row.empirical_gap)
+ for row in finite_rows
+ ]
+ )
+ ),
+ "corr": safe_corr(
+ np.asarray([getattr(row, field) for row in finite_rows]),
+ np.asarray([row.empirical_gap for row in finite_rows]),
+ ),
+ }
+ for predictor, field in (
+ ("frozen", "fixed_gap"),
+ ("linear_velocity", "velocity_gap"),
+ ("early_retangent", "retangent_gap"),
+ )
+ },
+ }
+
+ (args.outdir / "summary.json").write_text(json.dumps(summary, indent=2) + "\n")
+ print(f"summary: {args.outdir / 'summary.json'}")
+
+
+if __name__ == "__main__":
+ main()