From ca4578ac9526ee74d6ba8159b5bec3cadd25ac20 Mon Sep 17 00:00:00 2001 From: Pradyot Ranjan <99216956+pradyotRanjan@users.noreply.github.com> Date: Thu, 13 Aug 2026 02:03:29 +0530 Subject: [PATCH] local mlx testing changes Signed-off-by: Pradyot Ranjan <99216956+pradyotRanjan@users.noreply.github.com> --- array_api_tests/dtype_helpers.py | 5 ++- array_api_tests/hypothesis_helpers.py | 8 +++- array_api_tests/test_inspection_functions.py | 7 ++++ array_api_tests/test_signatures.py | 40 +++++++++++++++++--- reporting.py | 2 +- 5 files changed, 53 insertions(+), 9 deletions(-) diff --git a/array_api_tests/dtype_helpers.py b/array_api_tests/dtype_helpers.py index 9fe2c7b1..e2d595ad 100644 --- a/array_api_tests/dtype_helpers.py +++ b/array_api_tests/dtype_helpers.py @@ -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) diff --git a/array_api_tests/hypothesis_helpers.py b/array_api_tests/hypothesis_helpers.py index 9753c3ec..02d8c275 100644 --- a/array_api_tests/hypothesis_helpers.py +++ b/array_api_tests/hypothesis_helpers.py @@ -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) ] diff --git a/array_api_tests/test_inspection_functions.py b/array_api_tests/test_inspection_functions.py index ae9362b5..e89cf221 100644 --- a/array_api_tests/test_inspection_functions.py +++ b/array_api_tests/test_inspection_functions.py @@ -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") diff --git a/array_api_tests/test_signatures.py b/array_api_tests/test_signatures.py index defaedbb..dc0df977 100644 --- a/array_api_tests/test_signatures.py +++ b/array_api_tests/test_signatures.py @@ -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,], ( @@ -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): @@ -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: diff --git a/reporting.py b/reporting.py index 579aa211..8570a7c8 100644 --- a/reporting.py +++ b/reporting.py @@ -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__