summaryrefslogtreecommitdiff
path: root/run_mvp.py
diff options
context:
space:
mode:
Diffstat (limited to 'run_mvp.py')
-rw-r--r--run_mvp.py53
1 files changed, 53 insertions, 0 deletions
diff --git a/run_mvp.py b/run_mvp.py
new file mode 100644
index 0000000..c55b33c
--- /dev/null
+++ b/run_mvp.py
@@ -0,0 +1,53 @@
+#!/usr/bin/env python3
+"""Run the anonymous KAFT MVP from the command line."""
+
+from __future__ import annotations
+
+import argparse
+
+import pandas as pd
+
+from kaft_mvp import MVPConfig, run_mvp
+
+
+def main() -> None:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--device", default="cpu")
+ parser.add_argument("--seeds", default="0,1,2")
+ parser.add_argument("--epochs", type=int, default=200)
+ parser.add_argument("--diagnostic-epochs", type=int, default=100)
+ parser.add_argument("--output-dir", default="artifacts")
+ args = parser.parse_args()
+ seeds = tuple(int(value) for value in args.seeds.split(","))
+ config = MVPConfig(
+ seeds=seeds,
+ epochs=args.epochs,
+ diagnostic_epochs=args.diagnostic_epochs,
+ device=args.device,
+ )
+ payload = run_mvp(config=config, output_dir=args.output_dir)
+ print("\nBP versus KAFT")
+ print(pd.DataFrame(payload["summary"]).to_string(index=False))
+ diagnostic = pd.DataFrame(
+ [
+ {
+ "seed": row["seed"],
+ "all_weight_grads_zero": row[
+ "all_weight_gradients_exact_zero"
+ ],
+ "output_adjacent_error": row[
+ "output_adjacent_error_frobenius"
+ ],
+ "hidden_probe_percent": 100.0
+ * row["standardized_penultimate_probe_accuracy"],
+ }
+ for row in payload["gradient_diagnostics"]
+ ]
+ )
+ print("\n10-layer BP diagnostic")
+ print(diagnostic.to_string(index=False))
+ print(f"\nArtifacts written to {args.output_dir}/")
+
+
+if __name__ == "__main__":
+ main()