mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-09 11:24:55 +08:00
Added support for declaring PyTree structures.
This commit is contained in:
@@ -27,7 +27,7 @@ from typing import Any, Literal, NoReturn, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ._decorator import storage
|
||||
from ._decorator import storage_get, storage_set
|
||||
|
||||
|
||||
try:
|
||||
@@ -198,24 +198,10 @@ class _MetaAbstractArray(type):
|
||||
if not in_dtypes:
|
||||
return False
|
||||
|
||||
no_temp_memo = hasattr(storage, "memo_stack") and len(storage.memo_stack) != 0
|
||||
|
||||
if no_temp_memo:
|
||||
single_memo, variadic_memo = storage.memo_stack[-1]
|
||||
# Make a copy so we don't mutate the original memo during the shape check.
|
||||
single_memo = single_memo.copy()
|
||||
variadic_memo = variadic_memo.copy()
|
||||
else:
|
||||
# `isinstance` happening outside any @jaxtyped decorators, e.g. at the
|
||||
# global scope. In this case just create a temporary memo, since we're not
|
||||
# going to be comparing against any stored values anyway.
|
||||
single_memo = {}
|
||||
variadic_memo = {}
|
||||
|
||||
single_memo, variadic_memo, pytree_memo = storage_get()
|
||||
if cls._check_shape(obj, single_memo, variadic_memo):
|
||||
# We update the memo every time we successfully pass a shape check
|
||||
if no_temp_memo:
|
||||
storage.memo_stack[-1] = single_memo, variadic_memo
|
||||
storage_set(single_memo, variadic_memo, pytree_memo)
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
|
||||
+34
-8
@@ -35,7 +35,33 @@ else:
|
||||
traceback_util.register_exclusion(__file__)
|
||||
|
||||
|
||||
storage = threading.local()
|
||||
_storage = threading.local()
|
||||
|
||||
|
||||
def _no_temp_memo():
|
||||
return hasattr(_storage, "memo_stack") and len(_storage.memo_stack) != 0
|
||||
|
||||
|
||||
def storage_get():
|
||||
if _no_temp_memo():
|
||||
single_memo, variadic_memo, pytree_memo = _storage.memo_stack[-1]
|
||||
# Make a copy so we don't mutate the original memo during the shape check.
|
||||
single_memo = single_memo.copy()
|
||||
variadic_memo = variadic_memo.copy()
|
||||
pytree_memo = pytree_memo.copy()
|
||||
else:
|
||||
# `isinstance` happening outside any @jaxtyped decorators, e.g. at the
|
||||
# global scope. In this case just create a temporary memo, since we're not
|
||||
# going to be comparing against any stored values anyway.
|
||||
single_memo = {}
|
||||
variadic_memo = {}
|
||||
pytree_memo = {}
|
||||
return single_memo, variadic_memo, pytree_memo
|
||||
|
||||
|
||||
def storage_set(single_memo, variadic_memo, pytree_memo):
|
||||
if _no_temp_memo():
|
||||
_storage.memo_stack[-1] = single_memo, variadic_memo, pytree_memo
|
||||
|
||||
|
||||
_jaxtyped_fns = weakref.WeakSet()
|
||||
@@ -125,10 +151,10 @@ def jaxtyped(fn):
|
||||
@ft.wraps(fn)
|
||||
def wrapped_fn(*args, **kwargs):
|
||||
try:
|
||||
memo_stack = storage.memo_stack
|
||||
memo_stack = _storage.memo_stack
|
||||
except AttributeError:
|
||||
memo_stack = storage.memo_stack = []
|
||||
memo_stack.append(({}, {}))
|
||||
memo_stack = _storage.memo_stack = []
|
||||
memo_stack.append(({}, {}, {}))
|
||||
try:
|
||||
return fn(*args, **kwargs)
|
||||
finally:
|
||||
@@ -141,13 +167,13 @@ def jaxtyped(fn):
|
||||
class _JaxtypingContext:
|
||||
def __enter__(self):
|
||||
try:
|
||||
memo_stack = storage.memo_stack
|
||||
memo_stack = _storage.memo_stack
|
||||
except AttributeError:
|
||||
memo_stack = storage.memo_stack = []
|
||||
memo_stack.append(({}, {}))
|
||||
memo_stack = _storage.memo_stack = []
|
||||
memo_stack.append(({}, {}, {}))
|
||||
|
||||
def __exit__(self, exc_type, exc_value, exc_tb):
|
||||
storage.memo_stack.pop()
|
||||
_storage.memo_stack.pop()
|
||||
|
||||
|
||||
@jaxtyped
|
||||
|
||||
+163
-11
@@ -24,6 +24,8 @@ from typing import Generic, TypeVar
|
||||
import jax.tree_util as jtu
|
||||
import typeguard
|
||||
|
||||
from ._decorator import storage_get, storage_set
|
||||
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
@@ -46,7 +48,7 @@ class _MetaPyTree(type):
|
||||
return True # Just `isinstance(x, PyTree)`
|
||||
|
||||
# We could use `isinstance` here but that would fail for more complicated
|
||||
# types, e.g. PyTree[Tuple[int]]. So at least internally we make a particular
|
||||
# types, e.g. PyTree[tuple[int]]. So at least internally we make a particular
|
||||
# choice of typechecker.
|
||||
#
|
||||
# Deliberately not using @jaxtyped so that we share the same `memo` as whatever
|
||||
@@ -63,8 +65,68 @@ class _MetaPyTree(type):
|
||||
else:
|
||||
return True
|
||||
|
||||
leaves = jtu.tree_leaves(obj, is_leaf=is_leaftype)
|
||||
return all(map(is_leaftype, leaves))
|
||||
leaves, structure = jtu.tree_flatten(obj, is_leaf=is_leaftype)
|
||||
single_memo, variadic_memo, pytree_memo = storage_get()
|
||||
if cls.structure is not None:
|
||||
if cls.structure.isidentifier():
|
||||
try:
|
||||
prev_structure = pytree_memo[cls.structure]
|
||||
except KeyError:
|
||||
pytree_memo[cls.structure] = structure
|
||||
else:
|
||||
if prev_structure != structure:
|
||||
return False
|
||||
else:
|
||||
named_pytree = 0
|
||||
pieces = cls.structure.split()
|
||||
if pieces[0] == "...":
|
||||
pieces = pieces[1:]
|
||||
prefix = False
|
||||
suffix = True
|
||||
elif pieces[-1] == "...":
|
||||
pieces = pieces[:-1]
|
||||
prefix = True
|
||||
suffix = False
|
||||
else:
|
||||
prefix = False
|
||||
suffix = False
|
||||
for identifier in pieces:
|
||||
try:
|
||||
prev_structure = pytree_memo[identifier]
|
||||
except KeyError as e:
|
||||
raise NameError(
|
||||
f"Cannot process composite structure '{cls.structure}' as "
|
||||
f"the structure name {identifier} has not been seen before."
|
||||
) from e
|
||||
# Not using `PyTreeDef.compose` due to JAX bug #18218.
|
||||
prev_pytree = jtu.tree_unflatten(
|
||||
prev_structure, [0] * prev_structure.num_leaves
|
||||
)
|
||||
named_pytree = jtu.tree_map(lambda _: prev_pytree, named_pytree)
|
||||
named_structure = jtu.tree_structure(named_pytree)
|
||||
if prefix:
|
||||
dummy_pytree = jtu.tree_unflatten(structure, [0] * len(leaves))
|
||||
dummy_named = jtu.tree_unflatten(
|
||||
named_structure, [0] * named_structure.num_leaves
|
||||
)
|
||||
try:
|
||||
jtu.tree_map(lambda _, __: 0, dummy_named, dummy_pytree)
|
||||
except ValueError:
|
||||
return False
|
||||
elif suffix:
|
||||
has_structure = lambda x: jtu.tree_structure(x) == named_structure
|
||||
dummy_pytree = jtu.tree_unflatten(structure, [0] * len(leaves))
|
||||
dummy_leaves = jtu.tree_leaves(dummy_pytree, is_leaf=has_structure)
|
||||
if any(not has_structure(x) for x in dummy_leaves):
|
||||
return False
|
||||
else:
|
||||
if structure != named_structure:
|
||||
return False
|
||||
for leaf_index, leaf in enumerate(leaves):
|
||||
if not is_leaftype(leaf):
|
||||
return False
|
||||
storage_set(single_memo, variadic_memo, pytree_memo)
|
||||
return True
|
||||
|
||||
# Can't return a generic (e.g. _FakePyTree[item]) because generic aliases don't do
|
||||
# the custom __instancecheck__ that we want.
|
||||
@@ -75,10 +137,54 @@ class _MetaPyTree(type):
|
||||
# has __module__ "types", e.g. we get types.PyTree[int].
|
||||
@ft.lru_cache(maxsize=None)
|
||||
def __getitem__(cls, item):
|
||||
name = str(_FakePyTree[item])
|
||||
if isinstance(item, tuple):
|
||||
if len(item) == 2:
|
||||
|
||||
class X(PyTree):
|
||||
leaftype = item
|
||||
class X(PyTree):
|
||||
leaftype = item[0]
|
||||
structure = item[1].strip()
|
||||
|
||||
if not isinstance(X.structure, str):
|
||||
raise ValueError(
|
||||
"The structure annotation `struct` in "
|
||||
"`jaxtyping.PyTree[leaftype, struct]` must be be a string, "
|
||||
f"e.g. `jaxtyping.PyTree[leaftype, 'T']`. Got '{X.structure}'."
|
||||
)
|
||||
pieces = X.structure.split()
|
||||
if len(pieces) == 0:
|
||||
raise ValueError(
|
||||
"The string `struct` in `jaxtyping.PyTree[leaftype, struct]` "
|
||||
"cannot be the empty string."
|
||||
)
|
||||
for piece_index, piece in enumerate(pieces):
|
||||
if (piece_index == 0) or (piece_index == len(pieces) - 1):
|
||||
if piece == "...":
|
||||
continue
|
||||
if not piece.isidentifier():
|
||||
raise ValueError(
|
||||
"The string `struct` in "
|
||||
"`jaxtyping.PyTree[leaftype, struct]` must be be a "
|
||||
"whitespace-separated sequence of identifiers, e.g. "
|
||||
"`jaxtyping.PyTree[leaftype, 'T']` or "
|
||||
"`jaxtyping.PyTree[leaftype, 'foo bar']`.\n"
|
||||
"(Here, 'identifier' is used in the same sense as in "
|
||||
"regular Python, i.e. a valid variable name.)\n"
|
||||
f"Got piece '{piece}' in overall structure '{X.structure}'."
|
||||
)
|
||||
name = str(_FakePyTree[item[0]])[:-1] + ', "' + item[1].strip() + '"]'
|
||||
else:
|
||||
raise ValueError(
|
||||
"The subscript `foo` in `jaxtyping.PyTree[foo]` must either be a "
|
||||
"leaf type, e.g. `PyTree[int]`, or a 2-tuple of leaf and "
|
||||
"structure, e.g. `PyTree[int, 'T']`. Received a tuple of length "
|
||||
f"{len(item)}."
|
||||
)
|
||||
else:
|
||||
name = str(_FakePyTree[item])
|
||||
|
||||
class X(PyTree):
|
||||
leaftype = item
|
||||
structure = None
|
||||
|
||||
X.__name__ = name
|
||||
X.__qualname__ = name
|
||||
@@ -107,10 +213,56 @@ else:
|
||||
PyTree.__module__ = "jaxtyping"
|
||||
PyTree.__doc__ = """Represents a PyTree.
|
||||
|
||||
Each PyTree is denoted by a type `PyTree[LeafType]`, such as `PyTree[int]` or
|
||||
`PyTree[Union[str, Float32[Array, "b c"]]]`.
|
||||
Annotations of the following sorts are supported:
|
||||
```python
|
||||
a: PyTree
|
||||
b: PyTree[LeafType]
|
||||
c: PyTree[LeafType, "T"]
|
||||
d: PyTree[LeafType, "S T"]
|
||||
e: PyTree[LeafType, "... T"]
|
||||
f: PyTree[LeafType, "T ..."]
|
||||
```
|
||||
|
||||
You can leave off the `[...]`, in which case `PyTree` is simply a suggestively-named
|
||||
alternative to `Any`.
|
||||
([By definition all types are PyTrees.](https://jax.readthedocs.io/en/latest/pytrees.html))
|
||||
These correspond to:
|
||||
|
||||
a. A plain `PyTree` can be used an annotation, in which case `PyTree` is simply a
|
||||
suggestively-named alternative to `Any`.
|
||||
([By definition all types are PyTrees.](https://jax.readthedocs.io/en/latest/pytrees.html))
|
||||
|
||||
b. `PyTree[LeafType]` denotes a PyTree all of whose leaves match `LeafType`. For
|
||||
example, `PyTree[int]` or `PyTree[Union[str, Float32[Array, "b c"]]]`.
|
||||
|
||||
c. A structure name can also be passed. In this case
|
||||
`jax.tree_util.tree_structure(...)` will be called, and bound to the structure name.
|
||||
This can be used to mark that multiple PyTrees all have the same structure:
|
||||
```python
|
||||
def f(x: PyTree[int, "T"], y: PyTree[int, "T"]):
|
||||
...
|
||||
```
|
||||
Structures are bound to names in the same way as array shape annotations, i.e.
|
||||
within the thread-local dynamic context of a [`jaxtyping.jaxtyped`][] decorator.
|
||||
|
||||
d. A composite structure can be declared. In this case the variable must have a PyTree
|
||||
structure each to the composition of multiple previously-bound PyTree structures.
|
||||
For example:
|
||||
```python
|
||||
def f(x: PyTree[int, "T"], y: PyTree[int, "S"], z: PyTree[int, "S T"]):
|
||||
...
|
||||
|
||||
x = (1, 2)
|
||||
y = {"key": 3}
|
||||
z = {"key": (4, 5)} # structure is the composition of the structures of `y` and `z`
|
||||
f(x, y, z)
|
||||
```
|
||||
When performing runtime type-checking, all the individual pieces must have already
|
||||
been bound to structures, otherwise the composite structure check will throw an error.
|
||||
|
||||
e. A structure can begin with a `...`, to denote that the lower levels of the PyTree
|
||||
must match the declared structure, but the upper levels can be arbitrary. As in the
|
||||
previous case, all named pieces must already have been seen and their structures
|
||||
bound.
|
||||
|
||||
f. A structure can end with a `...`, to denote that the PyTree must be a prefix of the
|
||||
declared structure, but the lower levels can be arbitrary. As in the previous two
|
||||
cases, all named pieces must already have been seen and their structures bound.
|
||||
""" # noqa: E501
|
||||
|
||||
@@ -193,3 +193,80 @@ def test_subclass_pytree():
|
||||
assert issubclass(x, PyTree)
|
||||
assert issubclass(y, PyTree)
|
||||
assert not issubclass(int, PyTree)
|
||||
|
||||
|
||||
def test_structure_match(typecheck):
|
||||
@jaxtyped
|
||||
@typecheck
|
||||
def f(x: PyTree[int, " T"], y: PyTree[str, " T"]):
|
||||
pass
|
||||
|
||||
f(1, "hi")
|
||||
f((3, 4, {"a": 5}), ("a", "b", {"a": "c"}))
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f(1, ("hi",))
|
||||
|
||||
|
||||
def test_structure_prefix(typecheck):
|
||||
@jaxtyped
|
||||
@typecheck
|
||||
def f(x: PyTree[int, " T"], y: PyTree[str, "T ..."]):
|
||||
pass
|
||||
|
||||
f(1, "hi")
|
||||
f((3, 4, {"a": 5}), ("a", "b", {"a": "c"}))
|
||||
f(1, ("hi",))
|
||||
f((1, 2), ({"a": "hi"}, {"a": "bye"}))
|
||||
f((1, 2), ({"a": "hi"}, {"not-a": "bye"}))
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f((1, 2), ({"a": "hi"}, {"a": "bye"}, {"a": "oh-no"}))
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f((3, 4, 5), {"a": ("hi", "bye")})
|
||||
|
||||
|
||||
def test_structure_suffix(typecheck):
|
||||
@jaxtyped
|
||||
@typecheck
|
||||
def f(x: PyTree[int, " T"], y: PyTree[str, "... T"]):
|
||||
pass
|
||||
|
||||
f(1, "hi")
|
||||
f((3, 4, {"a": 5}), ("a", "b", {"a": "c"}))
|
||||
f(1, ("hi",))
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f((3, 4), {"a": (1, 2)})
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f((3, 4, 5), {"a": ("hi", "bye")})
|
||||
|
||||
|
||||
def test_structure_compose(typecheck):
|
||||
@jaxtyped
|
||||
@typecheck
|
||||
def f(x: PyTree[int, " T"], y: PyTree[int, " S"], z: PyTree[str, "S T"]):
|
||||
pass
|
||||
|
||||
f(1, 2, "hi")
|
||||
f((1, 2), 2, ("a", "b"))
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f((1, 2), 2, (1, 2))
|
||||
|
||||
f((1, 2), {"a": 3}, {"a": ("hi", "bye")})
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
f((1, 2), {"a": 3}, ({"a": "hi"}, {"a": "bye"}))
|
||||
|
||||
@jaxtyped
|
||||
@typecheck
|
||||
def g(x: PyTree[int, " T"], y: PyTree[int, " S"], z: PyTree[str, "T S"]):
|
||||
pass
|
||||
|
||||
with pytest.raises(ParamError):
|
||||
g((1, 2), {"a": 3}, {"a": ("hi", "bye")})
|
||||
|
||||
g((1, 2), {"a": 3}, ({"a": "hi"}, {"a": "bye"}))
|
||||
|
||||
Reference in New Issue
Block a user