#!/usr/bin/env python3 """Print parameter and serialization sizes for NeuralGCM checkpoints.""" from __future__ import annotations import argparse import pickle import sys try: from common import PROJECT_ROOT, resolve_path except ModuleNotFoundError: # supports ``python -m scripts.checkpoint_info`` from scripts.common import PROJECT_ROOT, resolve_path if str(PROJECT_ROOT) not in sys.path: sys.path.insert(0, str(PROJECT_ROOT)) from model.NeuralGCM import checkpoint_mode, format_parameter_summary def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("checkpoints", nargs="+") args = parser.parse_args() for value in args.checkpoints: path = resolve_path(value) with path.open("rb") as handle: payload = pickle.load(handle) if not isinstance(payload, dict) or "params" not in payload: raise ValueError(f"{path} does not contain an official params tree") mode = payload.get("mode") or checkpoint_mode(payload) or "unknown" training_state = payload.get("training_state") resume_text = ( f"resumable=true step={training_state.get('step')}" if isinstance(training_state, dict) else "resumable=false" ) print( f"checkpoint={path.name} mode={mode} " f"file.bytes={path.stat().st_size:,} " f"{resume_text} {format_parameter_summary(payload['params'])}" ) if __name__ == "__main__": main()