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