Have PyTreeDef appear correctly in docs

This commit is contained in:
Patrick Kidger
2023-06-14 10:42:14 -07:00
parent 356f5b7f7b
commit 926dc53856
+2
View File
@@ -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