mirror of
https://github.com/wassname/jaxtyping.git
synced 2026-08-21 11:16:29 +08:00
These flags are `JAXTYPING_DISABLE` and `JAXTYPING_REMOVE_TYPECHECKER_STACK`. In addition, have now added warnings when using old-style double-decorator syntax, which also serves to guard against the easy mistake of ```python @jaxtyped(typechecker) def foo(...) ``` which actually decorates the `typechecker`, not `foo`.
85 lines
2.5 KiB
Python
85 lines
2.5 KiB
Python
# 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 random
|
|
|
|
import jax.random as jr
|
|
import pytest
|
|
import typeguard
|
|
|
|
|
|
try:
|
|
import beartype
|
|
except ImportError:
|
|
|
|
def skip(*args, **kwargs):
|
|
pytest.skip("Beartype not installed")
|
|
|
|
typecheck_params = [typeguard.typechecked, skip]
|
|
else:
|
|
typecheck_params = [typeguard.typechecked, beartype.beartype]
|
|
|
|
|
|
@pytest.fixture(params=typecheck_params)
|
|
def typecheck(request):
|
|
return request.param
|
|
|
|
|
|
@pytest.fixture(params=(False, True))
|
|
def jaxtyp(request):
|
|
import jaxtyping
|
|
|
|
if request.param:
|
|
# New-style
|
|
# @jaxtyping.jaxtyped(typechecker=typechecker)
|
|
# def f(...)
|
|
return lambda typechecker: jaxtyping.jaxtyped(typechecker=typechecker)
|
|
else:
|
|
# Old-style
|
|
# @jaxtyping.jaxtyped
|
|
# @typechecker
|
|
# def f(...)
|
|
def impl(typechecker):
|
|
def decorator(fn):
|
|
with pytest.warns(match="As of jaxtyping version 0.2.24"):
|
|
return jaxtyping.jaxtyped(typechecker(fn))
|
|
|
|
return decorator
|
|
|
|
return impl
|
|
|
|
|
|
@pytest.fixture()
|
|
def getkey():
|
|
def _getkey():
|
|
# Not sure what the maximum actually is but this will do
|
|
return jr.PRNGKey(random.randint(0, 2**31 - 1))
|
|
|
|
return _getkey
|
|
|
|
|
|
@pytest.fixture(scope="module")
|
|
def beartype_or_skip():
|
|
yield pytest.importorskip("beartype")
|
|
|
|
|
|
@pytest.fixture(scope="module")
|
|
def typeguard_or_skip():
|
|
yield pytest.importorskip("typeguard")
|