Skip to content

Sort along dims with RadixSort - #129

Merged
maleadt merged 5 commits into
JuliaGPU:mainfrom
shreyas-omkar:sh/seg-radix-main
Oct 1, 2026
Merged

maleadt merged 5 commits into
JuliaGPU:mainfrom
shreyas-omkar:sh/seg-radix-main

Conversation

@shreyas-omkar

@shreyas-omkar shreyas-omkar commented Sep 23, 2026 •

Copy link
Copy Markdown
Member

RadixSort only sorts whole arrays. Asking it to sort along a dimension is an error:

julia> A = CuArray(rand(Float32, 1_048_576, 16));

julia> AK.sort!(A; dims=1, alg=AK.RadixSort())
ERROR: ArgumentError: RadixSort does not support sorting along `dims`

So sort!(A; dims) has merge sort and bitonic sort, and both slow down as slices get longer: merge sort needs about log2(len) merge passes, bitonic sort about log2(len)^2 / 2 compare-exchange steps. Radix sort makes one pass per byte of the key, whatever the slice length. With this PR, RadixSort sorts along any dimension, and for long slices it is the fastest of the three:

16M Float32s along dims=1 slice length RadixSort MergeSort BitonicSort
RTX 5080 (CUDA) 16K 3.11 3.54 2.11
256K 3.04 5.69 4.23
4M 3.16 8.24 6.98
M1 (Metal) 16K 110.3 125.0 110.9
256K 102.5 202.9 234.9
4M 96.3 288.9 354.4

Times are in ms, the minimum of 10 sorts with a preallocated workspace, and the merge and bitonic sort times already include the speedup described below. With 64-bit keys radix sort makes twice as many passes: for 16M Int64s on the RTX 5080 it only wins from slices of about 1M. For short slices it is far slower than bitonic sort (1M Float32s in slices of 16 take 2.25 ms against 0.024 ms on CUDA). Auto therefore still picks radix sort only for whole arrays. Along dims, the crossover depends on the element type and the total size, not just the slice length, and a single length threshold in the sort tuning can't capture that.

How it works

A radix pass has each block count how often every digit occurs in its part of the input. An exclusive scan over all the counts gives every block the output positions for its elements, and the blocks scatter them there. For slices, no block straddles two slices and the counts are laid out slice by slice, so the same single scan puts each slice's elements in that slice's range. The scatter just subtracts where the slice starts.

Instead of adding separate kernels for this, the existing radix kernels now take the slice layouts that merge and bitonic sort already use; a whole array is one slice. Sorting along dims thus gets what the whole-array sort has: the portable kernels for backends without local-memory atomics, and the single-block sort in local memory, which now sorts short slices one per block. The price is some indexing overhead. A dedicated dims=1 kernel that indexes its input directly, as in @shreyas-omkar's segmented radix sort that this PR started from, is 8 to 13% faster on the RTX 5080, and 10 to 13% faster on the M1 for 16M elements (more for smaller inputs).

This also changes three other things:

  • Slices along the first dimension are contiguous, yet they were indexed through a stride like any other slice. That costs a 64-bit multiplication per element access and 64-bit divisions per thread, which GPUs do in software. Contiguous slices now have their own layout, indexed by offset. Merge and bitonic sort benefit too: for 16M Float32s in slices of 4M, merge sort goes from 10.1 to 8.2 ms and bitonic sort from 9.0 to 7.0 ms on the RTX 5080, and from 330 to 289 ms and 533 to 354 ms on the M1.
  • Radix sort computes positions as UInt32s, so an array of more than 2^32 elements would come out wrong. It now throws an ArgumentError for those, and for slices that long.
  • The radix kernels rank digits in chunks whose width was a constant, 32. The best width depends on the device: 64 speeds radix sort up by 7 to 13% on Metal and oneAPI, whole arrays included, but slows it down by about 6% on CUDA. Like block_size and items_per_thread, it is now a field of the algorithm, RadixSort(chunk_size=64), filled from the sort tuning, which keeps 32 as the default for now.

Testing

The tests pass at every commit with threads, CUDA and PoCL, and with AcceleratedKernels' kernels on KernelAbstractions 0.10's host backend (Julia 1.13 and 1.10). At the head they also pass on Metal (M1), oneAPI (Iris Xe) and OpenCL on Intel's driver. The new tests sort every supported element type in both orders, in slices short enough for local memory and longer ones, along every dimension of 2D and 3D arrays, including contiguous slices after singleton dimensions. They also cover the portable kernels, empty and singleton slices, and workspaces.

@shreyas-omkar

Copy link
Copy Markdown
Member Author

@christiangnrd @maleadt Please take a look

