From e03c1c329e131e8b92cf8be96822bf4adbe4a19e Mon Sep 17 00:00:00 2001 From: Patrick Kidger <33688385+patrick-kidger@users.noreply.github.com> Date: Sat, 4 Mar 2023 17:29:04 +0000 Subject: [PATCH] We now have `Float[np.ndarray, ...] <: np.ndarray`. Added basic torch tests. (#68) This required quite a lot of refactoring! JAX supports virtual subclass registration (its metaclass is ABCMeta) but NumPy does not, so we have to actually subclass `np.ndarray`. Simple stuff like __base__ hacking fails due to deallocator conflicts. --- .github/workflows/release.yml | 3 +- .github/workflows/run_tests.yml | 3 +- CONTRIBUTING.md | 3 +- jaxtyping/__init__.py | 2 + jaxtyping/array_types.py | 383 ++++++++++++++++---------------- test/test_array.py | 135 ++++------- test/test_serialisation.py | 6 + 7 files changed, 252 insertions(+), 283 deletions(-) diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index ad9c799..de15d24 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -33,7 +33,8 @@ jobs: with: python-version: "3.8" 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 pytest pypi-token: ${{ secrets.pypi_token }} diff --git a/.github/workflows/run_tests.yml b/.github/workflows/run_tests.yml index 9f1a1bf..596e025 100644 --- a/.github/workflows/run_tests.yml +++ b/.github/workflows/run_tests.yml @@ -42,7 +42,8 @@ jobs: - name: Install dependencies run: | 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 uses: pre-commit/action@v2.0.3 diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 79957ad..163070a 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -28,7 +28,8 @@ Now make your changes. Make sure to include additional tests if necessary. Next verify the tests all pass: ```bash -pip install pytest +pip install pytest cloudpickle +pip install torch --extra-index-url https://download.pytorch.org/whl/cpu pytest ``` diff --git a/jaxtyping/__init__.py b/jaxtyping/__init__.py index 24b5d45..af19fa9 100644 --- a/jaxtyping/__init__.py +++ b/jaxtyping/__init__.py @@ -90,6 +90,8 @@ from .import_hook import install_import_hook as install_import_hook if typing.TYPE_CHECKING: # Set up to deliberately confuse a static type checker. + import typing_extensions + PyTree: typing_extensions.TypeAlias = getattr(typing, "foo" + "bar") # What's going on with this madness? # diff --git a/jaxtyping/array_types.py b/jaxtyping/array_types.py index 49dc63e..8a9f066 100644 --- a/jaxtyping/array_types.py +++ b/jaxtyping/array_types.py @@ -17,7 +17,6 @@ # 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 abc import enum import functools as ft import typing @@ -244,6 +243,14 @@ class _MetaAbstractArray(type): 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 False @@ -257,6 +264,186 @@ class AbstractArray(metaclass=_MetaAbstractArray): 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, "", "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): def __instancecheck__(cls, obj: Any) -> NoReturn: raise RuntimeError( @@ -265,8 +452,7 @@ class _MetaAbstractDtype(type): f'`jaxtyping.{cls.__name__}[jnp.ndarray, "..."]`.' ) - @ft.lru_cache(maxsize=None) - def __getitem__(cls, item: Tuple[Any, str]) -> _MetaAbstractArray: + def __getitem__(cls, item: Tuple[Any, str]): if not isinstance(item, tuple) or len(item) != 2: raise ValueError( "As of jaxtyping v0.2.0, type annotations must now include an explicit " @@ -274,196 +460,15 @@ class _MetaAbstractDtype(type): ) array_type, dim_str = item del item - 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, "", "eval") - elem = _SymbolicDim(elem, broadcastable) - dims.append(elem) - dims = tuple(dims) - - _not_made = object() - - def _make(x): - # 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 x is bool: - if _check_scalar("bool", cls.dtypes, dims): - return x - else: - return _not_made - elif x is int: - if _check_scalar("int", cls.dtypes, dims): - return x - else: - return _not_made - elif x is float: - if _check_scalar("float", cls.dtypes, dims): - return x - else: - return _not_made - elif x is complex: - if _check_scalar("complex", cls.dtypes, dims): - return x - else: - return _not_made - try: - type_str = x.__name__ - except AttributeError: - type_str = repr(x) - if _array_name_format == "dtype_and_shape": - name = f"{cls.__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" - ) - out = _MetaAbstractArray( - name, - (AbstractArray,), - dict( - array_type=x, - dtypes=cls.dtypes, - dims=dims, - index_variadic=index_variadic, - ), - ) - if getattr(typing, "GENERATING_DOCUMENTATION", False): - out.__module__ = "builtins" - else: - out.__module__ = "jaxtyping" - return out - if typing.get_origin(array_type) is typing.Union: - out = [_make(x) for x in typing.get_args(array_type)] + out = [ + _make_array(x, dim_str, cls.dtypes, cls.__name__) + for x in typing.get_args(array_type) + ] out = tuple(x for x in out if x is not _not_made) out = Union[out] else: - out = _make(array_type) - # So that `issubclass(Float[Array, ""], Array) == True`. - if isinstance(array_type, abc.ABCMeta): - array_type.register(out) + out = _make_array(array_type, dim_str, cls.dtypes, cls.__name__) return out diff --git a/test/test_array.py b/test/test_array.py index 7cf8a8f..9653e5b 100644 --- a/test/test_array.py +++ b/test/test_array.py @@ -23,17 +23,9 @@ import jax.numpy as jnp import jax.random as jr import numpy as np import pytest +import torch -from jaxtyping import ( - AbstractArray, - AbstractDtype, - Array, - ArrayLike, - Float, - Float32, - jaxtyped, - Shaped, -) +from jaxtyping import AbstractDtype, Array, ArrayLike, Float, Float32, jaxtyped, Shaped from .helpers import ParamError, ReturnError @@ -423,36 +415,6 @@ def test_incomplete_symbolic(typecheck, getkey): foo(x) -def _eq(x, y): - assert type(x) is set - assert type(y) is set - assert len(x) == len(y) - for xi in x: - assert any(_eq_impl(xi, yi) for yi in y) - - -def _eq_impl(x, y): - if issubclass(x, AbstractArray): - if type(x) is not type(y): - return False - if x.array_type is not y.array_type: - return False - if x.dtypes != y.dtypes: - return False - if x.index_variadic != y.index_variadic: - return False - if len(x.dims) != len(y.dims): - return False - for x_dim, y_dim in zip(x.dims, y.dims): - if type(x_dim) is not type(y_dim): - return False - if x_dim.__dict__ != y_dim.__dict__: - return False - return True - else: - return x is y - - def test_arraylike(typecheck, getkey): floatlike1 = Float32[ArrayLike, ""] floatlike2 = Float[ArrayLike, ""] @@ -461,59 +423,50 @@ def test_arraylike(typecheck, getkey): assert get_origin(floatlike1) is Union assert get_origin(floatlike2) is Union assert get_origin(floatlike3) is Union - _eq( - set(get_args(floatlike1)), - { - Float32[Array, ""], - Float32[np.ndarray, ""], - Float32[np.bool_, ""], - Float32[np.number, ""], - float, - }, - ) - _eq( - set(get_args(floatlike2)), - { - Float[Array, ""], - Float[np.ndarray, ""], - Float[np.bool_, ""], - Float[np.number, ""], - float, - }, - ) - _eq( - set(get_args(floatlike3)), - { - Float32[Array, "4"], - Float32[np.ndarray, "4"], - Float32[np.bool_, "4"], - Float32[np.number, "4"], - }, - ) + 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 - _eq( - set(get_args(shaped1)), - { - Shaped[Array, ""], - Shaped[np.ndarray, ""], - Shaped[np.bool_, ""], - Shaped[np.number, ""], - bool, - int, - float, - complex, - }, - ) - _eq( - set(get_args(shaped2)), - { - Shaped[Array, "4"], - Shaped[np.ndarray, "4"], - Shaped[np.bool_, "4"], - Shaped[np.number, "4"], - }, - ) + 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) diff --git a/test/test_serialisation.py b/test/test_serialisation.py index bb118af..f76f1bc 100644 --- a/test/test_serialisation.py +++ b/test/test_serialisation.py @@ -1,4 +1,6 @@ import cloudpickle +import numpy as np +import torch from jaxtyping import AbstractArray, Array, Shaped @@ -6,5 +8,9 @@ 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)