40 lines
1.3 KiB
Python
40 lines
1.3 KiB
Python
"""Report safetensors metadata without materializing model weights."""
|
|
|
|
import argparse
|
|
import json
|
|
import struct
|
|
from collections import Counter
|
|
from pathlib import Path
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("checkpoint", type=Path)
|
|
parser.add_argument("--output", type=Path)
|
|
args = parser.parse_args()
|
|
|
|
with args.checkpoint.open("rb") as file:
|
|
header_size = struct.unpack("<Q", file.read(8))[0]
|
|
header = json.loads(file.read(header_size))
|
|
|
|
tensors = {name: spec for name, spec in header.items() if name != "__metadata__"}
|
|
report = {
|
|
"checkpoint": str(args.checkpoint),
|
|
"metadata": header.get("__metadata__", {}),
|
|
"tensor_count": len(tensors),
|
|
"dtypes": dict(sorted(Counter(spec["dtype"] for spec in tensors.values()).items())),
|
|
"tensors": {
|
|
name: {key: spec[key] for key in ("dtype", "shape", "data_offsets") if key in spec}
|
|
for name, spec in tensors.items()
|
|
},
|
|
}
|
|
rendered = json.dumps(report, indent=2, sort_keys=True)
|
|
if args.output:
|
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
args.output.write_text(rendered + "\n", encoding="utf-8")
|
|
else:
|
|
print(rendered)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|