From 55a0e954beaae323c4d18342c6473430d0732430 Mon Sep 17 00:00:00 2001 From: kyinhub Date: Sat, 25 Jul 2026 11:05:56 -0700 Subject: [PATCH] fix(transforms): restore tracing state after exceptions Signed-off-by: kyinhub --- monai/transforms/inverse.py | 15 ++++++++++++--- .../inverse/test_traceable_transform.py | 12 ++++++++++++ 2 files changed, 24 insertions(+), 3 deletions(-) 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()