Skip to content

Commit f7d7ec8

Browse files
committed
Generation of compatable device and dtype pairs for test_from_dlpack
Signed-off-by: Pradyot Ranjan <99216956+pradyotRanjan@users.noreply.github.com>
1 parent a109d8f commit f7d7ec8

3 files changed

Lines changed: 14 additions & 5 deletions

File tree

‎array_api_tests/dtype_helpers.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -198,6 +198,13 @@ def get_scalar_type(dtype: DataType) -> ScalarType:
198198
def is_scalar(x):
199199
return isinstance(x, (int, float, complex, bool))
200200

201+
def is_dtype_device_compatible(dtype, device):
202+
try:
203+
xp.asarray([0], dtype=dtype, device=device)
204+
except Exception:
205+
return False
206+
else:
207+
return True
201208

202209
def complex_dtype_for(dtyp):
203210
"""Complex dtype for a float or complex."""

‎array_api_tests/hypothesis_helpers.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,7 @@ def arrays(dtype, *args, elements=None, **kwargs) -> SearchStrategy[Array]:
9393

9494
_dtype_categories = [(xp.bool,), dh.uint_dtypes, dh.int_dtypes, dh.real_float_dtypes, dh.complex_dtypes]
9595
_sorted_dtypes = [d for category in _dtype_categories for d in category]
96+
_device_dtype_pairs = [(d, device) for d in _sorted_dtypes for device in xp.__array_namespace_info__().devices() if dh.is_dtype_device_compatible(d, device)]
9697

9798
def _dtypes_sorter(dtype_pair: Tuple[DataType, DataType]):
9899
dtype1, dtype2 = dtype_pair
@@ -199,6 +200,7 @@ def oneway_broadcastable_shapes(draw) -> OnewayBroadcastableShapes:
199200
# Use these instead of xps.scalar_dtypes, etc. because it skips dtypes from
200201
# ARRAY_API_TESTS_SKIP_DTYPES
201202
all_dtypes = sampled_from(_sorted_dtypes)
203+
device_dtype_pairs = sampled_from(_device_dtype_pairs)
202204
all_int_dtypes = sampled_from(dh.all_int_dtypes)
203205
int_dtypes = sampled_from(dh.int_dtypes) # signed ints
204206
uint_dtypes = sampled_from(dh.uint_dtypes)

‎array_api_tests/test_dlpack.py‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -88,22 +88,22 @@ def test_dunder_dlpack(x, copy_kw, max_version_kw, dl_device_kw, data):
8888

8989

9090
@given(
91-
x=hh.arrays(dtype=hh.all_dtypes, shape=hh.shapes(min_dims=1, max_side=2)),
91+
dtype_device_pair = hh.device_dtype_pairs,
9292
copy_kw=hh.kwargs(copy=st.booleans()),
9393
data=st.data()
9494
)
95-
def test_from_dlpack(x, copy_kw, data):
95+
def test_from_dlpack(copy_kw, data, dtype_device_pair):
9696
# TODO: 1. test copy; 2. generate inputs on non-default devices;
9797
# 3. test for copy=False cross-device transfers
9898
# 4. test 0D arrays / numpy scalars (the latter do not support dlpack ATM)
99-
99+
dtype, device = dtype_device_pair
100+
x = data.draw(hh.arrays(dtype=dtype, shape=hh.shapes(min_dims=1, max_side=2)))
100101
copy = copy_kw["copy"] if copy_kw else None
101102
if copy is False:
102103
# XXX there is no way to tell if a no-copy cross-device transfer is meant to succeed
103104
devices = [x.device]
104105
else:
105-
devices = xp.__array_namespace_info__().devices()
106-
devices = _compatible_devices(devices)
106+
devices = [device]
107107

108108
tgt_device_kw = data.draw(
109109
hh.kwargs(device=st.sampled_from(devices) | st.none())

0 commit comments

Comments
 (0)