Skip to content

Commit b5b4230

Browse files
committed
fix(transforms): restore tracing state after exceptions
Signed-off-by: kyinhub <kevinpyin@gmail.com>
1 parent 3ee058b commit b5b4230

2 files changed

Lines changed: 15 additions & 2 deletions

File tree

monai/transforms/inverse.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -405,8 +405,10 @@ def trace_transform(self, to_trace: bool):
405405
"""Temporarily set the tracing status of a transform with a context manager."""
406406
prev = self.tracing
407407
self.tracing = to_trace
408-
yield
409-
self.tracing = prev
408+
try:
409+
yield
410+
finally:
411+
self.tracing = prev
410412

411413

412414
class InvertibleTransform(TraceableTransform, InvertibleTrait):

tests/transforms/inverse/test_traceable_transform.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,17 @@ def pop(self, data):
2929

3030
class TestTraceable(unittest.TestCase):
3131

32+
def test_trace_transform_restores_state_after_exception(self):
33+
transform = _TraceTest()
34+
transform.tracing = True
35+
36+
with self.assertRaisesRegex(RuntimeError, "expected failure"):
37+
with transform.trace_transform(False):
38+
self.assertFalse(transform.tracing)
39+
raise RuntimeError("expected failure")
40+
41+
self.assertTrue(transform.tracing)
42+
3243
def test_default(self):
3344
expected_key = "_transforms"
3445
a = _TraceTest()

0 commit comments

Comments
 (0)