h3-blackwell-runtime/tools/inspect_safetensors.py

41 lines
1.3 KiB
Python
Raw Normal View History

2026-08-12 14:12:42 +07:00
"""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()