diff --git a/jaxtyping/__init__.py b/jaxtyping/__init__.py index 5bd1371..e618daf 100644 --- a/jaxtyping/__init__.py +++ b/jaxtyping/__init__.py @@ -160,6 +160,8 @@ elif has_jax: from `jax.tree_util.tree_structure(...)`. """ + PyTreeDef.__module__ = "builtins" + else: from jax.tree_util import PyTreeDef as PyTreeDef