Explicitly restrict checkpoint loading - #4
Conversation
|
There was a problem hiding this comment.
Copilot review overview
🟢 Approval recommended
The change is straightforward; the only noted gap is a minor request for regression coverage.
Review effort: Lite
Findings: None
What changed in this PR
Restricts checkpoint loading to safe tensor/primitive deserialization and documents the supported format.
Changes:
- Sets
weights_only=Trueinload_checkpoint. - Documents restricted checkpoint behavior.
| File | Summary | Review note |
|---|---|---|
utils/training.py |
Restricts checkpoint deserialization. | Regression tests referenced in the PR description are not included. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| def load_checkpoint(path: Path | str) -> dict[str, Any]: | ||
| return torch.load(Path(path), map_location="cpu") | ||
| """Load tensor and primitive checkpoint state without arbitrary pickle objects.""" | ||
| return torch.load(Path(path), map_location="cpu", weights_only=True) |
There was a problem hiding this comment.
The new weights_only=True restriction has no committed regression coverage, despite the PR description stating that tests were added. Please add tests showing that tensor and primitive checkpoints still round-trip and that unsupported pickle objects are rejected; otherwise, this security-sensitive behavior could silently regress.
Explicitly set weights_only=True in load_checkpoint and add regression tests. Both tests pass on PyTorch 2.9.1.