Repository navigation
Implement sorting, reductions, scans and findall with AcceleratedKernels - #790
Merged
Merged
Conversation
1 task
Member
|
Great! I've been wondering if we shouldn't do this :) |
This was referenced Sep 24, 2026
`sort!`, `sort`, `sortperm!`, `sortperm`, `partialsort!` and `reverse(!)` on GPU arrays call AcceleratedKernels, which picks the algorithm and its settings for the array's device. Base's algorithm objects become requirements: `MergeSort`, `InsertionSort` and Base's default ask for a stable sort, `QuickSort` and `PartialQuickSort` allow an unstable one, and any other is an `ArgumentError`. An AcceleratedKernels algorithm passed as `alg` is used as given (`sort!(A; alg=AK.RadixSort())`). `scratch` is accepted and ignored. The methods follow Base's shapes: vectors take no `dims` and other arrays require it, `sortperm!` checks the axes of the index array, `partialsort!` returns an element for an integer and a view for a range (and also takes a positional ordering), and `reverse!(v, start, stop)` is a no-op for a trivial interval before it checks bounds. JLArrays keeps reference implementations on its storage, since AcceleratedKernels' kernels do not run on its back-end; they accept the same `alg` values. Co-authored-by: shreyas-omkar <shreyashegdeplus06@gmail.com>
Base's reductions of GPU arrays call AcceleratedKernels' primitives directly: `AK.mapreduce` for scalar results, and `AK.mapreducedim!` into a destination otherwise. AcceleratedKernels applies `init` once and needs no neutral element, and its results have its accumulator type; everything that exists only because Base says so is decided here, on the host: - Empty inputs give Base's result or error: `Base.mapreduce_empty_iter` for a scalar, and `Base.reducedim_init` on an empty host stand-in along `dims` (`maximum` of an empty array returned `typemin`). - A one-element input is one scalar read and `Base.mapreduce_first` (or `op(init, x)`), so it has Base's type. - A scalar result is converted to the type Base's fold settles on (`sum([1, 2]; init=Int8(0))` is an `Int`); along `dims` the result is `typeof(init)`, else the fold type, and `dims` is checked as in Base. - Sums along `dims` start from zero, as Base's do, so a slice of `-0.0`s sums to `0.0`. - An explicit `init=nothing` is an initial value, as in Base; it used to mean none. `Base.mapreducedim!` on GPU arrays folds into the destination's values. GPUArrays no longer calls its `mapreducedim!` hook, nor picks neutral elements for back-ends: their overrides of `GPUArrays.mapreducedim!` are no longer reached from Base's API. The function stays, implemented with AcceleratedKernels, so that back-ends that extend it keep loading until they delete their methods. `sum!`, `prod!`, `maximum!`, `minimum!`, `extrema!`, `any!`, `all!` and `count!` take Base's `init::Bool` without Base's pass that fills `r`: `init=true` reduces from the operator's identity in `eltype(r)` (given to AcceleratedKernels as `init`), or, for `maximum!`, `minimum!` and `extrema!`, overwrites `r`; `init=false` folds into `r`. An input whose every slice is empty gets Base's result from a host stand-in. `findmin` and `findmax` (and `argmin`, `argmax`) reduce `(f(x), i)` pairs without an `init`, since partial results start from their first element: an empty input throws Base's error instead of returning `typemax`, and `f` needs no `typemin`. Indices come from `keys(A)`. `argmin(f, A)` and `argmax(f, A)`, which Base computes by iterating, read the element `findmin(f, A)` or `findmax(f, A)` selects. `count` requires its predicate to return a `Bool`, as Base's `_bool` does. JLArrays runs AcceleratedKernels' host algorithms on its storage, through an internal indirection that is not an extension point for back-ends.
`accumulate!`, `accumulate`, `cumsum(!)` and `cumprod(!)` on GPU arrays call `AK.accumulate!` through the four `Base._accumulate!` methods, split as Base's are. `init` is passed on only when given, so an explicit `init=nothing` stays an initial value. Base's shape rules hold: an in-place scan without `dims` needs a vector, and an allocating one scans other arrays in linear order, keeping their shape, and a `dims` beyond the array's copies it, ignoring `init` (AcceleratedKernels would apply `init` to every element). `cumsum!` of floating-point vectors, which Base computes pairwise, uses the same scan. Base's front-end chooses the destination's element type, which sets the running values' type in AcceleratedKernels, so result types are Base's. Where a caller passes a wider destination to `accumulate!`, the scan runs in that type, while some of Base's code paths run in the elements' type and overflow: `accumulate!(*, zeros(Int, 4), fill(0x10, 4))` gives `[16, 256, 4096, 65536]`, where Base gives `[16, 0, 0, 0]`. JLArrays keeps reference implementations on its storage.
`findall(A)` and `findall(f, A)` on GPU arrays call `AK.findall`, which
selects the `items` GPUArrays passes: Base's indices, `keys(A)` (`Int` for
vectors, `CartesianIndex` otherwise), except `LinearIndices(A)` for the
predicate form of a 0-dimensional array, where Base's indices are linear.
Logical indexing (`A[mask]`, including mixed forms such as `A[mask, :]`)
turns the mask into the indices it selects, as the back-ends each did; a
mask of the array's shape instead selects the values themselves, in one
pass (`AK.findall(mask; items=A)`). A mask used as the only index that
does not fit the array is a `BoundsError`, as in Base; it used to select,
since GPUArrays' `checkindex` for index arrays on the device treated a
mask's elements as integer indices. A mask mixed with other indices is
still not checked.
Scalar `any` and `all` with a predicate that returns a `Bool` use
AcceleratedKernels' short-circuiting implementation. A predicate that can
return `missing` gets Base's three-valued logic from a reduction of codes
(`false` < `missing` < `true`) with `max` or `min`, which needs no array
of `Union{Missing, Bool}`. Along `dims` the predicate must return a
`Bool`, as in Base. A predicate that inference shows cannot return those
values is an `ArgumentError` before launching; an empty array is never
checked, as Base never calls the predicate on it.
JLArrays keeps reference implementations on its storage.
Back-ends' specializations of `GPUArrays.mapreducedim!` are no longer called, so the function goes. The released JLArrays only overrides that function, and fails to reduce, sort, scan or `findall` with this GPUArrays; JLArrays 0.4.1 requires GPUArrays 12.
maleadt
marked this pull request as ready for review
October 1, 2026 11:37
This was referenced Oct 6, 2026
Closed
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.
This implements Base's sorting, reduction, scan and
findallfunctions once for every GPU array, on top of AcceleratedKernels (AK). Today each back-end carries its own copies (CUDA.jl, Metal.jl, AMDGPU.jl and oneAPI.jl all definefindall, logical indexing, scans and sorting, and every back-end has its own reduction kernel), and some back-ends have none: OpenCL.jl has nosort!oraccumulate. With this PR a back-end gets all of them by defining nothing, and back-ends can delete their copies.It builds on the AK host-API rework (JuliaGPU/AcceleratedKernels.jl#133, released as AK 0.5). The two divide the work: AK provides primitives with contracts of their own, modelled on CUB's (
initapplied once, no neutral element needed, one documented accumulator type,findallselecting caller-chosen items), and GPUArrays turns them into Base's API, deciding on the host everything that exists only because Base says so: result types, empty results and their errors, the one-element result,dimsrules and index types. It supersedes #788 by @shreyas-omkar, whose tests and JLArrays parts are kept (co-authored); the algorithm selection #788 did here now lives in AK.What users get
Base's functions on GPU arrays, with Base's results:
Sorting is stable by default, as in Base. Base's algorithm objects become requirements (
MergeSortandInsertionSortask for a stable sort,QuickSortandPartialQuickSortallow an unstable one, others are anArgumentError), and AK algorithms are accepted as given:sort!(v; alg=AK.RadixSort()).Reductions apply
initonce, need no neutral element, and give Base's result types and empty results:In-place reductions (
sum!,prod!,maximum!,minimum!,extrema!,any!,all!,count!) take Base'sinit::Boolwithout Base's pass that fillsrfirst.Scans follow Base's shape rules: without
dims, an allocating scan of a matrix runs in linear order and keeps the shape, an in-place one is an error, and adimsbeyond the array's copies it.findallgives Base's index types, andA[mask]with a mask ofA's shape selects the values in one pass instead of computing indices and gathering.any/allshort-circuit on the GPU, and give Base's three-valued logic for predicates that returnmissing, without storingmissingin a device array.As with any GPU reduction,
opmust be associative and commutative (a scan's operator associative). The remaining differences from Base:accumulate!into a destination wider than the elements runs in the destination's type, where some of Base's paths overflow in the elements' type, and in-place reductions convert to the destination once rather than after every element.reduce(closure, A; dims), for which Base has noreducedim_init.any!into a non-Booldestination.A[mask, :]) that does not fit its dimension is not aBoundsError, as before.How it is built
mapreducecallsAK.mapreducefor scalar results andAK.mapreducedim!into an array of Base's element type otherwise (withinit, or in overwrite mode). Base's empty results come fromBase.mapreduce_empty_iter, or fromBase.reducedim_initon an empty host stand-in alongdims(with Base's errors); a one-element input is one scalar read andBase.mapreduce_first.Base.mapreducedim!folds into its destination.findmin/findmaxreduce(f(x), i)pairs withoutinit, and give Base's indices throughkeys(A).findallpasses AK the items to select:keys(A), orLinearIndices(A)for the predicate form of a 0-dimensional array.findall) on its storage, since AK's GPU kernels do not run on its back-end. The reductions go through an internal indirection (_ak_mapreduce,_ak_mapreducedim!) that is not an extension point for back-ends. JLArrays becomes 0.4.1 and requires GPUArrays 12: the released 0.4.0 only overridesGPUArrays.mapreducedim!, and with this GPUArrays fails to reduce, sort, scan orfindall("This kernel is unavailable for backend CPU"), which is what makes an 11.x release impossible.For back-end packages
This is a breaking release, GPUArrays 12.0. GPUArrays no longer calls the
GPUArrays.mapreducedim!hook, nor picks neutral elements for back-ends: every Base reduction goes to AK directly, and the hook is removed. Back-ends have to delete their methods for it and require GPUArrays 12; the other methods below can go at the same time (until they do, they shadow GPUArrays' and keep their old behaviour):GPUArrays.mapreducedim!and its kernels (CUDACore/src/mapreduce.jl),findall,to_index/to_indices, the_accumulate!methods,accumulateandaccumulate_pairwise!,sort!/sort/sortperm(!)/partialsort(!),reverse(!).GPUArrays.mapreducedim!and its kernels (src/mapreduce.jl),findall,to_index/to_indices, the_accumulate!methods,accumulateandaccumulate_pairwise!. Its MPS sorts can stay where they are faster, if they accept AK algorithms.GPUArrays.mapreducedim!(src/kernels/mapreduce.jl),findall,to_index/to_indices, thesort!/sortperm(!)andaccumulate(!)/cumsum/cumprodwrappers (which call AK 0.4's API and have to change anyway), andreverse(!).GPUArrays.mapreducedim!(src/mapreduce.jl),findall,to_index/to_indices, thesort!/sortperm(!)andaccumulate(!)/cumsum/cumprodwrappers (which call AK 0.4's API too).GPUArrays.mapreducedim!(src/mapreduce.jl).Interim slowdown for small reductions. Back-ends' reduction kernels are no longer used, and for small reductions CUDA.jl's is faster than AK's (table below): up to about 4× for a 1000×10 reduction along
dims=2, while AK is about 3× faster for a 1000×100000 one alongdims=1. The fix is to optimize AK's kernels, tracked in JuliaGPU/AcceleratedKernels.jl#135.Testing
After rebasing onto current main, with AK 0.5.0 and the 12.0 commit (
c846739): JLArray and Array on Julia 1.13, 27,347.Before the rebase, GPUArrays' testsuite (passing tests, Julia 1.13 unless noted):
a66d284sorting andreverse86920aareductions57a33a3scansf6492fbfindalland the predicatesOn the last commit, also:
reverse,findalland scan methods: Base's algorithm objects, empty and 0-dimensional arrays,init=nothing, mixed element types and an ambiguity withfindall(in(...), A). No reduction fails, since GPUArrays no longer calls CUDA.jl's reduction kernel.AMDGPU.jl, Metal.jl and oneAPI.jl were not tested with this branch.
Before merging
[sources]entry and set the compat toAcceleratedKernels = "0.5".Follow-ups:
findmin!/findmax!on GPU arrays (#792), and faster small reductions in AK (JuliaGPU/AcceleratedKernels.jl#135).Reduction timings: CUDA.jl's reduction kernel against AK's
sumofFloat32withinit=0f0on an RTX 5080 (Julia 1.13, CUDA.jl 6.4.0), in µs, minimum of 30 synchronized runs: CUDA.jl'sGPUArrays.mapreducedim!kernel, which GPUArrays used so far, against AK'smapreducedim!, and for "all" againstAK.mapreduce(what GPUArrays now calls for a scalar result; CUDA.jl's result is copied to the host too).dims=1dims=2dims=1dims=2dims=1dims=2