mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-09 11:24:55 +08:00
Threading fix (#24)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user