Files
geointel/scripts/inspect_torch_checkpoint.py
T

28 lines
720 B
Python

#!/usr/bin/env python3
"""Print non-tensor provenance metadata from a PyTorch checkpoint."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("checkpoint", type=Path)
args = parser.parse_args()
import torch
checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
payload = {
key: checkpoint.get(key)
for key in ("date", "version", "license", "docs", "train_args")
if checkpoint.get(key) is not None
}
print(json.dumps(payload, indent=2, default=str))
return 0
if __name__ == "__main__":
raise SystemExit(main())