diff --git a/API.md b/API.md index 7469eec..4e15298 100644 --- a/API.md +++ b/API.md @@ -144,13 +144,13 @@ from jaxtyping import install_import_hook # Plus any one of the following: # 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.bar # with the specified string, will automatically have both `@jaxtyped` import foo.bar.qux # and the specified typechecker applied to all of their functions. # 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) @@ -177,7 +177,7 @@ The import hook will automatically decorate all functions, and the `__init__` me ```python ### entry_point.py 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 ### do_stuff.py @@ -192,7 +192,7 @@ def g(x: Float32[Array, "..."]): ```python ### __init__.py 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 .another_subpackage import bar # full name is my_library_name.another_subpackage so will be hook'd. ``` diff --git a/jaxtyping/import_hook.py b/jaxtyping/import_hook.py index 070817e..ed3690f 100644 --- a/jaxtyping/import_hook.py +++ b/jaxtyping/import_hook.py @@ -55,7 +55,7 @@ from importlib.abc import MetaPathFinder from importlib.machinery import SourceFileLoader from importlib.util import cache_from_source, decode_source from inspect import isclass -from typing import Iterable, List, Optional, Tuple +from typing import Iterable, List, Optional, Tuple, Union from unittest.mock import patch @@ -70,7 +70,8 @@ def _optimized_cache_from_source(path, debug_override=None): # Version 3: now also annotating classes. # 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? - return cache_from_source(path, debug_override, optimization="jaxtyping4") + # Version 5: Added support for string-based `typechecker` argument. + return cache_from_source(path, debug_override, optimization="jaxtyping5") def _dot_lookup(*elements): @@ -80,6 +81,12 @@ def _dot_lookup(*elements): return out +def _str_lookup(string): + module = ast.parse(string) + (expr,) = module.body + return expr.value + + class _JaxtypingTransformer(ast.NodeVisitor): def __init__(self, *, typechecker) -> None: self._parents: List[ast.AST] = [] @@ -96,7 +103,7 @@ class _JaxtypingTransformer(ast.NodeVisitor): else: node.body.insert(i, ast.Import(names=[ast.alias("jaxtyping", None)])) if self._typechecker is not None: - typechecker_module, _ = self._typechecker + typechecker_module, _ = self._typechecker.split(".", 1) node.body.insert( i, ast.Import(names=[ast.alias(typechecker_module, None)]) ) @@ -112,7 +119,7 @@ class _JaxtypingTransformer(ast.NodeVisitor): if self._typechecker is None: args = [ast.Constant(None)] else: - args = [_dot_lookup(*self._typechecker)] + args = [_str_lookup(self._typechecker)] node.decorator_list.insert(0, ast.Call(func, args, keywords=[])) self._parents.append(node) self.generic_visit(node) @@ -137,7 +144,7 @@ class _JaxtypingTransformer(ast.NodeVisitor): # Place at the end of the decorator list, as decorators # frequently remove annotations from functions and we'd like to # use those annotations. - node.decorator_list.append(_dot_lookup(*self._typechecker)) + node.decorator_list.append(_str_lookup(self._typechecker)) self._parents.append(node) self.generic_visit(node) self._parents.pop() @@ -234,7 +241,7 @@ class ImportHookManager: # Deliberately no default for `typechecker` so that folks must opt-in to not having # a typechecker. def install_import_hook( - modules: Iterable[str], typechecker: Optional[Tuple[str, str]] + modules: Iterable[str], typechecker: Optional[Union[str, Tuple[str, str]]] ) -> ImportHookManager: """Automatically apply `@jaxtyped`, and optionally a type checker, to all classes and functions. @@ -246,9 +253,9 @@ def install_import_hook( - `packages`: the names of the modules in which to automatically apply `@jaxtyped` and `@typechecked`. - `typechecker`: the module and function of the typechecker you want to use, as a - 2-tuple of strings. For example `typechecker=("typeguard", "typechecked")` or - `typechecker=("beartype", "beartype")`. You may pass `typechecker=None` if you - do not want to automatically decorate with a typechecker as well. + string. For example `typechecker="typeguard.typechecked"`, or + `typechecker="beartype.beartype"`. You may pass `typechecker=None` if you do not + want to automatically decorate with a typechecker as well. 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. @@ -286,6 +293,10 @@ def install_import_hook( if isinstance(modules, str): modules = [modules] + # Support old less-flexible API. + if isinstance(typechecker, tuple): + typechecker = ".".join(typechecker) + for i, finder in enumerate(sys.meta_path): if ( isclass(finder) diff --git a/test/import_hook_tester_beartype_full.py b/test/import_hook_tester_beartype_full.py new file mode 100644 index 0000000..b77b756 --- /dev/null +++ b/test/import_hook_tester_beartype_full.py @@ -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)) diff --git a/test/import_hook_tester_beartype_old.py b/test/import_hook_tester_beartype_old.py new file mode 100644 index 0000000..b77b756 --- /dev/null +++ b/test/import_hook_tester_beartype_old.py @@ -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)) diff --git a/test/import_hook_tester_typeguard_old.py b/test/import_hook_tester_typeguard_old.py new file mode 100644 index 0000000..b77b756 --- /dev/null +++ b/test/import_hook_tester_typeguard_old.py @@ -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)) diff --git a/test/test_import_hook.py b/test/test_import_hook.py index 66c6cc6..9f35fe4 100644 --- a/test/test_import_hook.py +++ b/test/test_import_hook.py @@ -22,14 +22,35 @@ import pytest 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(): hook = install_import_hook( - "test.import_hook_tester_typeguard", ("typeguard", "typechecked") + "test.import_hook_tester_typeguard", "typeguard.typechecked" ) with hook: 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(): try: import beartype # noqa: F401 @@ -37,15 +58,27 @@ def test_import_hook_beartype(): pytest.skip("Beartype not installed") else: hook = install_import_hook( - "test.import_hook_tester_beartype", ("beartype", "beartype") + "test.import_hook_tester_beartype", "beartype.beartype" ) with hook: 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(): hook = install_import_hook( - "test.import_hook_tester_transitive", ("typeguard", "typechecked") + "test.import_hook_tester_transitive", "typeguard.typechecked" ) with hook: from . import import_hook_tester_transitive # noqa: F401 @@ -53,7 +86,7 @@ def test_import_hook_transitive(): def test_import_hook_broken_checker(): 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): from . import import_hook_tester_broken_checker # noqa: F401