From 7dba3516c2eda47b0458f38d8332db409caf3df6 Mon Sep 17 00:00:00 2001 From: Patrick Kidger <33688385+patrick-kidger@users.noreply.github.com> Date: Thu, 29 Dec 2022 17:46:25 +0000 Subject: [PATCH] Fixed working with the new (unreleased) version of typeguard (#53) --- jaxtyping/__init__.py | 2 +- jaxtyping/pytree_type.py | 10 +++++++++- test/helpers.py | 23 +++++++++++++++++++---- 3 files changed, 29 insertions(+), 6 deletions(-) diff --git a/jaxtyping/__init__.py b/jaxtyping/__init__.py index 06578ab..74f65d1 100644 --- a/jaxtyping/__init__.py +++ b/jaxtyping/__init__.py @@ -89,4 +89,4 @@ elif has_jax: del has_jax -__version__ = "0.2.9" +__version__ = "0.2.10" diff --git a/jaxtyping/pytree_type.py b/jaxtyping/pytree_type.py index 8dfe3dd..a10aa73 100644 --- a/jaxtyping/pytree_type.py +++ b/jaxtyping/pytree_type.py @@ -59,6 +59,14 @@ class _MetaPyTree(type): return out +try: + # new typeguard + _TypeCheckError = (TypeError, typeguard.TypeCheckError) +except AttributeError: + # old typeguard + _TypeCheckError = TypeError + + class _MetaSubscriptPyTree(type): def __call__(self, *args, **kwargs): raise RuntimeError("PyTree cannot be instantiated") @@ -77,7 +85,7 @@ class _MetaSubscriptPyTree(type): def is_leaftype(x): try: accepts_leaftype(x) - except TypeError: + except _TypeCheckError: return False else: return True diff --git a/test/helpers.py b/test/helpers.py index 3829bbf..e0fe431 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -18,16 +18,31 @@ # CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. import equinox as eqx +import typeguard +ParamError = [] +ReturnError = [] +ParamError.append(TypeError) # old typeguard +ReturnError.append(TypeError) # old typeguard + +try: + # new typeguard + ParamError.append(typeguard.TypeCheckError) + ReturnError.append(typeguard.TypeCheckError) +except AttributeError: + pass + try: import beartype except ImportError: - ParamError = TypeError - ReturnError = TypeError + pass else: - ParamError = (TypeError, beartype.roar.BeartypeCallHintParamViolation) - ReturnError = (TypeError, beartype.roar.BeartypeCallHintReturnViolation) + ParamError.append(beartype.roar.BeartypeCallHintParamViolation) + ReturnError.append(beartype.roar.BeartypeCallHintReturnViolation) + +ParamError = tuple(ParamError) +ReturnError = tuple(ReturnError) @eqx.filter_jit