diff --git a/monai/transforms/inverse.py b/monai/transforms/inverse.py index f250fdfaf6..16bc01946c 100644 --- a/monai/transforms/inverse.py +++ b/monai/transforms/inverse.py @@ -402,11 +402,20 @@ def pop_transform(self, data, key: Hashable = None, check: bool = True): @contextmanager def trace_transform(self, to_trace: bool): - """Temporarily set the tracing status of a transform with a context manager.""" + """Temporarily set the tracing status of a transform. + + The previous tracing state is restored when the context exits normally + or because of an exception. + + Args: + to_trace: tracing state to use within the context. + """ prev = self.tracing self.tracing = to_trace - yield - self.tracing = prev + try: + yield + finally: + self.tracing = prev class InvertibleTransform(TraceableTransform, InvertibleTrait): diff --git a/tests/transforms/inverse/test_traceable_transform.py b/tests/transforms/inverse/test_traceable_transform.py index 8ee7c9e62f..cdb3d48595 100644 --- a/tests/transforms/inverse/test_traceable_transform.py +++ b/tests/transforms/inverse/test_traceable_transform.py @@ -29,6 +29,18 @@ def pop(self, data): class TestTraceable(unittest.TestCase): + def test_trace_transform_restores_state_after_exception(self): + """Verify tracing state is restored after an exception.""" + transform = _TraceTest() + transform.tracing = True + + with self.assertRaisesRegex(RuntimeError, "expected failure"): + with transform.trace_transform(False): + self.assertFalse(transform.tracing) + raise RuntimeError("expected failure") + + self.assertTrue(transform.tracing) + def test_default(self): expected_key = "_transforms" a = _TraceTest()