Repository navigation
Sort along dims with RadixSort - #129
Merged
Merged
Conversation
Member
Author
|
@christiangnrd @maleadt Please take a look |
shreyas-omkar
force-pushed
the
sh/seg-radix-main
branch
2 times, most recently
from
September 23, 2026 12:18
e54b853 to
692daaa
Compare
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
force-pushed
the
sh/seg-radix-main
branch
from
September 30, 2026 17:22
692daaa to
6bbee5e
Compare
dims with RadixSort
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
force-pushed
the
sh/seg-radix-main
branch
from
September 30, 2026 18:58
6bbee5e to
86c2cb8
Compare
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
RadixSortonly sorts whole arrays. Asking it to sort along a dimension is an error:So
sort!(A; dims)has merge sort and bitonic sort, and both slow down as slices get longer: merge sort needs aboutlog2(len)merge passes, bitonic sort aboutlog2(len)^2 / 2compare-exchange steps. Radix sort makes one pass per byte of the key, whatever the slice length. With this PR,RadixSortsorts along any dimension, and for long slices it is the fastest of the three:dims=1RadixSortMergeSortBitonicSortTimes 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 (1MFloat32s in slices of 16 take 2.25 ms against 0.024 ms on CUDA).Autotherefore still picks radix sort only for whole arrays. Alongdims, 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
dimsthus 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 dedicateddims=1kernel 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:
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.UInt32s, so an array of more than 2^32 elements would come out wrong. It now throws anArgumentErrorfor those, and for slices that long.block_sizeanditems_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.