Skip to content
Open
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
15 changes: 12 additions & 3 deletions monai/transforms/inverse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
12 changes: 12 additions & 0 deletions tests/transforms/inverse/test_traceable_transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading