diff --git a/jaxtyping/__init__.py b/jaxtyping/__init__.py index d6ecfca..992e1dd 100644 --- a/jaxtyping/__init__.py +++ b/jaxtyping/__init__.py @@ -187,9 +187,11 @@ elif has_jax: # Conveniences if typing.TYPE_CHECKING: - from jax.random import PRNGKeyArray as PRNGKeyArray - - from ._indirection import Scalar as Scalar, ScalarLike as ScalarLike + from ._indirection import ( + PRNGKeyArray as PRNGKeyArray, + Scalar as Scalar, + ScalarLike as ScalarLike, + ) elif has_jax: from ._array_types import Scalar, ScalarLike # noqa: F401 diff --git a/jaxtyping/_indirection.py b/jaxtyping/_indirection.py index 33e6ea1..2504347 100644 --- a/jaxtyping/_indirection.py +++ b/jaxtyping/_indirection.py @@ -48,5 +48,8 @@ from typing import ( Annotated as UInt64, # noqa: F401 ) -from jax import Array as Scalar # noqa: F401 +from jax import ( + Array as PRNGKeyArray, # noqa: F401 + Array as Scalar, # noqa: F401 +) from jax.typing import ArrayLike as ScalarLike # noqa: F401 diff --git a/test/test_import_hook.py b/test/test_import_hook.py index 8ad2527..29abb8d 100644 --- a/test/test_import_hook.py +++ b/test/test_import_hook.py @@ -29,7 +29,7 @@ import pytest import jaxtyping -_here = pathlib.Path(__file__).resolve().parent +_here = pathlib.Path(__file__).parent try: