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()