From 1c64a329f31cef1e18baee283c56c5f1e1a90e1d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Hendrik=20M=C3=B6ller?= Date: Thu, 24 Sep 2026 09:47:20 +0000 Subject: [PATCH] fix: keep the config record when a run is resumed with a different one The trainer copied `SMAUGLAB_PARAMS_JSON` to `transform_params_used_for_training.json` unconditionally, so resuming a run with a different config replaced the record: the file then described a config only the later epochs saw, while still claiming to describe the run. That file is the only provenance a finished run carries. An identical config is still a no-op. A differing one leaves the original in place, writes itself beside it under a numbered name, and warns naming both. Co-Authored-By: Claude Opus 5 --- smauglab/trainers/nnUNetTrainerDAExt.py | 33 +++++++++- unit_tests/test_trainer_provenance.py | 84 +++++++++++++++++++++++++ 2 files changed, 116 insertions(+), 1 deletion(-) create mode 100644 unit_tests/test_trainer_provenance.py diff --git a/smauglab/trainers/nnUNetTrainerDAExt.py b/smauglab/trainers/nnUNetTrainerDAExt.py index ffc00a9..5ea62dd 100644 --- a/smauglab/trainers/nnUNetTrainerDAExt.py +++ b/smauglab/trainers/nnUNetTrainerDAExt.py @@ -16,6 +16,7 @@ renaming it would make several hundred trained models unloadable. """ +import filecmp import importlib import os import shutil @@ -75,6 +76,36 @@ def resolve_config_path() -> str: return str(importlib.resources.files(configs) / DEFAULT_CONFIG) +def _record_config(source: str, destination: str) -> None: + """Write down which config this run was trained with. + + Overwriting unconditionally destroys the record whenever a run is resumed + with a different SMAUGLAB_PARAMS_JSON: the earlier epochs were trained with + the old config and nothing says so any more. The file is the only provenance + a finished run carries, so a differing one is kept and the new one written + beside it under a numbered name, with a warning naming both. + """ + if not os.path.exists(destination): + shutil.copy(source, destination) + return + + if filecmp.cmp(source, destination, shallow=False): + return + + index = 1 + root, extension = os.path.splitext(destination) + while os.path.exists(f"{root}_{index}{extension}"): + index += 1 + kept = f"{root}_{index}{extension}" + shutil.copy(source, kept) + warnings.warn( + f"This run was previously trained with a different SmaugLab config. " + f"{destination} is unchanged and describes the earlier epochs; the config now in use " + f"was written to {kept}.", + stacklevel=2, + ) + + def _has_gpu_augmentations(config) -> bool: """Whether this config asks for anything on the GPU side.""" return bool(config.names(Backend.GPU)) @@ -113,7 +144,7 @@ def __init__(self, plans: dict, configuration: str, fold: int, dataset_json: dic f" CPU: {len(config.names(Backend.CPU))} augmentations, GPU: {len(config.names(Backend.GPU))}, mode: {config.pipeline_mode().value}" ) - shutil.copy(json_path, os.path.join(self.output_folder, "transform_params_used_for_training.json")) + _record_config(json_path, os.path.join(self.output_folder, "transform_params_used_for_training.json")) # A non-finite loss is otherwise invisible: `GradScaler` skips the step without a # word, so the run simply stops learning and the only symptom is `train_loss nan` diff --git a/unit_tests/test_trainer_provenance.py b/unit_tests/test_trainer_provenance.py new file mode 100644 index 0000000..b80a0c0 --- /dev/null +++ b/unit_tests/test_trainer_provenance.py @@ -0,0 +1,84 @@ +"""Resuming a run must not erase what the earlier epochs were trained with. + +The trainer copies the config it was given to +`transform_params_used_for_training.json` in the run folder. It did so +unconditionally, so resuming with a different `SMAUGLAB_PARAMS_JSON` replaced +the record: the file then described a config that only the later epochs saw, +while claiming to describe the run. + +That file is the only provenance a finished run carries. A differing config is +now written beside it under a numbered name, with a warning naming both. +""" + +from __future__ import annotations + +import importlib.util +import json +import tempfile +import unittest +import warnings +from pathlib import Path + +if importlib.util.find_spec("nnunetv2") is None: + raise unittest.SkipTest("the trainer needs the nnunetv2 extra") + +from smauglab.trainers.nnUNetTrainerDAExt import _record_config # noqa: E402 + +RECORD = "transform_params_used_for_training.json" + + +def write(path: Path, payload: dict) -> Path: + path.write_text(json.dumps(payload)) + return path + + +class TestRecordConfig(unittest.TestCase): + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.addCleanup(self._tmp.cleanup) + self.run = Path(self._tmp.name) + self.destination = self.run / RECORD + + def test_the_first_run_writes_the_record(self): + source = write(self.run / "a.json", {"GPU": {}}) + + _record_config(str(source), str(self.destination)) + + self.assertEqual(json.loads(self.destination.read_text()), {"GPU": {}}) + + def test_resuming_with_the_same_config_is_a_no_op(self): + source = write(self.run / "a.json", {"GPU": {}}) + _record_config(str(source), str(self.destination)) + + with warnings.catch_warnings(): + warnings.simplefilter("error") + _record_config(str(source), str(self.destination)) + + self.assertEqual(sorted(p.name for p in self.run.glob("transform_params_*")), [RECORD]) + + def test_resuming_with_a_different_config_keeps_the_original(self): + first = write(self.run / "a.json", {"GPU": {"RandomFlipTransformGPU": {"p": 1.0}}}) + _record_config(str(first), str(self.destination)) + second = write(self.run / "b.json", {"GPU": {}}) + + with self.assertWarns(UserWarning) as caught: + _record_config(str(second), str(self.destination)) + + self.assertEqual( + json.loads(self.destination.read_text()), + {"GPU": {"RandomFlipTransformGPU": {"p": 1.0}}}, + "the record of the earlier epochs was overwritten", + ) + kept = self.run / "transform_params_used_for_training_1.json" + self.assertEqual(json.loads(kept.read_text()), {"GPU": {}}) + self.assertIn(kept.name, str(caught.warning)) + + def test_a_third_config_does_not_clobber_the_second(self): + _record_config(str(write(self.run / "a.json", {"GPU": {"a": 1}})), str(self.destination)) + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + _record_config(str(write(self.run / "b.json", {"GPU": {"b": 2}})), str(self.destination)) + _record_config(str(write(self.run / "c.json", {"GPU": {"c": 3}})), str(self.destination)) + + self.assertEqual(json.loads((self.run / "transform_params_used_for_training_1.json").read_text()), {"GPU": {"b": 2}}) + self.assertEqual(json.loads((self.run / "transform_params_used_for_training_2.json").read_text()), {"GPU": {"c": 3}})