Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions array_api_tests/dtype_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,12 +203,13 @@ def is_device_dlpack_compatible(device):
"""If device is dlpack compatible, return True, else False"""
# XXX: there seems to be no better way than try-catch for __dlpack_device__()

x = xp.empty(2, device=device)
try:
x = xp.empty(2, device=device)
x.__dlpack_device__()
except:
except Exception:
# case in point: torch.device(type="meta") raises
# ValueError: Unknown device type meta for Dlpack
# also covers libraries whose creation functions don't accept device=
return False
else:
# no exception => device is compatible (or a cuda device)
Expand Down
8 changes: 7 additions & 1 deletion array_api_tests/hypothesis_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,10 +93,16 @@ def arrays(dtype, *args, elements=None, **kwargs) -> SearchStrategy[Array]:

_dtype_categories = [(xp.bool,), dh.uint_dtypes, dh.int_dtypes, dh.real_float_dtypes, dh.complex_dtypes]
_sorted_dtypes = [d for category in _dtype_categories for d in category]
try:
_devices = xp.__array_namespace_info__().devices()
except Exception:
# __array_namespace_info__ is not available (e.g. libraries that only
# partially implement the standard), fall back to the default device
_devices = [None]
_device_dtype_pairs = [
(dtype, device)
for dtype in _sorted_dtypes
for device in xp.__array_namespace_info__().devices()
for device in _devices
if dh.is_device_dlpack_compatible(device)
if dh.is_dtype_device_compatible(dtype, device)
]
Expand Down
7 changes: 7 additions & 0 deletions array_api_tests/test_inspection_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,13 @@
from . import xp
from . import api_version

# For local testing, should be removed.
if not hasattr(xp, "__array_namespace_info__"):
pytest.skip(
"__array_namespace_info__ is not defined in the array module",
allow_module_level=True,
)

pytestmark = pytest.mark.min_version("2023.12")


Expand Down
40 changes: 35 additions & 5 deletions array_api_tests/test_signatures.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,9 +84,11 @@ def _test_inspectable_func(sig: Signature, stub_sig: Signature):

kwonly_stub_params = stub_params[len(non_kwonly_stub_params) :]
for stub_param in kwonly_stub_params:
assert (
stub_param.name in sig.parameters.keys()
), f"Argument '{stub_param.name}' missing from signature"
# revert back before pushing
if stub_param.name not in sig.parameters.keys():
if stub_param.name == "device":
pytest.skip("'device' argument not supported by this array module")
assert False, f"Argument '{stub_param.name}' missing from signature"
param = next(p for p in params if p.name == stub_param.name)
f_stub_kind = kind_to_str[stub_param.kind]
assert param.kind in [stub_param.kind, Parameter.POSITIONAL_OR_KEYWORD,], (
Expand Down Expand Up @@ -239,14 +241,33 @@ def _test_uninspectable_func(func_name: str, func: Callable, stub_sig: Signature
else:
assert param.kind in VAR_KINDS # sanity check
pytest.skip(no_arg_msg.format(param.name))
# revert back before pushing
def call(*a, **kw):
try:
func(*a, **kw)
except TypeError as e:
if "device" in kw:
kw_wo_device = {k: v for k, v in kw.items() if k != "device"}
try:
func(*a, **kw_wo_device)
except TypeError:
raise e from None
else:
pytest.skip(
f"'device' argument not supported by this array module: {e}"
)
raise

if len(posorkw_args) == 0:
func(*posargs, **kwargs)
# revert back before pushing
call(*posargs, **kwargs)
else:
posorkw_name_to_arg_pairs = list(posorkw_args.items())
for i in range(len(posorkw_name_to_arg_pairs), -1, -1):
extra_posargs = [arg for _, arg in posorkw_name_to_arg_pairs[:i]]
extra_kwargs = dict(posorkw_name_to_arg_pairs[i:])
func(*posargs, *extra_posargs, **kwargs, **extra_kwargs)
# revert back before pushing
call(*posargs, *extra_posargs, **kwargs, **extra_kwargs)


def _test_func_signature(func: Callable, stub: FunctionType, is_method=False):
Expand All @@ -264,6 +285,15 @@ def _test_func_signature(func: Callable, stub: FunctionType, is_method=False):
try:
sig = signature(func)
except ValueError:
sig = None
if sig is not None:
params = list(sig.parameters.values())
if [p.kind for p in params] == [
Parameter.VAR_POSITIONAL,
Parameter.VAR_KEYWORD,
]:
sig = None
if sig is None:
try:
_test_uninspectable_func(stub.__name__, func, stub_sig)
except Exception as e:
Expand Down
2 changes: 1 addition & 1 deletion reporting.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
raise ImportError("pytest-json-report is required to run the array API tests")

def to_json_serializable(o):
if o in dtype_to_name:
if o is not None and o in dtype_to_name:
return dtype_to_name[o]
if isinstance(o, (BuiltinFunctionType, FunctionType, type)):
return o.__name__
Expand Down