NeuralGCM / scripts /checkpoint_info.py
yzt15806542928's picture
Upload folder using huggingface_hub
f4a39ee verified
Raw
History Blame Contribute Delete
1.52 kB
#!/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()