From 926dc53856c6fad68805b31da0f2f34b6f538348 Mon Sep 17 00:00:00 2001 From: Patrick Kidger <33688385+patrick-kidger@users.noreply.github.com> Date: Wed, 14 Jun 2023 10:42:14 -0700 Subject: [PATCH] Have PyTreeDef appear correctly in docs --- jaxtyping/__init__.py | 2 ++ 1 file changed, 2 insertions(+) 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