28 lines
720 B
Python
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())
|