Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 32 additions & 1 deletion smauglab/trainers/nnUNetTrainerDAExt.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
renaming it would make several hundred trained models unloadable.
"""

import filecmp
import importlib
import os
import shutil
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -116,7 +147,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`
Expand Down
84 changes: 84 additions & 0 deletions unit_tests/test_trainer_provenance.py
Original file line number Diff line number Diff line change
@@ -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}})
Loading