From e4a93ee218fd6012a42050d7a2199dc1ea18f548 Mon Sep 17 00:00:00 2001 From: Patrick Kidger <33688385+patrick-kidger@users.noreply.github.com> Date: Tue, 24 Oct 2023 10:25:42 -0700 Subject: [PATCH] Added `with jaxtyped("context"):`, for now undocumented. --- jaxtyping/_decorator.py | 14 ++++++++++++++ test/test_decorator.py | 12 +++++++++++- 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/jaxtyping/_decorator.py b/jaxtyping/_decorator.py index ca92643..d24ddfc 100644 --- a/jaxtyping/_decorator.py +++ b/jaxtyping/_decorator.py @@ -118,6 +118,8 @@ def jaxtyped(fn): else: fdel = jaxtyped(fn.fdel) return property(fget=fget, fset=fset, fdel=fdel) + elif fn == "context": + return _JaxtypingContext() else: @ft.wraps(fn) @@ -136,6 +138,18 @@ def jaxtyped(fn): return wrapped_fn +class _JaxtypingContext: + def __enter__(self): + try: + memo_stack = storage.memo_stack + except AttributeError: + memo_stack = storage.memo_stack = [] + memo_stack.append(({}, {})) + + def __exit__(self, exc_type, exc_value, exc_tb): + storage.memo_stack.pop() + + @jaxtyped def _check_dataclass_annotations(self, typechecker): for field in dataclasses.fields(self): diff --git a/test/test_decorator.py b/test/test_decorator.py index c655425..f64bf57 100644 --- a/test/test_decorator.py +++ b/test/test_decorator.py @@ -1,6 +1,8 @@ import abc -from jaxtyping import jaxtyped +import jax.random as jr + +from jaxtyping import Array, Float, jaxtyped class M(metaclass=abc.ABCMeta): @@ -75,3 +77,11 @@ def test_abstractmethod(): def test_property(): assert N().j1 == 3 assert N().j2 == 4 + + +def test_context(getkey): + a = jr.normal(getkey(), (3, 4)) + b = jr.normal(getkey(), (5,)) + with jaxtyped("context"): + assert isinstance(a, Float[Array, "foo bar"]) + assert not isinstance(b, Float[Array, "foo"])