diff --git a/jaxtyping/decorator.py b/jaxtyping/decorator.py index 3099007..68ceb0c 100644 --- a/jaxtyping/decorator.py +++ b/jaxtyping/decorator.py @@ -22,16 +22,19 @@ import threading storage = threading.local() -storage.memo_stack = [] def jaxtyped(fn): @ft.wraps(fn) def wrapper(*args, **kwargs): - storage.memo_stack.append(({}, {}, {})) + try: + memo_stack = storage.memo_stack + except AttributeError: + memo_stack = storage.memo_stack = [] + memo_stack.append(({}, {}, {})) try: return fn(*args, **kwargs) finally: - storage.memo_stack.pop() + memo_stack.pop() return wrapper diff --git a/test/test_threading.py b/test/test_threading.py new file mode 100644 index 0000000..24c50b3 --- /dev/null +++ b/test/test_threading.py @@ -0,0 +1,41 @@ +# Copyright (c) 2022 Google LLC +# +# Permission is hereby granted, free of charge, to any person obtaining a copy of +# this software and associated documentation files (the "Software"), to deal in +# the Software without restriction, including without limitation the rights to +# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of +# the Software, and to permit persons to whom the Software is furnished to do so, +# subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS +# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR +# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER +# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN +# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + +import threading + +import jax.numpy as jnp +from typeguard import typechecked + +from jaxtyping import Array, Float, jaxtyped + + +def test_threading(): + @jaxtyped + @typechecked + def add(x: Float[Array, "a b"], y: Float[Array, "a b"]) -> Float[Array, "a b"]: + return x + y + + def run(): + a = jnp.array([[1.0, 2.0]]) + b = jnp.array([[2.0, 3.0]]) + add(a, b) + + thread = threading.Thread(target=run) + thread.start() + thread.join()