Repository navigation
Add dims support to GPU sorting algorithms #59
Description
Activity
In AMDGPU.jl, this surfaced as JuliaGPU/AMDGPU.jl#1030 (
sort!(x; dims=2)hits aMethodErrorin_sort_impl!). We added a stopgap downstream; here is what we measured (if helpful to decide the AK-side API).Two ways to do it with what AK has today:
- Per-slice: loop over
CartesianIndicesof the non-sorted dimensions and callAK.sort!on eachview. Works out of the box, since the merge sort handles contiguous and stridedSubArrays alike, but serializes into one tiny kernel launch per slice. - Tag and sort once: tag every element with the index of its slice, then do a single
AK.sort!with a comparator ordering lexicographically by(slice, element). Slices come out grouped and internally sorted, then get scattered back.
RX 7900 XTX,
Float32, ROCm 6.4.4, wholesort!(A; dims)call:size dims slices (1) per-slice (2) tagged (100, 100)1 100 1.22 ms 0.14 ms (1024, 1024)1 1024 29.06 ms 0.92 ms (1024, 1024)2 1024 30.95 ms 0.81 ms (8192, 128)2 8192 115.79 ms 0.82 ms (128, 8192)1 8192 114.38 ms 0.76 ms The gap is occupancy: (1) runs one merge sort per slice with nothing else in flight, so it degrades with slice count, which is exactly what
dims=2on a tall matrix produces.(2) is a workaround with a cost: a global
O(N log²N)sort where the problem needs only per-sliceO(n log²n), plus a tag array and AK's temporary (aTuple{Int32,Float64}pads to 16 B, so roughly 4× the footprint). A segmented merge sort, one that never merges across segment boundaries, recovers both, anddimson a dense array is just a segmented sort with a fixed segment length. It would coversortperm(A; dims)and may provide the missing to JuliaGPU/GPUArrays.jl#608.(2) is a workaround with a cost: a global
O(N log²N)sort where the problem needs only per-sliceO(n log²n), plus a tag array and AK's temporary (aTuple{Int32,Float64}pads to 16 B, so roughly 4× the footprint). A segmented merge sort, one whose merge passes never cross a segment boundary, recovers both: fordims=1the segments are contiguous and of fixed length, and the general case is the same thing over strided segments. On a 1024x1024dims=1sort that is 1 global pass over the values instead of the 11 passes over tagged pairs that (2) does. It would coversortperm(A; dims)and give JuliaGPU/GPUArrays.jl#608.I could try to help with a PR draft if this direction is something you'd consider.
- Per-slice: loop over
See JuliaGPU/GPUArrays.jl#608.