Skip to content
Merged
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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
name = "GPUArrays"
uuid = "0c68f7d7-f131-5f86-a1c3-88cf8149b2d7"
version = "12.0.1"
version = "12.0.2"

[workspace]
projects = ["lib/GPUArraysCore", "lib/JLArrays", "test", "docs"]
Expand Down
14 changes: 5 additions & 9 deletions src/host/indexing.jl
Original file line number Diff line number Diff line change
Expand Up @@ -250,17 +250,13 @@ _findall_items(A) = ndims(A) == 0 ? LinearIndices(A) : keys(A)
# Those no longer carry the mask's shape, so a single mask is checked against the array first, as
# Base does; a mask mixed with other indices is not (as before).
Base.to_index(::AnyGPUArray, I::AbstractArray{Bool}) = findall(I)
@static if VERSION >= v"1.11.0-DEV.1157"
Base.to_indices(A::AnyGPUArray, I::Tuple{AbstractArray{Bool}}) =
(checkbounds(A, I[1]); (Base.to_index(A, I[1]),))
else
# (also reached for the last of several indices, whose `inds` are then not all of `A`'s)
_check_mask(A, inds, mask) = length(inds) == ndims(A) ? checkbounds(A, mask) : nothing
Base.to_indices(A::AnyGPUArray, I::Tuple{AbstractArray{Bool}}) =
(checkbounds(A, I[1]); (Base.to_index(A, I[1]),))
@static if VERSION < v"1.11.0-DEV.1157"
# Base turns a trailing mask into a `LogicalIndex`, bypassing `to_index`
Base.to_indices(A::AnyGPUArray, inds,
I::Tuple{Union{Array{Bool,N}, BitArray{N}}}) where {N} =
(_check_mask(A, inds, I[1]); (Base.to_index(A, I[1]),))
Base.to_indices(A::AnyGPUArray, inds, I::Tuple{AbstractArray{Bool}}) =
(_check_mask(A, inds, I[1]); (Base.to_index(A, I[1]),))
(Base.to_index(A, I[1]),)
end
# ... except that a mask of the array's shape selects the values themselves, in one pass
function Base.getindex(A::AbstractGPUArray, mask::AnyGPUArray{Bool})
Expand Down
4 changes: 2 additions & 2 deletions test/testsuite/findall.jl
Original file line number Diff line number Diff line change
Expand Up @@ -50,8 +50,8 @@ end
@test compare_exact((A, m) -> A[m], AT, a, m)
end
@test compare_exact((A, m) -> view(A, 1:5, :)[m], AT, rand(Float32, 10, 4), rand(Bool, 5, 4))
# ... and a mask that does not fit the array
for (a, m) in ((rand(Float32, 3), rand(Bool, 2)), (rand(Float32, 2, 3), rand(Bool, 3, 2)))
# ... and a mask that does not fit the array, though every index it selects is in bounds
for (a, m) in ((rand(Float32, 3), Bool[1, 0]), (rand(Float32, 2, 3), Bool[1 0; 1 0; 0 0]))
@test compare_exact((A, m) -> A[m], AT, a, m)
@test compare_exact(A -> A[m], AT, a) # (a host mask)
@test compare_exact(A -> A[view(m, :)], AT, a) # (a host view as the mask)
Expand Down
Loading