diff --git a/test/unit/test_function.py b/test/unit/test_function.py index 7292f04b52..f594a72db3 100644 --- a/test/unit/test_function.py +++ b/test/unit/test_function.py @@ -20,6 +20,13 @@ class TestFunction(unittest.TestCase): a = Tensor([1,2,3]) np.testing.assert_equal(f(a,a).numpy(), [2,4,6]) + def test_depth_restored_on_exception(self): + from tinygrad.function import _function + @function + def f(a:Tensor) -> Tensor: raise ValueError("error") + with self.assertRaises(ValueError): f(Tensor([1])) + self.assertEqual(_function.depth, 0) + def test_implicit(self): inp = Tensor([7,8,9]) @function(allow_implicit=True) diff --git a/tinygrad/function.py b/tinygrad/function.py index 2f4b868b46..7dace9694f 100644 --- a/tinygrad/function.py +++ b/tinygrad/function.py @@ -46,8 +46,10 @@ class _function(Generic[ReturnType]): # run it and do surgery later with Context(ALLOW_DEVICE_USAGE=getenv("DEVICE_IN_FUNCTION_BUG", 0)): _function.depth += 1 - ret = self.fxn(*args, **kwargs) - _function.depth -= 1 + try: + ret = self.fxn(*args, **kwargs) + finally: + _function.depth -= 1 if isinstance(ret, Tensor): uret = ret.uop elif isinstance(ret, tuple) and all(isinstance(x, Tensor) for x in ret):