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
12 changes: 11 additions & 1 deletion array_api_strict/_creation_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,16 @@ def _supports_buffer_protocol(obj: object) -> TypeIs[SupportsBufferProtocol]:
return True


def _contains_nested_array(obj: list[object] | tuple[object, ...]) -> bool:
from ._array_object import Array

return any(
isinstance(item, Array)
or isinstance(item, list | tuple) and _contains_nested_array(item)
for item in obj
)


def asarray(
obj: Array | complex | NestedSequence[complex] | SupportsBufferProtocol,
/,
Expand Down Expand Up @@ -97,7 +107,7 @@ def asarray(

if isinstance(obj, Array):
return Array._new(np.array(obj._array, copy=copy, dtype=_np_dtype), device=device)
elif isinstance(obj, list | tuple) and any(isinstance(x, Array) for x in obj):
elif isinstance(obj, list | tuple) and _contains_nested_array(obj):
raise TypeError("Nested Arrays are not allowed. Use `stack` instead.")

if dtype is None and isinstance(obj, int) and (obj > 2 ** 64 or obj < -(2 ** 63)):
Expand Down
27 changes: 27 additions & 0 deletions array_api_strict/tests/test_creation_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,33 @@ def test_asarray_nested_arrays():
asarray([1, asarray(1)])


@pytest.mark.parametrize("outer", [list, tuple])
@pytest.mark.parametrize("inner", [list, tuple])
@pytest.mark.parametrize("array_value", [1, [1]])
@pytest.mark.parametrize("depth", [2, 3])
def test_asarray_deeply_nested_arrays(outer, inner, array_value, depth):
obj = asarray(array_value)
for _ in range(depth - 1):
obj = inner([obj])
obj = outer([obj])
with pytest.raises(TypeError, match="Nested Arrays are not allowed"):
asarray(obj)


def test_asarray_nested_array_after_scalars():
with pytest.raises(TypeError, match="Nested Arrays are not allowed"):
asarray([[1, 2], [3, asarray(4)]])


@pytest.mark.parametrize("outer", [list, tuple])
@pytest.mark.parametrize("inner", [list, tuple])
def test_asarray_nested_scalars(outer, inner):
obj = outer([inner([1, 2]), inner([3, 4])])
res = asarray(obj)
assert res.shape == (2, 2)
assert all(res == asarray([[1, 2], [3, 4]]))


def test_asarray_device_inference():
assert asarray([1, 2, 3]).device == CPU_DEVICE

Expand Down
Loading