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