diff --git a/src/sort/sort.jl b/src/sort/sort.jl index 417c2061..91dc85ca 100644 --- a/src/sort/sort.jl +++ b/src/sort/sort.jl @@ -61,6 +61,9 @@ struct SampleSort <: SortAlgorithm end # Algorithm choice alg::Union{Nothing, SortAlgorithm}=nothing, + # Sort each slice along this dimension; `:` sorts the whole array flat + dims::Union{Colon, Integer}=Colon(), + # GPU settings block_size::Union{Nothing, Int}=nothing, @@ -71,6 +74,10 @@ struct SampleSort <: SortAlgorithm end Sorts the array `v` in-place using the specified backend. The `lt`, `by`, `rev`, and `order` arguments are the same as for `Base.sort`. +With the default `dims=:` the whole array is sorted as one flat vector. Pass an integer `dims` to +sort each 1D slice along that dimension independently, matching `Base.sort(A; dims)`. The `dims` +form always uses a comparison sort and so ignores `alg`. + ## CPU CPU settings: use at most `max_tasks` threads to sort the array such that at least `min_elems` elements are sorted by each thread. A parallel [`sample_sort!`](@ref) is used, processing @@ -138,12 +145,24 @@ function _sort_impl!( alg::Union{Nothing, SortAlgorithm}=nothing, + # Sort each 1D slice along this dimension; `:` sorts the whole array as one flat vector + dims::Union{Colon, Integer}=Colon(), + # GPU settings; nothing => each GPU algorithm picks its own tuned default block_size::Union{Nothing, Int}=nothing, # Temporary buffer, same size as `v` temp::Union{Nothing, AbstractArray}=nothing, ) + if !(dims isa Colon) + return _sort_dims!( + v, backend, Int(dims); + lt, by, rev, order, + max_tasks, min_elems, prefer_threads, + block_size, + ) + end + if use_gpu_algorithm(backend, prefer_threads) alg = isnothing(alg) ? MergeSort() : alg if alg isa MergeSort @@ -189,6 +208,54 @@ function _sort_impl!( end +# Sort each slice along `dim` on its own, like Base.sort(A; dims). +# We have no batched sort kernel, so tag every element with its slice, sort the whole array +# once by (slice, value), then scatter each element back to its place. Works on any backend. +function _sort_dims!( + v::AbstractArray{T, N}, backend::Backend, dim::Int; + lt, by, rev, order, + max_tasks, min_elems, prefer_threads, + block_size, +) where {T, N} + 1 <= dim <= N || throw(ArgumentError("dimension $dim is not 1 ≤ dims ≤ $N")) + slice_len = size(v, dim) + (length(v) <= 1 || slice_len <= 1) && return v # every slice is a singleton + + len = length(v) + bs = isnothing(block_size) ? 256 : block_size + + # a slice is picked by the other dims, so collapse `dim` to 1 + other_size = ntuple(d -> d == dim ? 1 : size(v, d), N) + slice_lin = LinearIndices(other_size) + slice_car = CartesianIndices(other_size) + elem_car = CartesianIndices(v) + + keys = similar(v, Tuple{Int, T}, len) + foreachindex(v, backend; max_tasks, min_elems, prefer_threads, block_size=bs) do i + ci = elem_car[i] + proj = ntuple(d -> d == dim ? 1 : ci[d], N) + sid = slice_lin[CartesianIndex(proj)] - 1 + @inbounds keys[i] = (sid, v[i]) + end + + # slices in order, values within a slice by the user's ordering + o = Base.Order.ord(lt, by, rev, order) + comp = (a, b) -> a[1] != b[1] ? a[1] < b[1] : Base.Order.lt(o, a[2], b[2]) + _sort_impl!(keys, backend; lt=comp, max_tasks, min_elems, prefer_threads, block_size) + + # sorted keys are grouped by slice, so slot k lands at row (k-1)%slice_len of its slice + foreachindex(keys, backend; max_tasks, min_elems, prefer_threads, block_size=bs) do k + s = (k - 1) ÷ slice_len + r = (k - 1) % slice_len + oc = slice_car[s + 1] + dst = ntuple(d -> d == dim ? r + 1 : oc[d], N) + @inbounds v[CartesianIndex(dst)] = keys[k][2] + end + + v +end + + """ sort( v::AbstractArray, backend::Backend=get_backend(v); @@ -244,6 +311,9 @@ end # Algorithm choice alg::Union{Nothing, SortAlgorithm}=nothing, + # Permute each slice along this dimension; `:` permutes the whole array flat + dims::Union{Colon, Integer}=Colon(), + # GPU settings block_size::Union{Nothing, Int}=nothing, @@ -255,6 +325,11 @@ Save into `ix` the index permutation of `v` such that `v[ix]` is sorted. The `lt `order` arguments are the same as for `Base.sortperm`. The same algorithms are used as for [`sort!`](@ref) with custom by-index comparators. +With the default `dims=:` the whole array is permuted as one flat vector. Pass an integer `dims` to +permute each 1D slice along that dimension independently, matching `Base.sortperm(A; dims)`; then +`ix` holds global linear indices and must have the same axes as `v`. The `dims` form always uses a +comparison sort and so ignores `alg`. + ## Algorithm choice By default, `sortperm!` uses [`sample_sortperm!`](@ref) on CPU backends and [`merge_sortperm!`](@ref) on GPU backends. Pass `alg=MergeSort(lowmem=true)` to use the lower-memory GPU permutation path. @@ -289,12 +364,24 @@ function _sortperm_impl!( alg::Union{Nothing, SortAlgorithm}=nothing, + # Permute each slice along this dimension; `:` permutes the whole array as one flat vector + dims::Union{Colon, Integer}=Colon(), + # GPU settings; nothing => merge sort's tuned default (sortperm is merge-only) block_size::Union{Nothing, Int}=nothing, # Temporary buffer, same size as `v` temp::Union{Nothing, AbstractArray}=nothing, ) + if !(dims isa Colon) + return _sortperm_dims!( + ix, v, backend, Int(dims); + lt, by, rev, order, + max_tasks, min_elems, prefer_threads, + block_size, + ) + end + if use_gpu_algorithm(backend, prefer_threads) alg = isnothing(alg) ? MergeSort() : alg bs = isnothing(block_size) ? 256 : block_size @@ -340,6 +427,65 @@ function _sortperm_impl!( end +# Same idea as _sort_dims!, but produce the permutation instead of sorting in place. +# We tag with (slice, value, index), sort, and write the original index into `ix`. The index +# also breaks ties, which keeps the permutation stable like Base.sortperm(A; dims). +function _sortperm_dims!( + ix::AbstractArray, v::AbstractArray{T, N}, backend::Backend, dim::Int; + lt, by, rev, order, + max_tasks, min_elems, prefer_threads, + block_size, +) where {T, N} + 1 <= dim <= N || throw(ArgumentError("dimension $dim is not 1 ≤ dims ≤ $N")) + axes(ix) == axes(v) || throw(ArgumentError("index array must have the same axes as the input")) + slice_len = size(v, dim) + len = length(v) + bs = isnothing(block_size) ? 256 : block_size + + if len <= 1 || slice_len <= 1 # each slice holds one element + foreachindex(v, backend; max_tasks, min_elems, prefer_threads, block_size=bs) do i + @inbounds ix[i] = i + end + return ix + end + + # a slice is picked by the other dims, so collapse `dim` to 1 + other_size = ntuple(d -> d == dim ? 1 : size(v, d), N) + slice_lin = LinearIndices(other_size) + slice_car = CartesianIndices(other_size) + elem_car = CartesianIndices(v) + + keys = similar(v, Tuple{Int, T, Int}, len) + foreachindex(v, backend; max_tasks, min_elems, prefer_threads, block_size=bs) do i + ci = elem_car[i] + proj = ntuple(d -> d == dim ? 1 : ci[d], N) + sid = slice_lin[CartesianIndex(proj)] - 1 + @inbounds keys[i] = (sid, v[i], i) + end + + # slices in order, then values, then index to break ties so the result stays stable + o = Base.Order.ord(lt, by, rev, order) + comp = (a, b) -> begin + a[1] != b[1] && return a[1] < b[1] + Base.Order.lt(o, a[2], b[2]) && return true + Base.Order.lt(o, b[2], a[2]) && return false + return a[3] < b[3] + end + _sort_impl!(keys, backend; lt=comp, max_tasks, min_elems, prefer_threads, block_size) + + # sorted keys are grouped by slice, so slot k lands at row (k-1)%slice_len of its slice + foreachindex(keys, backend; max_tasks, min_elems, prefer_threads, block_size=bs) do k + s = (k - 1) ÷ slice_len + r = (k - 1) % slice_len + oc = slice_car[s + 1] + dst = ntuple(d -> d == dim ? r + 1 : oc[d], N) + @inbounds ix[CartesianIndex(dst)] = keys[k][3] + end + + ix +end + + """ sortperm( v::AbstractArray, diff --git a/test/generic/sort.jl b/test/generic/sort.jl index d8b941de..39ffdd8d 100644 --- a/test/generic/sort.jl +++ b/test/generic/sort.jl @@ -831,4 +831,101 @@ end @test_throws ArgumentError AK.sort!(v; prefer_threads, alg=AK.RadixSort()) end end + + +@testset "sort_dims" begin + Random.seed!(0) + + # Fuzzy correctness against Base.sort(A; dims) for 2D and 3D arrays + for _ in 1:200 + nd = rand(2:3) + sz = ntuple(_ -> rand(1:15), nd) + for T in valid_backend_eltypes(BACKEND, (Int32, Float32)) + A_h = rand(T, sz...) + A = array_from_host(A_h) + for dim in 1:nd, rev in (false, true) + @test Array(AK.sort(A; prefer_threads, dims=dim, rev=rev)) == + sort(A_h; dims=dim, rev=rev) + end + end + end + + # by= and order= act on the values within each slice + A_h = rand(Float32, 17, 23) + A = array_from_host(A_h) + @test Array(AK.sort(A; prefer_threads, dims=1, by=x->-x)) == sort(A_h; dims=1, by=x->-x) + @test Array(AK.sort(A; prefer_threads, dims=2, + order=Base.Order.Reverse)) == sort(A_h; dims=2, order=Base.Order.Reverse) + + # In-place sorts each slice, leaves the array otherwise intact + A_h = rand(Int32, 40, 31) + A = array_from_host(A_h) + AK.sort!(A; prefer_threads, dims=2) + @test Array(A) == sort(A_h; dims=2) + + # dims=1 on a vector is a full sort + v_h = rand(Int32, 5000) + v = array_from_host(v_h) + @test Array(AK.sort(v; prefer_threads, dims=1)) == sort(v_h) + + # Singleton slice dimension is a no-op + A_h = rand(Float32, 1, 64) + A = array_from_host(A_h) + @test Array(AK.sort(A; prefer_threads, dims=1)) == A_h + + # Out-of-range dimension errors + A = array_from_host(rand(Float32, 8, 8)) + @test_throws ArgumentError AK.sort(A; prefer_threads, dims=3) + @test_throws ArgumentError AK.sort(A; prefer_threads, dims=0) +end + + +@testset "sortperm_dims" begin + Random.seed!(0) + + # Fuzzy correctness against Base.sortperm(A; dims); small integer ranges give many ties, so + # matching Base's index array exactly also checks that the permutation is stable + for _ in 1:200 + nd = rand(2:3) + sz = ntuple(_ -> rand(1:15), nd) + for T in valid_backend_eltypes(BACKEND, (Int32, Float32)) + A_h = T <: Integer ? rand(T(0):T(4), sz...) : rand(T, sz...) + A = array_from_host(A_h) + for dim in 1:nd, rev in (false, true) + ix = Array(AK.sortperm(A; prefer_threads, dims=dim, rev=rev)) + @test ix == sortperm(A_h; dims=dim, rev=rev) + @test A_h[ix] == sort(A_h; dims=dim, rev=rev) + end + end + end + + # by= and order= act on the values within each slice + A_h = rand(Float32, 17, 23) + A = array_from_host(A_h) + @test Array(AK.sortperm(A; prefer_threads, dims=1, by=x->-x)) == sortperm(A_h; dims=1, by=x->-x) + @test Array(AK.sortperm(A; prefer_threads, dims=2, + order=Base.Order.Reverse)) == sortperm(A_h; dims=2, order=Base.Order.Reverse) + + # In-place fills ix with the same global linear indices as Base + A_h = rand(Int32(0):Int32(5), 40, 31) + A = array_from_host(A_h) + ix = array_from_host(zeros(Int, 40, 31)) + AK.sortperm!(ix, A; prefer_threads, dims=2) + @test Array(ix) == sortperm(A_h; dims=2) + + # dims=1 on a vector is a full sortperm + v_h = rand(Int32(0):Int32(9), 5000) + v = array_from_host(v_h) + @test Array(AK.sortperm(v; prefer_threads, dims=1)) == sortperm(v_h) + + # Singleton slice dimension yields the identity index array + A_h = rand(Float32, 1, 64) + A = array_from_host(A_h) + @test Array(AK.sortperm(A; prefer_threads, dims=1)) == reshape(1:64, 1, 64) + + # Out-of-range dimension errors + A = array_from_host(rand(Float32, 8, 8)) + @test_throws ArgumentError AK.sortperm(A; prefer_threads, dims=3) + @test_throws ArgumentError AK.sortperm(A; prefer_threads, dims=0) +end end