@shreyas-omkar
shreyas-omkar force-pushed the sh/seg-radix-main branch 2 times, most recently from e54b853 to 692daaa Compare September 23, 2026 12:18
maleadt added a commit to shreyas-omkar/AcceleratedKernels.jl that referenced this pull request Sep 30, 2026
`AK.sort!(A; dims, alg=AK.RadixSort())` used to throw an ArgumentError.
It now sorts the slices along any dimension, with the kernels of the
previous commit: each radix pass sorts all slices at once, and where the
chunked kernels run and a slice of up to `2 * block_size` elements fits in
local memory, one block sorts each slice there.

For long slices this beats merge and bitonic sort. On an RTX 5080, 16M
Float32s in slices of 1M sort in 3.1 ms, against 6.8 ms for merge sort
and 5.4 ms for bitonic sort; on an M1, in 95 ms against 266 and 236 ms.
For short slices it is much slower than bitonic sort.

`Auto` still picks radix sort only for whole arrays: along `dims`, where
it pays off depends on the element type and the total size as well as
the slice length.

Based on the segmented radix sort of JuliaGPU#129, by @shreyas-omkar.

Co-authored-by: shreyas-omkar <shreyashegdeplus06@gmail.com>
@maleadt maleadt changed the title feat: Support sort(A; dims=1) in RadixSort via a segmented radix Sort along dims with RadixSort Sep 30, 2026
@maleadt
maleadt changed the base branch from main to tb/api September 30, 2026 17:22
maleadt and others added 4 commits September 30, 2026 20:09
Sorting along `dims` indexes each slice through its stride, with a 64-bit
multiplication per element access and 64-bit divisions per thread, which
GPUs do in software. Slices along the first dimension (or after dimensions
of size 1) have stride 1: they now get a layout of their own, whose slices
are indexed by offset only.

This speeds up merge and bitonic sort along `dims=1`. For 16M Float32s in
slices of 4M, merge sort goes from 10.1 to 8.2 ms and bitonic sort from
9.0 to 7.0 ms on an RTX 5080, and from 330 to 289 ms and 533 to 354 ms on
an M1.
The histogram, scatter and single-block kernels now take a slices.jl
layout and index the array through `slice`, like the merge and bitonic
sort kernels. A whole-array sort is one flat slice, for which nothing
changes: the histogram of block `b` for digit `d` stays at
`d * num_blocks + b`, and `slice` returns the array itself.

For several slices, the histograms are laid out slice-major, so the
exclusive scan over all of them gives each block its position in the
concatenation of the sorted slices; the scatter subtracts the slice's
start. The single-block kernel sorts one slice per block. Launch and
buffer sizes are computed with overflow checks, as tiny slices make them
exceed the number of elements.

Nothing sorts slices with radix sort yet: `RadixSort` still rejects
`dims`.
`AK.sort!(A; dims, alg=AK.RadixSort())` used to throw an ArgumentError.
It now sorts the slices along any dimension, with the kernels of the
previous commit: each radix pass sorts all slices at once, and where the
chunked kernels run and a slice of up to `2 * block_size` elements fits in
local memory, one block sorts each slice there.

For long slices this beats merge and bitonic sort. On an RTX 5080, 16M
Float32s in slices of 1M sort in 3.1 ms, against 6.9 ms for merge sort
and 5.6 ms for bitonic sort; on an M1, in 99 ms against 245 and 299 ms.
For short slices it is much slower than bitonic sort.

`Auto` still picks radix sort only for whole arrays: along `dims`, where
it pays off depends on the element type and the total size as well as
the slice length.

Based on the segmented radix sort of JuliaGPU#129, by @shreyas-omkar.

Co-authored-by: shreyas-omkar <shreyashegdeplus06@gmail.com>
Radix sort computes positions as UInt32s, so it gave wrong results for
longer arrays. Along `dims`, the scatter's positions wrap consistently,
so only the length of each slice is limited.
@maleadt
maleadt deleted the branch JuliaGPU:main September 30, 2026 19:37
@maleadt maleadt closed this Sep 30, 2026
@maleadt maleadt reopened this Sep 30, 2026
@maleadt
maleadt changed the base branch from tb/api to main September 30, 2026 19:40
The chunked radix kernels rank each element among the earlier elements
of its chunk, whose width was fixed at 32. The width trades comparisons
per element against the size of the chunk histograms in local memory,
and the best value depends on the device: 64 is 7 to 13% faster on an M1
and an Iris Xe, 32 about 6% faster on an RTX 5080.

Like `block_size` and `items_per_thread`, it is now a field of
`RadixSort`, `chunk_size`, filled from the sort tuning's
`radix_chunk_size`. That defaults to 32, so nothing changes unless a
backend's tuning or the caller picks another width.
@maleadt
maleadt merged commit 9590edd into JuliaGPU:main Oct 1, 2026
37 of 39 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants