Compare commits

..
4 Commits
Author SHA1 Message Date
Kevin P Murphy 784aa78f7c update jaxtyped decorator (#44)
* update jaxtyped decorator

* add newline character to pacify flake

* add precomit hooks

* add missing return statement

* replace typing_extensions>=4.2.0 with typing_extensions

* pin version range for typing_extensions

* set min version of typing-extensions but not max

* bump version number to 0.2.8
2022-11-13 22:25:28 -08:00
Patrick Kidger 3f877c0dbb Update array_types.py (#41) 2022-11-08 22:52:17 -08:00
Patrick Kidger 607f3c66b5 Silenced warning (#40) 2022-10-28 15:19:52 -07:00
Peter Roelants d3651ca70e NamedTuple example (#36) 2022-10-03 07:40:51 -07:00
6 changed files with 51 additions and 12 deletions
+1 -1
View File
@@ -67,4 +67,4 @@ from .import_hook import install_import_hook
from .pytree_type import PyTree
__version__ = "0.2.7"
__version__ = "0.2.8"
+6 -6
View File
@@ -383,15 +383,15 @@ class _MetaAbstractDtype(type):
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":
# In python 3.8, e.g., typing.Union lacks `__name__`.
try:
type_str = array_type.__name__
except AttributeError:
type_str = repr(array_type)
name = f"{cls.__name__}[{type_str}, '{dim_str}']"
elif _array_name_format == "array":
name = array_type.__name__
name = type_str
else:
raise ValueError(f"array_name_format {_array_name_format} not recognised")
out = _MetaAbstractArray(
+7 -1
View File
@@ -18,6 +18,7 @@
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
import functools as ft
import inspect
import threading
@@ -44,4 +45,9 @@ class _Jaxtyped:
def jaxtyped(fn):
return ft.wraps(fn)(_Jaxtyped(fn))
if inspect.isclass(fn): # allow decorators on class definitions
init = jaxtyped(fn.__init__)
fn.__init__ = init
return fn
else:
return ft.wraps(fn)(_Jaxtyped(fn))
+2 -2
View File
@@ -22,7 +22,7 @@ import typing
from typing import Generic, TYPE_CHECKING, TypeVar
from typing_extensions import Protocol
import jax
import jax.tree_util as jtu
import typeguard
@@ -83,7 +83,7 @@ class _MetaSubscriptPyTree(type):
else:
return True
leaves = jax.tree_leaves(obj, is_leaf=is_leaftype)
leaves = jtu.tree_leaves(obj, is_leaf=is_leaftype)
return all(map(is_leaftype, leaves))
+4 -1
View File
@@ -67,11 +67,14 @@ python_requires = "~=3.7"
# We use typeguard internally (in a fairly minimal way), but it's not required that
# end users make the same choice.
# For typing_extensions, we choose versions that match
# https://github.com/explosion/confection/blob/main/setup.cfg#L33 used in colab
install_requires = [
"jax>=0.3.4",
"numpy>=1.20.0",
"typeguard>=2.13.3",
"typing_extensions>=4.2.0",
"typing_extensions>=3.7.4.1",
]
entry_points = dict(pytest11=["jaxtyping = jaxtyping.pytest_plugin"])
+31 -1
View File
@@ -17,7 +17,7 @@
# 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.
from typing import Tuple, Union
from typing import NamedTuple, Tuple, Union
import equinox as eqx
import jax
@@ -155,3 +155,33 @@ def test_pytree_tuple(typecheck):
g([1, 1])
with pytest.raises(ParamError):
g([(1, 1), "hi"])
def test_pytree_namedtuple(typecheck):
class CustomNamedTuple(NamedTuple):
x: Float[jnp.ndarray, "a b"]
y: Float[jnp.ndarray, "b c"]
class OtherCustomNamedTuple(NamedTuple):
x: Float[jnp.ndarray, "a b"]
y: Float[jnp.ndarray, "b c"]
@typecheck
def g(x: PyTree[CustomNamedTuple]):
...
g(
CustomNamedTuple(
x=jax.random.normal(jax.random.PRNGKey(42), (3, 2)),
y=jax.random.normal(jax.random.PRNGKey(420), (2, 5)),
)
)
with pytest.raises(ParamError):
g(object())
with pytest.raises(ParamError):
g(
OtherCustomNamedTuple(
x=jax.random.normal(jax.random.PRNGKey(42), (3, 2)),
y=jax.random.normal(jax.random.PRNGKey(420), (2, 5)),
)
)