File size: 1,515 Bytes
f4a39ee | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 | #!/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()
|