Object Detection
mmdet
mmdetection
lung-nodule
medical-imaging
LungNoduleAgent / DetectionModel /tools /sanitize_checkpoint.py
YangC777's picture
Add sanitized Nodule Spotter detection models
f49837b
Raw
History Blame Contribute Delete
4.48 kB
#!/usr/bin/env python3
"""Create an inference-only MMEngine checkpoint from a trusted training checkpoint.
The source checkpoints used for this release contain MMEngine runtime objects,
the full training configuration, optimizer/runtime state, and local filesystem
paths. This script loads those trusted files with an explicit safe-global
allowlist and writes only:
* ``state_dict`` tensors on CPU
* minimal, non-sensitive dataset metadata
* the source epoch and a format version
Never use this script on an untrusted checkpoint.
"""
from __future__ import annotations
import argparse
import builtins
import hashlib
import zipfile
from collections import OrderedDict
from pathlib import Path
import numpy as np
import torch
from mmengine.logging.history_buffer import HistoryBuffer
_DTYPES = (
np.float16,
np.float32,
np.float64,
np.int8,
np.int16,
np.int32,
np.int64,
np.uint8,
np.uint16,
np.uint32,
np.uint64,
np.bool_,
)
_SAFE_GLOBALS = [
HistoryBuffer,
(np._core.multiarray._reconstruct, "numpy.core.multiarray._reconstruct"),
(np._core.multiarray.scalar, "numpy.core.multiarray.scalar"),
np.ndarray,
np.dtype,
*{type(np.dtype(dtype)) for dtype in _DTYPES},
builtins.getattr,
]
_FORBIDDEN_MARKERS = (
b"/data",
b"HistoryBuffer",
b"message_hub",
b"optimizer",
b"experiment_name",
)
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _pickle_payload(path: Path) -> bytes:
with zipfile.ZipFile(path) as archive:
pickle_names = [name for name in archive.namelist() if name.endswith("/data.pkl")]
if len(pickle_names) != 1:
raise RuntimeError(f"Expected one data.pkl member, found {pickle_names}")
return archive.read(pickle_names[0])
def sanitize(source: Path, destination: Path) -> None:
torch.serialization.add_safe_globals(_SAFE_GLOBALS)
checkpoint = torch.load(source, map_location="cpu", weights_only=True)
if not isinstance(checkpoint, dict) or "state_dict" not in checkpoint:
raise TypeError("Source is not an MMEngine checkpoint with a state_dict")
source_state = checkpoint["state_dict"]
if not isinstance(source_state, dict):
raise TypeError("state_dict must be a mapping")
if not source_state or not all(torch.is_tensor(value) for value in source_state.values()):
raise TypeError("state_dict must contain tensors only")
state_dict = OrderedDict(
(name, tensor.detach().cpu().contiguous())
for name, tensor in source_state.items()
)
source_meta = checkpoint.get("meta") or {}
source_epoch = int(source_meta.get("epoch", 0))
clean_checkpoint = {
"state_dict": state_dict,
"meta": {
"dataset_meta": {
"classes": ("nodule",),
"palette": [(220, 20, 60)],
},
"source_epoch": source_epoch,
"format_version": 1,
},
}
destination.parent.mkdir(parents=True, exist_ok=True)
torch.save(clean_checkpoint, destination)
verified = torch.load(destination, map_location="cpu", weights_only=True)
verified_state = verified.get("state_dict", {})
if list(verified_state) != list(state_dict):
raise RuntimeError("State-dict keys changed during serialization")
for name, tensor in state_dict.items():
if not torch.equal(tensor, verified_state[name]):
raise RuntimeError(f"Tensor changed during serialization: {name}")
payload = _pickle_payload(destination)
leaked = [marker.decode("ascii") for marker in _FORBIDDEN_MARKERS if marker in payload]
if leaked:
raise RuntimeError(f"Sanitized checkpoint still contains forbidden markers: {leaked}")
print(f"source_sha256={_sha256(source)}")
print(f"output_sha256={_sha256(destination)}")
print(f"source_epoch={source_epoch}")
print(f"tensor_count={len(state_dict)}")
print(f"parameter_numel={sum(tensor.numel() for tensor in state_dict.values())}")
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("source", type=Path)
parser.add_argument("destination", type=Path)
args = parser.parse_args()
sanitize(args.source, args.destination)
if __name__ == "__main__":
main()