From f540f99b53f8a23b208e114db85d54ef5833098d Mon Sep 17 00:00:00 2001 From: Walter Simson Date: Tue, 22 Sep 2026 14:02:42 -0700 Subject: [PATCH 1/2] Explicitly restrict checkpoint loading to weights-only state --- tests/test_checkpoint_loading.py | 70 ++++++++++++++++++++++++++++++++ utils/training.py | 3 +- 2 files changed, 72 insertions(+), 1 deletion(-) create mode 100644 tests/test_checkpoint_loading.py diff --git a/tests/test_checkpoint_loading.py b/tests/test_checkpoint_loading.py new file mode 100644 index 0000000..7498e90 --- /dev/null +++ b/tests/test_checkpoint_loading.py @@ -0,0 +1,70 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU-only regression tests for the public training checkpoint format.""" + +import os +import pickle +from pathlib import Path +from tempfile import TemporaryDirectory +import unittest +from unittest.mock import patch + +import torch + +from utils.training import load_checkpoint, save_checkpoint, setup_run + + +class UnsupportedMetadata: + """An inert custom object that must never be allowlisted by the loader.""" + + +class CheckpointLoadingTests(unittest.TestCase): + def test_training_state_round_trip_and_resume(self): + torch.manual_seed(0) + model = torch.nn.Linear(2, 1) + optimizer = torch.optim.AdamW(model.parameters()) + model(torch.ones(1, 2)).sum().backward() + optimizer.step() + optimizer.zero_grad() + state = { + "model": model.state_dict(), + "optimizer": optimizer.state_dict(), + "epoch": 1, + "step": 1, + "best_val_loss": 0.25, + "phase1_ckpt": "phase1/best.ckpt", + } + with TemporaryDirectory() as directory: + run = setup_run("phase2", base_dir=directory) + checkpoint = save_checkpoint(state, run) + loaded = load_checkpoint(str(checkpoint)) + + restored = torch.nn.Linear(2, 1) + restored.load_state_dict(loaded["model"]) + restored_optimizer = torch.optim.AdamW(restored.parameters()) + restored_optimizer.load_state_dict(loaded["optimizer"]) + for key in ("epoch", "step", "best_val_loss", "phase1_ckpt"): + self.assertEqual(loaded[key], state[key]) + for value in loaded["model"].values(): + self.assertEqual(value.device.type, "cpu") + # A resumed optimizer must produce the same next update. + for network, optim in ((model, optimizer), (restored, restored_optimizer)): + network(torch.ones(1, 2)).sum().backward() + optim.step() + for actual, expected in zip(restored.parameters(), model.parameters()): + torch.testing.assert_close(actual, expected) + + def test_custom_objects_rejected_even_if_unsafe_default_requested(self): + with TemporaryDirectory() as directory: + checkpoint = Path(directory) / "unsupported.ckpt" + torch.save({"model": {}, "metadata": UnsupportedMetadata()}, checkpoint) + # Explicit weights_only=True must win over this PyTorch override. + with patch.dict(os.environ, {"TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD": "1", + "TORCH_FORCE_WEIGHTS_ONLY_LOAD": "0"}): + with self.assertRaises(pickle.UnpicklingError): + load_checkpoint(checkpoint) + + +if __name__ == "__main__": + unittest.main() diff --git a/utils/training.py b/utils/training.py index f79e0da..f148a14 100644 --- a/utils/training.py +++ b/utils/training.py @@ -58,7 +58,8 @@ def save_checkpoint(state: dict[str, Any], run: RunPaths, *, name: str = "last.c 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) def save_config_snapshot(run: RunPaths, cfg: dict[str, Any], *, name: str = "config.yaml", overwrite: bool = False) -> Path: From 192cec8a5de126bd6ae28c47127922da38406738 Mon Sep 17 00:00:00 2001 From: Walter Simson Date: Tue, 22 Sep 2026 14:05:17 -0700 Subject: [PATCH 2/2] Remove checkpoint regression tests --- tests/test_checkpoint_loading.py | 70 -------------------------------- 1 file changed, 70 deletions(-) delete mode 100644 tests/test_checkpoint_loading.py diff --git a/tests/test_checkpoint_loading.py b/tests/test_checkpoint_loading.py deleted file mode 100644 index 7498e90..0000000 --- a/tests/test_checkpoint_loading.py +++ /dev/null @@ -1,70 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 - -"""CPU-only regression tests for the public training checkpoint format.""" - -import os -import pickle -from pathlib import Path -from tempfile import TemporaryDirectory -import unittest -from unittest.mock import patch - -import torch - -from utils.training import load_checkpoint, save_checkpoint, setup_run - - -class UnsupportedMetadata: - """An inert custom object that must never be allowlisted by the loader.""" - - -class CheckpointLoadingTests(unittest.TestCase): - def test_training_state_round_trip_and_resume(self): - torch.manual_seed(0) - model = torch.nn.Linear(2, 1) - optimizer = torch.optim.AdamW(model.parameters()) - model(torch.ones(1, 2)).sum().backward() - optimizer.step() - optimizer.zero_grad() - state = { - "model": model.state_dict(), - "optimizer": optimizer.state_dict(), - "epoch": 1, - "step": 1, - "best_val_loss": 0.25, - "phase1_ckpt": "phase1/best.ckpt", - } - with TemporaryDirectory() as directory: - run = setup_run("phase2", base_dir=directory) - checkpoint = save_checkpoint(state, run) - loaded = load_checkpoint(str(checkpoint)) - - restored = torch.nn.Linear(2, 1) - restored.load_state_dict(loaded["model"]) - restored_optimizer = torch.optim.AdamW(restored.parameters()) - restored_optimizer.load_state_dict(loaded["optimizer"]) - for key in ("epoch", "step", "best_val_loss", "phase1_ckpt"): - self.assertEqual(loaded[key], state[key]) - for value in loaded["model"].values(): - self.assertEqual(value.device.type, "cpu") - # A resumed optimizer must produce the same next update. - for network, optim in ((model, optimizer), (restored, restored_optimizer)): - network(torch.ones(1, 2)).sum().backward() - optim.step() - for actual, expected in zip(restored.parameters(), model.parameters()): - torch.testing.assert_close(actual, expected) - - def test_custom_objects_rejected_even_if_unsafe_default_requested(self): - with TemporaryDirectory() as directory: - checkpoint = Path(directory) / "unsupported.ckpt" - torch.save({"model": {}, "metadata": UnsupportedMetadata()}, checkpoint) - # Explicit weights_only=True must win over this PyTorch override. - with patch.dict(os.environ, {"TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD": "1", - "TORCH_FORCE_WEIGHTS_ONLY_LOAD": "0"}): - with self.assertRaises(pickle.UnpicklingError): - load_checkpoint(checkpoint) - - -if __name__ == "__main__": - unittest.main()