mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-09 11:24:55 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ee46c57e53 | ||
|
|
38be24f9c8 | ||
|
|
e718f00cc5 | ||
|
|
c232eeaa89 | ||
|
|
e03c1c329e | ||
|
|
fef81cf0a0 | ||
|
|
bf241b4e27 | ||
|
|
5600a1aac8 | ||
|
|
2b339715f9 | ||
|
|
8c86958b77 | ||
|
|
ffc56bf782 | ||
|
|
5c25da278a | ||
|
|
e2f004afd4 |
@@ -33,7 +33,8 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: "3.8"
|
python-version: "3.8"
|
||||||
test-script: |
|
test-script: |
|
||||||
python -m pip install pytest beartype equinox jaxlib
|
python -m pip install pytest beartype equinox jaxlib cloudpickle
|
||||||
|
python -m pip install torch --extra-index-url https://download.pytorch.org/whl/cpu
|
||||||
cp -r ${{ github.workspace }}/test ./test
|
cp -r ${{ github.workspace }}/test ./test
|
||||||
pytest
|
pytest
|
||||||
pypi-token: ${{ secrets.pypi_token }}
|
pypi-token: ${{ secrets.pypi_token }}
|
||||||
|
|||||||
@@ -42,7 +42,8 @@ jobs:
|
|||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: |
|
run: |
|
||||||
python -m pip install --upgrade pip
|
python -m pip install --upgrade pip
|
||||||
python -m pip install pytest wheel beartype equinox jaxlib
|
python -m pip install pytest wheel beartype equinox jaxlib cloudpickle
|
||||||
|
python -m pip install torch --extra-index-url https://download.pytorch.org/whl/cpu
|
||||||
|
|
||||||
- name: Checks with pre-commit
|
- name: Checks with pre-commit
|
||||||
uses: pre-commit/action@v2.0.3
|
uses: pre-commit/action@v2.0.3
|
||||||
|
|||||||
@@ -62,6 +62,8 @@ Float32[Array, "some_shape"]
|
|||||||
|
|
||||||
The array should typically be a `jaxtyping.Array`, which is an alias for `jax.numpy.ndarray`.
|
The array should typically be a `jaxtyping.Array`, which is an alias for `jax.numpy.ndarray`.
|
||||||
|
|
||||||
|
`jaxtyping.ArrayLike` is also available, which is an alias for `jax.typing.ArrayLike`. This is a union over JAX arrays and the builtin `bool`/`int`/`float`/`complex`.
|
||||||
|
|
||||||
But you can use other types as well. `jaxtyping` has support for JAX, NumPy, TensorFlow, and PyTorch, e.g.:
|
But you can use other types as well. `jaxtyping` has support for JAX, NumPy, TensorFlow, and PyTorch, e.g.:
|
||||||
```python
|
```python
|
||||||
Float[np.ndarray, "..."]
|
Float[np.ndarray, "..."]
|
||||||
@@ -142,13 +144,13 @@ from jaxtyping import install_import_hook
|
|||||||
# Plus any one of the following:
|
# Plus any one of the following:
|
||||||
|
|
||||||
# decorate @jaxtyped and @typeguard.typechecked
|
# decorate @jaxtyped and @typeguard.typechecked
|
||||||
with install_import_hook("foo", ("typeguard", "typechecked")):
|
with install_import_hook("foo", "typeguard.typechecked"):
|
||||||
import foo # Any module imported inside this `with` block, whose name begins
|
import foo # Any module imported inside this `with` block, whose name begins
|
||||||
import foo.bar # with the specified string, will automatically have both `@jaxtyped`
|
import foo.bar # with the specified string, will automatically have both `@jaxtyped`
|
||||||
import foo.bar.qux # and the specified typechecker applied to all of their functions.
|
import foo.bar.qux # and the specified typechecker applied to all of their functions.
|
||||||
|
|
||||||
# decorate @jaxtyped and @beartype.beartype
|
# decorate @jaxtyped and @beartype.beartype
|
||||||
with install_import_hook("foo", ("beartype", "beartype")):
|
with install_import_hook("foo", "beartype.beartype"):
|
||||||
...
|
...
|
||||||
|
|
||||||
# decorate only @jaxtyped (if you want that for some reason)
|
# decorate only @jaxtyped (if you want that for some reason)
|
||||||
@@ -175,7 +177,7 @@ The import hook will automatically decorate all functions, and the `__init__` me
|
|||||||
```python
|
```python
|
||||||
### entry_point.py
|
### entry_point.py
|
||||||
from jaxtyping import install_import_hook
|
from jaxtyping import install_import_hook
|
||||||
with install_import_hook("do_stuff", ("typeguard", "typechecked")):
|
with install_import_hook("do_stuff", "typeguard.typechecked"):
|
||||||
import do_stuff
|
import do_stuff
|
||||||
|
|
||||||
### do_stuff.py
|
### do_stuff.py
|
||||||
@@ -190,7 +192,7 @@ def g(x: Float32[Array, "..."]):
|
|||||||
```python
|
```python
|
||||||
### __init__.py
|
### __init__.py
|
||||||
from jaxtyping import install_import_hook
|
from jaxtyping import install_import_hook
|
||||||
with install_import_hook("my_library_name", ("beartype", "beartype")):
|
with install_import_hook("my_library_name", "beartype.beartype"):
|
||||||
from .subpackage import foo # full name is my_library_name.subpackage so will be hook'd
|
from .subpackage import foo # full name is my_library_name.subpackage so will be hook'd
|
||||||
from .another_subpackage import bar # full name is my_library_name.another_subpackage so will be hook'd.
|
from .another_subpackage import bar # full name is my_library_name.another_subpackage so will be hook'd.
|
||||||
```
|
```
|
||||||
|
|||||||
+2
-1
@@ -28,7 +28,8 @@ Now make your changes. Make sure to include additional tests if necessary.
|
|||||||
Next verify the tests all pass:
|
Next verify the tests all pass:
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
pip install pytest
|
pip install pytest cloudpickle
|
||||||
|
pip install torch --extra-index-url https://download.pytorch.org/whl/cpu
|
||||||
pytest
|
pytest
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|||||||
@@ -29,6 +29,10 @@ def accepts_pytree_of_arrays(x: PyTree[Float[Array, "batch c1 c2"]]):
|
|||||||
pip install jaxtyping
|
pip install jaxtyping
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Requires Python 3.8+.
|
||||||
|
|
||||||
|
JAX is an optional dependency, required for `jaxtyping.{Array, ArrayLike, PyTree}`. If JAX is not installed then these types will not be available, but you may still use jaxtyping alongside PyTorch/NumPy/etc.
|
||||||
|
|
||||||
Also install your favourite runtime type-checking package. The two most popular are [typeguard](https://github.com/agronholm/typeguard) (which exhaustively checks every argument) and [beartype](https://github.com/beartype/beartype) (which checks random pieces of arguments).
|
Also install your favourite runtime type-checking package. The two most popular are [typeguard](https://github.com/agronholm/typeguard) (which exhaustively checks every argument) and [beartype](https://github.com/beartype/beartype) (which checks random pieces of arguments).
|
||||||
|
|
||||||
## Documentation
|
## Documentation
|
||||||
@@ -49,12 +53,6 @@ Computer vision models: [Eqxvision](https://github.com/paganpasta/eqxvision).
|
|||||||
|
|
||||||
SymPy<->JAX conversion; train symbolic expressions via gradient descent: [sympy2jax](https://github.com/google/sympy2jax).
|
SymPy<->JAX conversion; train symbolic expressions via gradient descent: [sympy2jax](https://github.com/google/sympy2jax).
|
||||||
|
|
||||||
### Acknowledgements
|
|
||||||
|
|
||||||
Shape annotations + runtime type checking is inspired by [TorchTyping](https://github.com/patrick-kidger/torchtyping).
|
|
||||||
|
|
||||||
The concise syntax is partially inspired by [etils.array_types](https://github.com/google/etils/tree/main/etils/array_types).
|
|
||||||
|
|
||||||
### Disclaimer
|
### Disclaimer
|
||||||
|
|
||||||
This is not an official Google product.
|
This is not an official Google product.
|
||||||
|
|||||||
+15
-3
@@ -18,7 +18,6 @@
|
|||||||
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
import typing
|
import typing
|
||||||
import typing_extensions
|
|
||||||
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -35,6 +34,7 @@ if typing.TYPE_CHECKING:
|
|||||||
# For imports, we need to explicitly `import X as X` in order for Pyright to see
|
# For imports, we need to explicitly `import X as X` in order for Pyright to see
|
||||||
# them as public. See discussion at https://github.com/microsoft/pyright/issues/2277
|
# them as public. See discussion at https://github.com/microsoft/pyright/issues/2277
|
||||||
from jax import Array as Array
|
from jax import Array as Array
|
||||||
|
from jax.typing import ArrayLike as ArrayLike
|
||||||
elif has_jax:
|
elif has_jax:
|
||||||
if getattr(typing, "GENERATING_DOCUMENTATION", False):
|
if getattr(typing, "GENERATING_DOCUMENTATION", False):
|
||||||
|
|
||||||
@@ -42,9 +42,19 @@ elif has_jax:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
Array.__module__ = "builtins"
|
Array.__module__ = "builtins"
|
||||||
|
|
||||||
|
class ArrayLike:
|
||||||
|
pass
|
||||||
|
|
||||||
|
ArrayLike.__module__ = "builtins"
|
||||||
else:
|
else:
|
||||||
from jax import Array as Array
|
from jax import Array as Array
|
||||||
|
|
||||||
|
try:
|
||||||
|
from jax.typing import ArrayLike as ArrayLike
|
||||||
|
except (ModuleNotFoundError, ImportError):
|
||||||
|
pass
|
||||||
|
|
||||||
from .array_types import (
|
from .array_types import (
|
||||||
AbstractArray as AbstractArray,
|
AbstractArray as AbstractArray,
|
||||||
AbstractDtype as AbstractDtype,
|
AbstractDtype as AbstractDtype,
|
||||||
@@ -80,7 +90,9 @@ from .import_hook import install_import_hook as install_import_hook
|
|||||||
|
|
||||||
if typing.TYPE_CHECKING:
|
if typing.TYPE_CHECKING:
|
||||||
# Set up to deliberately confuse a static type checker.
|
# Set up to deliberately confuse a static type checker.
|
||||||
PyTree = getattr(typing, "foo" + "bar")
|
import typing_extensions
|
||||||
|
|
||||||
|
PyTree: typing_extensions.TypeAlias = getattr(typing, "foo" + "bar")
|
||||||
# What's going on with this madness?
|
# What's going on with this madness?
|
||||||
#
|
#
|
||||||
# At static-type-checking-time, we want `PyTree` to be a type for which both
|
# At static-type-checking-time, we want `PyTree` to be a type for which both
|
||||||
@@ -102,4 +114,4 @@ elif has_jax:
|
|||||||
|
|
||||||
del has_jax
|
del has_jax
|
||||||
|
|
||||||
__version__ = "0.2.12"
|
__version__ = "0.2.14"
|
||||||
|
|||||||
+223
-156
@@ -20,8 +20,17 @@
|
|||||||
import enum
|
import enum
|
||||||
import functools as ft
|
import functools as ft
|
||||||
import typing
|
import typing
|
||||||
from typing import Any, Dict, List, NoReturn, Optional, Tuple, TYPE_CHECKING, Union
|
from typing import (
|
||||||
from typing_extensions import Literal
|
Any,
|
||||||
|
Dict,
|
||||||
|
List,
|
||||||
|
Literal,
|
||||||
|
NoReturn,
|
||||||
|
Optional,
|
||||||
|
Tuple,
|
||||||
|
TYPE_CHECKING,
|
||||||
|
Union,
|
||||||
|
)
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
@@ -234,6 +243,20 @@ class _MetaAbstractArray(type):
|
|||||||
assert False
|
assert False
|
||||||
|
|
||||||
|
|
||||||
|
@ft.lru_cache(maxsize=None)
|
||||||
|
def _make_metaclass(base_metaclass):
|
||||||
|
class MetaAbstractArray(_MetaAbstractArray, base_metaclass):
|
||||||
|
pass
|
||||||
|
|
||||||
|
return MetaAbstractArray
|
||||||
|
|
||||||
|
|
||||||
|
def _check_scalar(dtype, dtypes, dims):
|
||||||
|
if len(dims) != 0:
|
||||||
|
return dims == (_anonymous_variadic_dim,)
|
||||||
|
return (_any_dtype is dtypes) or any(d.startswith(dtype) for d in dtypes)
|
||||||
|
|
||||||
|
|
||||||
class AbstractArray(metaclass=_MetaAbstractArray):
|
class AbstractArray(metaclass=_MetaAbstractArray):
|
||||||
array_type: Any
|
array_type: Any
|
||||||
dtypes: List[str]
|
dtypes: List[str]
|
||||||
@@ -241,6 +264,186 @@ class AbstractArray(metaclass=_MetaAbstractArray):
|
|||||||
index_variadic: Optional[int]
|
index_variadic: Optional[int]
|
||||||
|
|
||||||
|
|
||||||
|
_not_made = object()
|
||||||
|
|
||||||
|
|
||||||
|
@ft.lru_cache(maxsize=None)
|
||||||
|
def _make_array(array_type, dim_str, dtypes, name):
|
||||||
|
if not isinstance(dim_str, str):
|
||||||
|
raise ValueError(
|
||||||
|
"Shape specification must be a string. Axes should be separated with "
|
||||||
|
"spaces."
|
||||||
|
)
|
||||||
|
dims = []
|
||||||
|
index_variadic = None
|
||||||
|
for index, elem in enumerate(dim_str.split()):
|
||||||
|
if "," in elem:
|
||||||
|
# Common mistake
|
||||||
|
raise ValueError("Dimensions should be separated with spaces, not commas")
|
||||||
|
if elem.endswith("#"):
|
||||||
|
raise ValueError(
|
||||||
|
"As of jaxtyping v0.1.0, broadcastable dimensions are now denoted "
|
||||||
|
"with a # at the start, rather than at the end"
|
||||||
|
)
|
||||||
|
|
||||||
|
if "..." in elem:
|
||||||
|
if elem != "...":
|
||||||
|
raise ValueError(
|
||||||
|
"Anonymous multiple dimension '...' must be used on its own; "
|
||||||
|
f"got {elem}"
|
||||||
|
)
|
||||||
|
broadcastable = False
|
||||||
|
variadic = True
|
||||||
|
anonymous = True
|
||||||
|
dim_type = _DimType.named
|
||||||
|
else:
|
||||||
|
broadcastable = False
|
||||||
|
variadic = False
|
||||||
|
anonymous = False
|
||||||
|
while True:
|
||||||
|
if len(elem) == 0:
|
||||||
|
# This branch needed as just `_` is valid
|
||||||
|
break
|
||||||
|
first_char = elem[0]
|
||||||
|
if first_char == "#":
|
||||||
|
if broadcastable:
|
||||||
|
raise ValueError(
|
||||||
|
"Do not use # twice to denote broadcastability, e.g. "
|
||||||
|
"`##foo` is not allowed"
|
||||||
|
)
|
||||||
|
broadcastable = True
|
||||||
|
elem = elem[1:]
|
||||||
|
elif first_char == "*":
|
||||||
|
if variadic:
|
||||||
|
raise ValueError(
|
||||||
|
"Do not use * twice to denote accepting multiple "
|
||||||
|
"dimensions, e.g. `**foo` is not allowed"
|
||||||
|
)
|
||||||
|
variadic = True
|
||||||
|
elem = elem[1:]
|
||||||
|
elif first_char == "_":
|
||||||
|
if anonymous:
|
||||||
|
raise ValueError(
|
||||||
|
"Do not use _ twice to denote anonymity, e.g. `__foo` "
|
||||||
|
"is not allowed"
|
||||||
|
)
|
||||||
|
anonymous = True
|
||||||
|
elem = elem[1:]
|
||||||
|
else:
|
||||||
|
break
|
||||||
|
try:
|
||||||
|
elem = int(elem)
|
||||||
|
except ValueError:
|
||||||
|
if len(elem) == 0 or elem.isidentifier():
|
||||||
|
dim_type = _DimType.named
|
||||||
|
else:
|
||||||
|
dim_type = _DimType.symbolic
|
||||||
|
else:
|
||||||
|
dim_type = _DimType.fixed
|
||||||
|
|
||||||
|
if variadic:
|
||||||
|
if index_variadic is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot use multiple-dimension specifiers (`*name` or `...`) "
|
||||||
|
"more than once"
|
||||||
|
)
|
||||||
|
index_variadic = index
|
||||||
|
|
||||||
|
if dim_type is _DimType.fixed:
|
||||||
|
if variadic:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot have a fixed axis bind to multiple dimensions, e.g. "
|
||||||
|
"`*4` is not allowed"
|
||||||
|
)
|
||||||
|
if anonymous:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot have a fixed axis be anonymous, e.g. `_4` is not " "allowed"
|
||||||
|
)
|
||||||
|
elem = _FixedDim(elem, broadcastable)
|
||||||
|
elif dim_type is _DimType.named:
|
||||||
|
if anonymous:
|
||||||
|
if broadcastable:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot have a dimension be both anonymous and "
|
||||||
|
"broadcastable, e.g. `#_` is not allowed"
|
||||||
|
)
|
||||||
|
if variadic:
|
||||||
|
elem = _anonymous_variadic_dim
|
||||||
|
else:
|
||||||
|
elem = _anonymous_dim
|
||||||
|
else:
|
||||||
|
if variadic:
|
||||||
|
elem = _NamedVariadicDim(elem, broadcastable)
|
||||||
|
else:
|
||||||
|
elem = _NamedDim(elem, broadcastable)
|
||||||
|
else:
|
||||||
|
assert dim_type is _DimType.symbolic
|
||||||
|
if anonymous:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot have a symbolic dimension be anonymous, e.g. "
|
||||||
|
"`_foo+bar` is not allowed"
|
||||||
|
)
|
||||||
|
if variadic:
|
||||||
|
raise ValueError(
|
||||||
|
"Cannot have symbolic multiple-dimensions, e.g. "
|
||||||
|
"`*foo+bar` is not allowed"
|
||||||
|
)
|
||||||
|
elem = compile(elem, "<string>", "eval")
|
||||||
|
elem = _SymbolicDim(elem, broadcastable)
|
||||||
|
dims.append(elem)
|
||||||
|
dims = tuple(dims)
|
||||||
|
|
||||||
|
# Allow Python built-in numeric types.
|
||||||
|
# TODO: do something more generic than this? Should we _make all types
|
||||||
|
# that have `shape` and `dtype` attributes or something?
|
||||||
|
if array_type is bool:
|
||||||
|
if _check_scalar("bool", dtypes, dims):
|
||||||
|
return array_type
|
||||||
|
else:
|
||||||
|
return _not_made
|
||||||
|
elif array_type is int:
|
||||||
|
if _check_scalar("int", dtypes, dims):
|
||||||
|
return array_type
|
||||||
|
else:
|
||||||
|
return _not_made
|
||||||
|
elif array_type is float:
|
||||||
|
if _check_scalar("float", dtypes, dims):
|
||||||
|
return array_type
|
||||||
|
else:
|
||||||
|
return _not_made
|
||||||
|
elif array_type is complex:
|
||||||
|
if _check_scalar("complex", dtypes, dims):
|
||||||
|
return array_type
|
||||||
|
else:
|
||||||
|
return _not_made
|
||||||
|
try:
|
||||||
|
type_str = array_type.__name__
|
||||||
|
except AttributeError:
|
||||||
|
type_str = repr(array_type)
|
||||||
|
if _array_name_format == "dtype_and_shape":
|
||||||
|
name = f"{name}[{type_str}, '{dim_str}']"
|
||||||
|
elif _array_name_format == "array":
|
||||||
|
name = type_str
|
||||||
|
else:
|
||||||
|
raise ValueError(f"array_name_format {_array_name_format} not recognised")
|
||||||
|
metaclass = _make_metaclass(type(array_type))
|
||||||
|
out = metaclass(
|
||||||
|
name,
|
||||||
|
(array_type, AbstractArray),
|
||||||
|
dict(
|
||||||
|
array_type=array_type,
|
||||||
|
dtypes=dtypes,
|
||||||
|
dims=dims,
|
||||||
|
index_variadic=index_variadic,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
if getattr(typing, "GENERATING_DOCUMENTATION", False):
|
||||||
|
out.__module__ = "builtins"
|
||||||
|
else:
|
||||||
|
out.__module__ = "jaxtyping"
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
class _MetaAbstractDtype(type):
|
class _MetaAbstractDtype(type):
|
||||||
def __instancecheck__(cls, obj: Any) -> NoReturn:
|
def __instancecheck__(cls, obj: Any) -> NoReturn:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -249,8 +452,7 @@ class _MetaAbstractDtype(type):
|
|||||||
f'`jaxtyping.{cls.__name__}[jnp.ndarray, "..."]`.'
|
f'`jaxtyping.{cls.__name__}[jnp.ndarray, "..."]`.'
|
||||||
)
|
)
|
||||||
|
|
||||||
@ft.lru_cache(maxsize=None)
|
def __getitem__(cls, item: Tuple[Any, str]):
|
||||||
def __getitem__(cls, item: Tuple[Any, str]) -> _MetaAbstractArray:
|
|
||||||
if not isinstance(item, tuple) or len(item) != 2:
|
if not isinstance(item, tuple) or len(item) != 2:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"As of jaxtyping v0.2.0, type annotations must now include an explicit "
|
"As of jaxtyping v0.2.0, type annotations must now include an explicit "
|
||||||
@@ -258,156 +460,17 @@ class _MetaAbstractDtype(type):
|
|||||||
)
|
)
|
||||||
array_type, dim_str = item
|
array_type, dim_str = item
|
||||||
del item
|
del item
|
||||||
if not isinstance(dim_str, str):
|
if typing.get_origin(array_type) is typing.Union:
|
||||||
raise ValueError(
|
out = [
|
||||||
"Shape specification must be a string. Axes should be separated with "
|
_make_array(x, dim_str, cls.dtypes, cls.__name__)
|
||||||
"spaces."
|
for x in typing.get_args(array_type)
|
||||||
)
|
]
|
||||||
dims = []
|
out = tuple(x for x in out if x is not _not_made)
|
||||||
index_variadic = None
|
out = Union[out]
|
||||||
for index, elem in enumerate(dim_str.split()):
|
|
||||||
if "," in elem:
|
|
||||||
# Common mistake
|
|
||||||
raise ValueError(
|
|
||||||
"Dimensions should be separated with spaces, not commas"
|
|
||||||
)
|
|
||||||
if elem.endswith("#"):
|
|
||||||
raise ValueError(
|
|
||||||
"As of jaxtyping v0.1.0, broadcastable dimensions are now denoted "
|
|
||||||
"with a # at the start, rather than at the end"
|
|
||||||
)
|
|
||||||
|
|
||||||
if "..." in elem:
|
|
||||||
if elem != "...":
|
|
||||||
raise ValueError(
|
|
||||||
"Anonymous multiple dimension '...' must be used on its own; "
|
|
||||||
f"got {elem}"
|
|
||||||
)
|
|
||||||
broadcastable = False
|
|
||||||
variadic = True
|
|
||||||
anonymous = True
|
|
||||||
dim_type = _DimType.named
|
|
||||||
else:
|
|
||||||
broadcastable = False
|
|
||||||
variadic = False
|
|
||||||
anonymous = False
|
|
||||||
while True:
|
|
||||||
if len(elem) == 0:
|
|
||||||
# This branch needed as just `_` is valid
|
|
||||||
break
|
|
||||||
first_char = elem[0]
|
|
||||||
if first_char == "#":
|
|
||||||
if broadcastable:
|
|
||||||
raise ValueError(
|
|
||||||
"Do not use # twice to denote broadcastability, e.g. "
|
|
||||||
"`##foo` is not allowed"
|
|
||||||
)
|
|
||||||
broadcastable = True
|
|
||||||
elem = elem[1:]
|
|
||||||
elif first_char == "*":
|
|
||||||
if variadic:
|
|
||||||
raise ValueError(
|
|
||||||
"Do not use * twice to denote accepting multiple "
|
|
||||||
"dimensions, e.g. `**foo` is not allowed"
|
|
||||||
)
|
|
||||||
variadic = True
|
|
||||||
elem = elem[1:]
|
|
||||||
elif first_char == "_":
|
|
||||||
if anonymous:
|
|
||||||
raise ValueError(
|
|
||||||
"Do not use _ twice to denote anonymity, e.g. `__foo` "
|
|
||||||
"is not allowed"
|
|
||||||
)
|
|
||||||
anonymous = True
|
|
||||||
elem = elem[1:]
|
|
||||||
else:
|
|
||||||
break
|
|
||||||
try:
|
|
||||||
elem = int(elem)
|
|
||||||
except ValueError:
|
|
||||||
if len(elem) == 0 or elem.isidentifier():
|
|
||||||
dim_type = _DimType.named
|
|
||||||
else:
|
|
||||||
dim_type = _DimType.symbolic
|
|
||||||
else:
|
|
||||||
dim_type = _DimType.fixed
|
|
||||||
|
|
||||||
if variadic:
|
|
||||||
if index_variadic is not None:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot use multiple-dimension specifiers (`*name` or `...`) "
|
|
||||||
"more than once"
|
|
||||||
)
|
|
||||||
index_variadic = index
|
|
||||||
|
|
||||||
if dim_type is _DimType.fixed:
|
|
||||||
if variadic:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot have a fixed axis bind to multiple dimensions, e.g. "
|
|
||||||
"`*4` is not allowed"
|
|
||||||
)
|
|
||||||
if anonymous:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot have a fixed axis be anonymous, e.g. `_4` is not "
|
|
||||||
"allowed"
|
|
||||||
)
|
|
||||||
elem = _FixedDim(elem, broadcastable)
|
|
||||||
elif dim_type is _DimType.named:
|
|
||||||
if anonymous:
|
|
||||||
if broadcastable:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot have a dimension be both anonymous and "
|
|
||||||
"broadcastable, e.g. `#_` is not allowed"
|
|
||||||
)
|
|
||||||
if variadic:
|
|
||||||
elem = _anonymous_variadic_dim
|
|
||||||
else:
|
|
||||||
elem = _anonymous_dim
|
|
||||||
else:
|
|
||||||
if variadic:
|
|
||||||
elem = _NamedVariadicDim(elem, broadcastable)
|
|
||||||
else:
|
|
||||||
elem = _NamedDim(elem, broadcastable)
|
|
||||||
else:
|
|
||||||
assert dim_type is _DimType.symbolic
|
|
||||||
if anonymous:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot have a symbolic dimension be anonymous, e.g. "
|
|
||||||
"`_foo+bar` is not allowed"
|
|
||||||
)
|
|
||||||
if variadic:
|
|
||||||
raise ValueError(
|
|
||||||
"Cannot have symbolic multiple-dimensions, e.g. "
|
|
||||||
"`*foo+bar` is not allowed"
|
|
||||||
)
|
|
||||||
elem = compile(elem, "<string>", "eval")
|
|
||||||
elem = _SymbolicDim(elem, broadcastable)
|
|
||||||
dims.append(elem)
|
|
||||||
# In python 3.8, e.g., typing.Union lacks `__name__`.
|
|
||||||
try:
|
|
||||||
type_str = array_type.__name__
|
|
||||||
except AttributeError:
|
|
||||||
type_str = repr(array_type)
|
|
||||||
if _array_name_format == "dtype_and_shape":
|
|
||||||
name = f"{cls.__name__}[{type_str}, '{dim_str}']"
|
|
||||||
elif _array_name_format == "array":
|
|
||||||
name = type_str
|
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"array_name_format {_array_name_format} not recognised")
|
out = _make_array(array_type, dim_str, cls.dtypes, cls.__name__)
|
||||||
out = _MetaAbstractArray(
|
if out is _not_made:
|
||||||
name,
|
raise ValueError("Invalid jaxtyping type annotation.")
|
||||||
(AbstractArray,),
|
|
||||||
dict(
|
|
||||||
array_type=array_type,
|
|
||||||
dtypes=cls.dtypes,
|
|
||||||
dims=dims,
|
|
||||||
index_variadic=index_variadic,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
if getattr(typing, "GENERATING_DOCUMENTATION", False):
|
|
||||||
out.__module__ = "builtins"
|
|
||||||
else:
|
|
||||||
out.__module__ = "jaxtyping"
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
@@ -425,7 +488,9 @@ class AbstractDtype(metaclass=_MetaAbstractDtype):
|
|||||||
|
|
||||||
dtypes: Union[Literal[_any_dtype], str, List[str]] = cls.dtypes
|
dtypes: Union[Literal[_any_dtype], str, List[str]] = cls.dtypes
|
||||||
if isinstance(dtypes, str):
|
if isinstance(dtypes, str):
|
||||||
dtypes = [dtypes]
|
dtypes = (dtypes,)
|
||||||
|
elif dtypes is not _any_dtype:
|
||||||
|
dtypes = tuple(dtypes)
|
||||||
cls.dtypes = dtypes
|
cls.dtypes = dtypes
|
||||||
|
|
||||||
|
|
||||||
@@ -459,7 +524,8 @@ if TYPE_CHECKING:
|
|||||||
Annotated as UInt64,
|
Annotated as UInt64,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
_bool = "bool_"
|
_bool = "bool"
|
||||||
|
_bool_ = "bool_"
|
||||||
_uint8 = "uint8"
|
_uint8 = "uint8"
|
||||||
_uint16 = "uint16"
|
_uint16 = "uint16"
|
||||||
_uint32 = "uint32"
|
_uint32 = "uint32"
|
||||||
@@ -502,6 +568,7 @@ else:
|
|||||||
Complex64 = _make_dtype(_complex64, "Complex64")
|
Complex64 = _make_dtype(_complex64, "Complex64")
|
||||||
Complex128 = _make_dtype(_complex128, "Complex128")
|
Complex128 = _make_dtype(_complex128, "Complex128")
|
||||||
|
|
||||||
|
bools = [_bool, _bool_]
|
||||||
uints = [_uint8, _uint16, _uint32, _uint64]
|
uints = [_uint8, _uint16, _uint32, _uint64]
|
||||||
ints = [_int8, _int16, _int32, _int64]
|
ints = [_int8, _int16, _int32, _int64]
|
||||||
floats = [_bfloat16, _float16, _float32, _float64]
|
floats = [_bfloat16, _float16, _float32, _float64]
|
||||||
@@ -510,7 +577,7 @@ else:
|
|||||||
# We match NumPy's type hierarachy in what types to provide. See the diagram at
|
# We match NumPy's type hierarachy in what types to provide. See the diagram at
|
||||||
# https://numpy.org/doc/stable/reference/arrays.scalars.html#scalars
|
# https://numpy.org/doc/stable/reference/arrays.scalars.html#scalars
|
||||||
|
|
||||||
Bool = _make_dtype(_bool, "Bool")
|
Bool = _make_dtype(bools, "Bool")
|
||||||
UInt = _make_dtype(uints, "UInt")
|
UInt = _make_dtype(uints, "UInt")
|
||||||
Int = _make_dtype(ints, "Int")
|
Int = _make_dtype(ints, "Int")
|
||||||
Integer = _make_dtype(uints + ints, "Integer")
|
Integer = _make_dtype(uints + ints, "Integer")
|
||||||
|
|||||||
+15
-5
@@ -21,17 +21,24 @@ import dataclasses
|
|||||||
import functools as ft
|
import functools as ft
|
||||||
import inspect
|
import inspect
|
||||||
import threading
|
import threading
|
||||||
|
import weakref
|
||||||
|
|
||||||
|
|
||||||
storage = threading.local()
|
storage = threading.local()
|
||||||
|
|
||||||
|
|
||||||
|
_fns = weakref.WeakKeyDictionary()
|
||||||
|
|
||||||
|
|
||||||
class _Jaxtyped:
|
class _Jaxtyped:
|
||||||
def __init__(self, fn):
|
def __init__(self, fn):
|
||||||
self.fn = fn
|
# Stored externally so that it doesn't get blatted in the `ft.wraps` below by
|
||||||
|
# a function that already has a `fn` attribute.
|
||||||
|
_fns[self] = fn
|
||||||
|
|
||||||
def __get__(self, instance, owner):
|
def __get__(self, instance, owner):
|
||||||
return ft.wraps(self.fn)(_Jaxtyped(self.fn.__get__(instance, owner)))
|
fn = _fns[self]
|
||||||
|
return ft.wraps(fn)(_Jaxtyped(fn.__get__(instance, owner)))
|
||||||
|
|
||||||
def __call__(self, *args, **kwargs):
|
def __call__(self, *args, **kwargs):
|
||||||
try:
|
try:
|
||||||
@@ -39,8 +46,9 @@ class _Jaxtyped:
|
|||||||
except AttributeError:
|
except AttributeError:
|
||||||
memo_stack = storage.memo_stack = []
|
memo_stack = storage.memo_stack = []
|
||||||
memo_stack.append(({}, {}, {}))
|
memo_stack.append(({}, {}, {}))
|
||||||
|
fn = _fns[self]
|
||||||
try:
|
try:
|
||||||
return self.fn(*args, **kwargs)
|
return fn(*args, **kwargs)
|
||||||
finally:
|
finally:
|
||||||
memo_stack.pop()
|
memo_stack.pop()
|
||||||
|
|
||||||
@@ -69,8 +77,10 @@ def _jaxtyped_typechecker(typechecker):
|
|||||||
def _wrapper(kls):
|
def _wrapper(kls):
|
||||||
assert inspect.isclass(kls)
|
assert inspect.isclass(kls)
|
||||||
if dataclasses.is_dataclass(kls):
|
if dataclasses.is_dataclass(kls):
|
||||||
init = jaxtyped(typechecker(kls.__init__))
|
if type(kls.__init__) is not _Jaxtyped:
|
||||||
kls.__init__ = init
|
# Extra `if` check to work around beartype bug #211
|
||||||
|
init = jaxtyped(typechecker(kls.__init__))
|
||||||
|
kls.__init__ = init
|
||||||
return kls
|
return kls
|
||||||
|
|
||||||
return _wrapper
|
return _wrapper
|
||||||
|
|||||||
+28
-11
@@ -50,12 +50,13 @@
|
|||||||
|
|
||||||
|
|
||||||
import ast
|
import ast
|
||||||
|
import functools as ft
|
||||||
import sys
|
import sys
|
||||||
from importlib.abc import MetaPathFinder
|
from importlib.abc import MetaPathFinder
|
||||||
from importlib.machinery import SourceFileLoader
|
from importlib.machinery import SourceFileLoader
|
||||||
from importlib.util import cache_from_source, decode_source
|
from importlib.util import cache_from_source, decode_source
|
||||||
from inspect import isclass
|
from inspect import isclass
|
||||||
from typing import Iterable, List, Optional, Tuple
|
from typing import Iterable, List, Optional, Tuple, Union
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
|
||||||
@@ -64,13 +65,18 @@ def _call_with_frames_removed(f, *args, **kwargs):
|
|||||||
return f(*args, **kwargs)
|
return f(*args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def _optimized_cache_from_source(path, debug_override=None):
|
def _optimized_cache_from_source(typechecker_hash, /, path, debug_override=None):
|
||||||
# Version 2: change the position of the `@jaxtyped` decorator, so need a
|
# Version 2: change the position of the `@jaxtyped` decorator, so need a
|
||||||
# different name to avoid hitting old __pycache__.
|
# different name to avoid hitting old __pycache__.
|
||||||
# Version 3: now also annotating classes.
|
# Version 3: now also annotating classes.
|
||||||
# Version 4: I'm honestly not sure, but bumping this fixed some kind of odd error.
|
# Version 4: I'm honestly not sure, but bumping this fixed some kind of odd error.
|
||||||
# Maybe I changed something with hte classes part way through version 3?
|
# Maybe I changed something with hte classes part way through version 3?
|
||||||
return cache_from_source(path, debug_override, optimization="jaxtyping4")
|
# Version 5: Added support for string-based `typechecker` argument.
|
||||||
|
# Version 6: optimization tag now depends on `typechecker` argument, so that
|
||||||
|
# changing the typechecker will hit a different cache.
|
||||||
|
return cache_from_source(
|
||||||
|
path, debug_override, optimization=f"jaxtyping6{typechecker_hash}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _dot_lookup(*elements):
|
def _dot_lookup(*elements):
|
||||||
@@ -80,6 +86,12 @@ def _dot_lookup(*elements):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def _str_lookup(string):
|
||||||
|
module = ast.parse(string)
|
||||||
|
(expr,) = module.body
|
||||||
|
return expr.value
|
||||||
|
|
||||||
|
|
||||||
class _JaxtypingTransformer(ast.NodeVisitor):
|
class _JaxtypingTransformer(ast.NodeVisitor):
|
||||||
def __init__(self, *, typechecker) -> None:
|
def __init__(self, *, typechecker) -> None:
|
||||||
self._parents: List[ast.AST] = []
|
self._parents: List[ast.AST] = []
|
||||||
@@ -96,7 +108,7 @@ class _JaxtypingTransformer(ast.NodeVisitor):
|
|||||||
else:
|
else:
|
||||||
node.body.insert(i, ast.Import(names=[ast.alias("jaxtyping", None)]))
|
node.body.insert(i, ast.Import(names=[ast.alias("jaxtyping", None)]))
|
||||||
if self._typechecker is not None:
|
if self._typechecker is not None:
|
||||||
typechecker_module, _ = self._typechecker
|
typechecker_module, _ = self._typechecker.split(".", 1)
|
||||||
node.body.insert(
|
node.body.insert(
|
||||||
i, ast.Import(names=[ast.alias(typechecker_module, None)])
|
i, ast.Import(names=[ast.alias(typechecker_module, None)])
|
||||||
)
|
)
|
||||||
@@ -112,7 +124,7 @@ class _JaxtypingTransformer(ast.NodeVisitor):
|
|||||||
if self._typechecker is None:
|
if self._typechecker is None:
|
||||||
args = [ast.Constant(None)]
|
args = [ast.Constant(None)]
|
||||||
else:
|
else:
|
||||||
args = [_dot_lookup(*self._typechecker)]
|
args = [_str_lookup(self._typechecker)]
|
||||||
node.decorator_list.insert(0, ast.Call(func, args, keywords=[]))
|
node.decorator_list.insert(0, ast.Call(func, args, keywords=[]))
|
||||||
self._parents.append(node)
|
self._parents.append(node)
|
||||||
self.generic_visit(node)
|
self.generic_visit(node)
|
||||||
@@ -137,7 +149,7 @@ class _JaxtypingTransformer(ast.NodeVisitor):
|
|||||||
# Place at the end of the decorator list, as decorators
|
# Place at the end of the decorator list, as decorators
|
||||||
# frequently remove annotations from functions and we'd like to
|
# frequently remove annotations from functions and we'd like to
|
||||||
# use those annotations.
|
# use those annotations.
|
||||||
node.decorator_list.append(_dot_lookup(*self._typechecker))
|
node.decorator_list.append(_str_lookup(self._typechecker))
|
||||||
self._parents.append(node)
|
self._parents.append(node)
|
||||||
self.generic_visit(node)
|
self.generic_visit(node)
|
||||||
self._parents.pop()
|
self._parents.pop()
|
||||||
@@ -148,6 +160,7 @@ class _JaxtypingLoader(SourceFileLoader):
|
|||||||
def __init__(self, *args, typechecker, **kwargs):
|
def __init__(self, *args, typechecker, **kwargs):
|
||||||
super().__init__(*args, **kwargs)
|
super().__init__(*args, **kwargs)
|
||||||
self._typechecker = typechecker
|
self._typechecker = typechecker
|
||||||
|
self._typechecker_hash = str(abs(hash(self._typechecker)))
|
||||||
|
|
||||||
def source_to_code(self, data, path, *, _optimize=-1):
|
def source_to_code(self, data, path, *, _optimize=-1):
|
||||||
source = decode_source(data)
|
source = decode_source(data)
|
||||||
@@ -171,7 +184,7 @@ class _JaxtypingLoader(SourceFileLoader):
|
|||||||
# patch safe
|
# patch safe
|
||||||
with patch(
|
with patch(
|
||||||
"importlib._bootstrap_external.cache_from_source",
|
"importlib._bootstrap_external.cache_from_source",
|
||||||
_optimized_cache_from_source,
|
ft.partial(_optimized_cache_from_source, self._typechecker_hash),
|
||||||
):
|
):
|
||||||
return super().exec_module(module)
|
return super().exec_module(module)
|
||||||
|
|
||||||
@@ -234,7 +247,7 @@ class ImportHookManager:
|
|||||||
# Deliberately no default for `typechecker` so that folks must opt-in to not having
|
# Deliberately no default for `typechecker` so that folks must opt-in to not having
|
||||||
# a typechecker.
|
# a typechecker.
|
||||||
def install_import_hook(
|
def install_import_hook(
|
||||||
modules: Iterable[str], typechecker: Optional[Tuple[str, str]]
|
modules: Iterable[str], typechecker: Optional[Union[str, Tuple[str, str]]]
|
||||||
) -> ImportHookManager:
|
) -> ImportHookManager:
|
||||||
"""Automatically apply `@jaxtyped`, and optionally a type checker, to all classes
|
"""Automatically apply `@jaxtyped`, and optionally a type checker, to all classes
|
||||||
and functions.
|
and functions.
|
||||||
@@ -246,9 +259,9 @@ def install_import_hook(
|
|||||||
- `packages`: the names of the modules in which to automatically apply `@jaxtyped`
|
- `packages`: the names of the modules in which to automatically apply `@jaxtyped`
|
||||||
and `@typechecked`.
|
and `@typechecked`.
|
||||||
- `typechecker`: the module and function of the typechecker you want to use, as a
|
- `typechecker`: the module and function of the typechecker you want to use, as a
|
||||||
2-tuple of strings. For example `typechecker=("typeguard", "typechecked")` or
|
string. For example `typechecker="typeguard.typechecked"`, or
|
||||||
`typechecker=("beartype", "beartype")`. You may pass `typechecker=None` if you
|
`typechecker="beartype.beartype"`. You may pass `typechecker=None` if you do not
|
||||||
do not want to automatically decorate with a typechecker as well.
|
want to automatically decorate with a typechecker as well.
|
||||||
|
|
||||||
If the function already has any decorators on it, then both the `@jaxtyped` and the
|
If the function already has any decorators on it, then both the `@jaxtyped` and the
|
||||||
typechecker decorators will go at the bottom of the decorator list, e.g.
|
typechecker decorators will go at the bottom of the decorator list, e.g.
|
||||||
@@ -286,6 +299,10 @@ def install_import_hook(
|
|||||||
if isinstance(modules, str):
|
if isinstance(modules, str):
|
||||||
modules = [modules]
|
modules = [modules]
|
||||||
|
|
||||||
|
# Support old less-flexible API.
|
||||||
|
if isinstance(typechecker, tuple):
|
||||||
|
typechecker = ".".join(typechecker)
|
||||||
|
|
||||||
for i, finder in enumerate(sys.meta_path):
|
for i, finder in enumerate(sys.meta_path):
|
||||||
if (
|
if (
|
||||||
isclass(finder)
|
isclass(finder)
|
||||||
|
|||||||
@@ -52,4 +52,4 @@ def pytest_configure(config):
|
|||||||
)
|
)
|
||||||
raise RuntimeError(message.format(", ".join(already_imported_packages)))
|
raise RuntimeError(message.format(", ".join(already_imported_packages)))
|
||||||
|
|
||||||
install_import_hook(packages, typechecker.rsplit(".", 1))
|
install_import_hook(packages, typechecker)
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ classifiers = [
|
|||||||
"Topic :: Scientific/Engineering :: Mathematics",
|
"Topic :: Scientific/Engineering :: Mathematics",
|
||||||
]
|
]
|
||||||
|
|
||||||
python_requires = "~=3.7"
|
python_requires = "~=3.8"
|
||||||
|
|
||||||
# We use typeguard internally (in a fairly minimal way), but it's not required that
|
# We use typeguard internally (in a fairly minimal way), but it's not required that
|
||||||
# end users make the same choice.
|
# end users make the same choice.
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
# Copyright (c) 2022 Google LLC
|
||||||
|
#
|
||||||
|
# Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||||
|
# this software and associated documentation files (the "Software"), to deal in
|
||||||
|
# the Software without restriction, including without limitation the rights to
|
||||||
|
# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of
|
||||||
|
# the Software, and to permit persons to whom the Software is furnished to do so,
|
||||||
|
# subject to the following conditions:
|
||||||
|
#
|
||||||
|
# The above copyright notice and this permission notice shall be included in all
|
||||||
|
# copies or substantial portions of the Software.
|
||||||
|
#
|
||||||
|
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||||
|
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
|
||||||
|
# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
||||||
|
# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
|
||||||
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
|
import equinox as eqx
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from jaxtyping import Float32
|
||||||
|
|
||||||
|
from .helpers import ParamError
|
||||||
|
|
||||||
|
|
||||||
|
def g(x: Float32[jnp.ndarray, " b"]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
g(jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
g(jnp.array(1))
|
||||||
|
|
||||||
|
|
||||||
|
class M(eqx.Module):
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
M(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
M(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
M(1, jnp.array(1.0))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class D:
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
D(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1, jnp.array(1.0))
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
# Copyright (c) 2022 Google LLC
|
||||||
|
#
|
||||||
|
# Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||||
|
# this software and associated documentation files (the "Software"), to deal in
|
||||||
|
# the Software without restriction, including without limitation the rights to
|
||||||
|
# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of
|
||||||
|
# the Software, and to permit persons to whom the Software is furnished to do so,
|
||||||
|
# subject to the following conditions:
|
||||||
|
#
|
||||||
|
# The above copyright notice and this permission notice shall be included in all
|
||||||
|
# copies or substantial portions of the Software.
|
||||||
|
#
|
||||||
|
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||||
|
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
|
||||||
|
# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
||||||
|
# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
|
||||||
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
|
import equinox as eqx
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from jaxtyping import Float32
|
||||||
|
|
||||||
|
from .helpers import ParamError
|
||||||
|
|
||||||
|
|
||||||
|
def g(x: Float32[jnp.ndarray, " b"]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
g(jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
g(jnp.array(1))
|
||||||
|
|
||||||
|
|
||||||
|
class M(eqx.Module):
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
M(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
M(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
M(1, jnp.array(1.0))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class D:
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
D(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1, jnp.array(1.0))
|
||||||
@@ -0,0 +1,62 @@
|
|||||||
|
# Copyright (c) 2022 Google LLC
|
||||||
|
#
|
||||||
|
# Permission is hereby granted, free of charge, to any person obtaining a copy of
|
||||||
|
# this software and associated documentation files (the "Software"), to deal in
|
||||||
|
# the Software without restriction, including without limitation the rights to
|
||||||
|
# use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of
|
||||||
|
# the Software, and to permit persons to whom the Software is furnished to do so,
|
||||||
|
# subject to the following conditions:
|
||||||
|
#
|
||||||
|
# The above copyright notice and this permission notice shall be included in all
|
||||||
|
# copies or substantial portions of the Software.
|
||||||
|
#
|
||||||
|
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||||
|
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS
|
||||||
|
# FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR
|
||||||
|
# COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER
|
||||||
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
import dataclasses
|
||||||
|
|
||||||
|
import equinox as eqx
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from jaxtyping import Float32
|
||||||
|
|
||||||
|
from .helpers import ParamError
|
||||||
|
|
||||||
|
|
||||||
|
def g(x: Float32[jnp.ndarray, " b"]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
g(jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
g(jnp.array(1))
|
||||||
|
|
||||||
|
|
||||||
|
class M(eqx.Module):
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
M(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
M(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
M(1, jnp.array(1.0))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class D:
|
||||||
|
foo: int
|
||||||
|
bar: Float32[jnp.ndarray, " a"]
|
||||||
|
|
||||||
|
|
||||||
|
D(1, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1.0, jnp.array([1.0]))
|
||||||
|
with pytest.raises(ParamError):
|
||||||
|
D(1, jnp.array(1.0))
|
||||||
@@ -2,3 +2,4 @@ equinox>=0.5.3
|
|||||||
pytest>=7.0.1
|
pytest>=7.0.1
|
||||||
beartype>=0.10.4
|
beartype>=0.10.4
|
||||||
typeguard>=2.13.3
|
typeguard>=2.13.3
|
||||||
|
cloudpickle>=2.2.1
|
||||||
|
|||||||
+62
-1
@@ -17,11 +17,15 @@
|
|||||||
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
# IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN
|
||||||
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
|
from typing import get_args, get_origin, Union
|
||||||
|
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import jax.random as jr
|
import jax.random as jr
|
||||||
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
from jaxtyping import AbstractDtype, Array, Float, Float32, jaxtyped, Shaped
|
from jaxtyping import AbstractDtype, Array, ArrayLike, Float, Float32, jaxtyped, Shaped
|
||||||
|
|
||||||
from .helpers import ParamError, ReturnError
|
from .helpers import ParamError, ReturnError
|
||||||
|
|
||||||
@@ -409,3 +413,60 @@ def test_incomplete_symbolic(typecheck, getkey):
|
|||||||
x = jr.normal(getkey(), (4,))
|
x = jr.normal(getkey(), (4,))
|
||||||
with pytest.raises(NameError):
|
with pytest.raises(NameError):
|
||||||
foo(x)
|
foo(x)
|
||||||
|
|
||||||
|
|
||||||
|
def test_arraylike(typecheck, getkey):
|
||||||
|
floatlike1 = Float32[ArrayLike, ""]
|
||||||
|
floatlike2 = Float[ArrayLike, ""]
|
||||||
|
floatlike3 = Float32[ArrayLike, "4"]
|
||||||
|
|
||||||
|
assert get_origin(floatlike1) is Union
|
||||||
|
assert get_origin(floatlike2) is Union
|
||||||
|
assert get_origin(floatlike3) is Union
|
||||||
|
assert set(get_args(floatlike1)) == {
|
||||||
|
Float32[Array, ""],
|
||||||
|
Float32[np.ndarray, ""],
|
||||||
|
Float32[np.bool_, ""],
|
||||||
|
Float32[np.number, ""],
|
||||||
|
float,
|
||||||
|
}
|
||||||
|
assert set(get_args(floatlike2)) == {
|
||||||
|
Float[Array, ""],
|
||||||
|
Float[np.ndarray, ""],
|
||||||
|
Float[np.bool_, ""],
|
||||||
|
Float[np.number, ""],
|
||||||
|
float,
|
||||||
|
}
|
||||||
|
assert set(get_args(floatlike3)) == {
|
||||||
|
Float32[Array, "4"],
|
||||||
|
Float32[np.ndarray, "4"],
|
||||||
|
Float32[np.bool_, "4"],
|
||||||
|
Float32[np.number, "4"],
|
||||||
|
}
|
||||||
|
|
||||||
|
shaped1 = Shaped[ArrayLike, ""]
|
||||||
|
shaped2 = Shaped[ArrayLike, "4"]
|
||||||
|
assert get_origin(shaped1) is Union
|
||||||
|
assert get_origin(shaped2) is Union
|
||||||
|
assert set(get_args(shaped1)) == {
|
||||||
|
Shaped[Array, ""],
|
||||||
|
Shaped[np.ndarray, ""],
|
||||||
|
Shaped[np.bool_, ""],
|
||||||
|
Shaped[np.number, ""],
|
||||||
|
bool,
|
||||||
|
int,
|
||||||
|
float,
|
||||||
|
complex,
|
||||||
|
}
|
||||||
|
assert set(get_args(shaped2)) == {
|
||||||
|
Shaped[Array, "4"],
|
||||||
|
Shaped[np.ndarray, "4"],
|
||||||
|
Shaped[np.bool_, "4"],
|
||||||
|
Shaped[np.number, "4"],
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_subclass():
|
||||||
|
assert issubclass(Float[Array, ""], Array)
|
||||||
|
assert issubclass(Float[np.ndarray, ""], np.ndarray)
|
||||||
|
assert issubclass(Float[torch.Tensor, ""], torch.Tensor)
|
||||||
|
|||||||
@@ -22,14 +22,35 @@ import pytest
|
|||||||
from jaxtyping import install_import_hook
|
from jaxtyping import install_import_hook
|
||||||
|
|
||||||
|
|
||||||
|
def test_import_hook_typeguard_old():
|
||||||
|
hook = install_import_hook(
|
||||||
|
"test.import_hook_tester_typeguard_old", ("typeguard", "typechecked")
|
||||||
|
)
|
||||||
|
with hook:
|
||||||
|
from . import import_hook_tester_typeguard_old # noqa: F401
|
||||||
|
|
||||||
|
|
||||||
def test_import_hook_typeguard():
|
def test_import_hook_typeguard():
|
||||||
hook = install_import_hook(
|
hook = install_import_hook(
|
||||||
"test.import_hook_tester_typeguard", ("typeguard", "typechecked")
|
"test.import_hook_tester_typeguard", "typeguard.typechecked"
|
||||||
)
|
)
|
||||||
with hook:
|
with hook:
|
||||||
from . import import_hook_tester_typeguard # noqa: F401
|
from . import import_hook_tester_typeguard # noqa: F401
|
||||||
|
|
||||||
|
|
||||||
|
def test_import_hook_beartype_old():
|
||||||
|
try:
|
||||||
|
import beartype # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
pytest.skip("Beartype not installed")
|
||||||
|
else:
|
||||||
|
hook = install_import_hook(
|
||||||
|
"test.import_hook_tester_beartype_old", ("beartype", "beartype")
|
||||||
|
)
|
||||||
|
with hook:
|
||||||
|
from . import import_hook_tester_beartype_old # noqa: F401
|
||||||
|
|
||||||
|
|
||||||
def test_import_hook_beartype():
|
def test_import_hook_beartype():
|
||||||
try:
|
try:
|
||||||
import beartype # noqa: F401
|
import beartype # noqa: F401
|
||||||
@@ -37,15 +58,27 @@ def test_import_hook_beartype():
|
|||||||
pytest.skip("Beartype not installed")
|
pytest.skip("Beartype not installed")
|
||||||
else:
|
else:
|
||||||
hook = install_import_hook(
|
hook = install_import_hook(
|
||||||
"test.import_hook_tester_beartype", ("beartype", "beartype")
|
"test.import_hook_tester_beartype", "beartype.beartype"
|
||||||
)
|
)
|
||||||
with hook:
|
with hook:
|
||||||
from . import import_hook_tester_beartype # noqa: F401
|
from . import import_hook_tester_beartype # noqa: F401
|
||||||
|
|
||||||
|
|
||||||
|
def test_import_hook_beartype_full():
|
||||||
|
try:
|
||||||
|
import beartype # noqa: F401
|
||||||
|
except ImportError:
|
||||||
|
pytest.skip("Beartype not installed")
|
||||||
|
else:
|
||||||
|
bearchecker = "beartype.beartype(conf=beartype.BeartypeConf(strategy=beartype.BeartypeStrategy.On))" # noqa: E501
|
||||||
|
hook = install_import_hook("test.import_hook_tester_beartype_full", bearchecker)
|
||||||
|
with hook:
|
||||||
|
from . import import_hook_tester_beartype_full # noqa: F401
|
||||||
|
|
||||||
|
|
||||||
def test_import_hook_transitive():
|
def test_import_hook_transitive():
|
||||||
hook = install_import_hook(
|
hook = install_import_hook(
|
||||||
"test.import_hook_tester_transitive", ("typeguard", "typechecked")
|
"test.import_hook_tester_transitive", "typeguard.typechecked"
|
||||||
)
|
)
|
||||||
with hook:
|
with hook:
|
||||||
from . import import_hook_tester_transitive # noqa: F401
|
from . import import_hook_tester_transitive # noqa: F401
|
||||||
@@ -53,7 +86,7 @@ def test_import_hook_transitive():
|
|||||||
|
|
||||||
def test_import_hook_broken_checker():
|
def test_import_hook_broken_checker():
|
||||||
hook = install_import_hook(
|
hook = install_import_hook(
|
||||||
"test.import_hook_tester_broken_checker", ("jaxtyping", "does_not_exist")
|
"test.import_hook_tester_broken_checker", "jaxtyping.does_not_exist"
|
||||||
)
|
)
|
||||||
with hook, pytest.raises(AttributeError):
|
with hook, pytest.raises(AttributeError):
|
||||||
from . import import_hook_tester_broken_checker # noqa: F401
|
from . import import_hook_tester_broken_checker # noqa: F401
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
import cloudpickle
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from jaxtyping import AbstractArray, Array, Shaped
|
||||||
|
|
||||||
|
|
||||||
|
def test_pickle():
|
||||||
|
x = cloudpickle.dumps(Shaped[Array, ""])
|
||||||
|
y = cloudpickle.dumps(AbstractArray)
|
||||||
|
z = cloudpickle.dumps(Shaped[np.ndarray, ""])
|
||||||
|
w = cloudpickle.dumps(Shaped[torch.Tensor, ""])
|
||||||
|
cloudpickle.loads(x)
|
||||||
|
cloudpickle.loads(y)
|
||||||
|
cloudpickle.loads(z)
|
||||||
|
cloudpickle.loads(w)
|
||||||
@@ -31,8 +31,6 @@ class _ErrorableThread(threading.Thread):
|
|||||||
super().run()
|
super().run()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.exc = e
|
self.exc = e
|
||||||
finally:
|
|
||||||
del self._target, self._args, self._kwargs
|
|
||||||
|
|
||||||
def join(self, timeout=None):
|
def join(self, timeout=None):
|
||||||
super().join(timeout)
|
super().join(timeout)
|
||||||
|
|||||||
Reference in New Issue
Block a user