mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-09-11 12:21:38 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7dba3516c2 | ||
|
|
2b1be5eb0a |
@@ -28,8 +28,6 @@ def accepts_pytree_of_arrays(x: PyTree[Float[Array, "batch c1 c2"]]):
|
|||||||
pip install jaxtyping
|
pip install jaxtyping
|
||||||
```
|
```
|
||||||
|
|
||||||
Requires JAX 0.3.4+.
|
|
||||||
|
|
||||||
Also install your favourite runtime type-checking package. The two most popular are [typeguard](https://github.com/agronholm/typeguard) (which exhaustively checks every argument) and [beartype](https://github.com/beartype/beartype) (which checks random pieces of arguments).
|
Also install your favourite runtime type-checking package. The two most popular are [typeguard](https://github.com/agronholm/typeguard) (which exhaustively checks every argument) and [beartype](https://github.com/beartype/beartype) (which checks random pieces of arguments).
|
||||||
|
|
||||||
## Documentation
|
## Documentation
|
||||||
|
|||||||
@@ -89,4 +89,4 @@ elif has_jax:
|
|||||||
|
|
||||||
del has_jax
|
del has_jax
|
||||||
|
|
||||||
__version__ = "0.2.9"
|
__version__ = "0.2.10"
|
||||||
|
|||||||
@@ -59,6 +59,14 @@ class _MetaPyTree(type):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
try:
|
||||||
|
# new typeguard
|
||||||
|
_TypeCheckError = (TypeError, typeguard.TypeCheckError)
|
||||||
|
except AttributeError:
|
||||||
|
# old typeguard
|
||||||
|
_TypeCheckError = TypeError
|
||||||
|
|
||||||
|
|
||||||
class _MetaSubscriptPyTree(type):
|
class _MetaSubscriptPyTree(type):
|
||||||
def __call__(self, *args, **kwargs):
|
def __call__(self, *args, **kwargs):
|
||||||
raise RuntimeError("PyTree cannot be instantiated")
|
raise RuntimeError("PyTree cannot be instantiated")
|
||||||
@@ -77,7 +85,7 @@ class _MetaSubscriptPyTree(type):
|
|||||||
def is_leaftype(x):
|
def is_leaftype(x):
|
||||||
try:
|
try:
|
||||||
accepts_leaftype(x)
|
accepts_leaftype(x)
|
||||||
except TypeError:
|
except _TypeCheckError:
|
||||||
return False
|
return False
|
||||||
else:
|
else:
|
||||||
return True
|
return True
|
||||||
|
|||||||
+19
-4
@@ -18,16 +18,31 @@
|
|||||||
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
# CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
|
||||||
|
|
||||||
import equinox as eqx
|
import equinox as eqx
|
||||||
|
import typeguard
|
||||||
|
|
||||||
|
|
||||||
|
ParamError = []
|
||||||
|
ReturnError = []
|
||||||
|
ParamError.append(TypeError) # old typeguard
|
||||||
|
ReturnError.append(TypeError) # old typeguard
|
||||||
|
|
||||||
|
try:
|
||||||
|
# new typeguard
|
||||||
|
ParamError.append(typeguard.TypeCheckError)
|
||||||
|
ReturnError.append(typeguard.TypeCheckError)
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
|
||||||
try:
|
try:
|
||||||
import beartype
|
import beartype
|
||||||
except ImportError:
|
except ImportError:
|
||||||
ParamError = TypeError
|
pass
|
||||||
ReturnError = TypeError
|
|
||||||
else:
|
else:
|
||||||
ParamError = (TypeError, beartype.roar.BeartypeCallHintParamViolation)
|
ParamError.append(beartype.roar.BeartypeCallHintParamViolation)
|
||||||
ReturnError = (TypeError, beartype.roar.BeartypeCallHintReturnViolation)
|
ReturnError.append(beartype.roar.BeartypeCallHintReturnViolation)
|
||||||
|
|
||||||
|
ParamError = tuple(ParamError)
|
||||||
|
ReturnError = tuple(ReturnError)
|
||||||
|
|
||||||
|
|
||||||
@eqx.filter_jit
|
@eqx.filter_jit
|
||||||
|
|||||||
Reference in New Issue
Block a user