Skip to content

Add dims support to GPU sorting algorithms #59

Description

@christiangnrd

Activity

  1. luraess commented on Aug 17, 2026

    @luraess
    Member

    In AMDGPU.jl, this surfaced as JuliaGPU/AMDGPU.jl#1030 (sort!(x; dims=2) hits a MethodError in _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:

    1. Per-slice: loop over CartesianIndices of the non-sorted dimensions and call AK.sort! on each view. Works out of the box, since the merge sort handles contiguous and strided SubArrays alike, but serializes into one tiny kernel launch per slice.
    2. 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, whole sort!(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=2 on a tall matrix produces.

    (2) is a workaround with a cost: a global O(N log²N) sort where the problem needs only per-slice O(n log²n), plus a tag array and AK's temporary (a Tuple{Int32,Float64} pads to 16 B, so roughly 4× the footprint). A segmented merge sort, one that never merges across segment boundaries, recovers both, and dims on a dense array is just a segmented sort with a fixed segment length. It would cover sortperm(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-slice O(n log²n), plus a tag array and AK's temporary (a Tuple{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: for dims=1 the segments are contiguous and of fixed length, and the general case is the same thing over strided segments. On a 1024x1024 dims=1 sort that is 1 global pass over the values instead of the 11 passes over tagged pairs that (2) does. It would cover sortperm(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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions