Skip to content

Implement sorting, reductions, scans and findall with AcceleratedKernels - #790

Merged
maleadt merged 5 commits into
mainfrom
tb/ak-glue
Oct 1, 2026
Merged

maleadt merged 5 commits into
mainfrom
tb/ak-glue

Conversation

@maleadt

@maleadt maleadt commented Sep 24, 2026 •

Copy link
Copy Markdown
Member

This implements Base's sorting, reduction, scan and findall functions 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 define findall, logical indexing, scans and sorting, and every back-end has its own reduction kernel), and some back-ends have none: OpenCL.jl has no sort! or accumulate. 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 (init applied once, no neutral element needed, one documented accumulator type, findall selecting 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, dims rules 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:

sort!(v); sort(A; dims=2); sortperm(v); partialsort!(v, 1:3)
reverse!(v, 3, 8); reverse(A; dims=1)
cumsum(A; dims=2); accumulate(*, v; init=2)
findall(x -> x > 0, A); A[A .> 0]
any(isnan, A); all(A; dims=1); findmax(A; dims=1)
reduce((a, b) -> a + b, A)          # no neutral element needed any more
sum!(r, A; init=false)              # folds into r
  • Sorting is stable by default, as in Base. Base's algorithm objects become requirements (MergeSort and InsertionSort ask for a stable sort, QuickSort and PartialQuickSort allow an unstable one, others are an ArgumentError), and AK algorithms are accepted as given: sort!(v; alg=AK.RadixSort()).

  • Reductions apply init once, need no neutral element, and give Base's result types and empty results:

    sum(CuArray([1, 2]); init=Int8(0))   # 3, an Int (was an Int8)
    sum(A; init=10)                      # 10 added once
    maximum(CuArray(Int[]))              # ArgumentError, as Base (was typemin(Int))
    sum(A; init=nothing)                 # an error, as Base: `nothing` is an initial value
    reduce((a, b) -> a + b, CuArray([true]))   # true, as Base
    sum(CuArray([-0.0]))                 # -0.0, as Base; along `dims` 0.0, as Base
    findmin(CuArray(Int[]))              # Base's error (was (typemax(Int), 1))
  • In-place reductions (sum!, prod!, maximum!, minimum!, extrema!, any!, all!, count!) take Base's init::Bool without Base's pass that fills r first.

  • 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 a dims beyond the array's copies it.

  • findall gives Base's index types, and A[mask] with a mask of A's shape selects the values in one pass instead of computing indices and gathering.

  • any/all short-circuit on the GPU, and give Base's three-valued logic for predicates that return missing, without storing missing in a device array.

As with any GPU reduction, op must be associative and commutative (a scan's operator associative). The remaining differences from Base:

  • Base's own code paths disagree or overflow: for example, 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.
  • Calls Base cannot make work, e.g. reduce(closure, A; dims), for which Base has no reducedim_init.
  • Calls outside Base's contract give other errors or results, e.g. an operator that always throws for the element types, or any! into a non-Bool destination.
  • A logical mask mixed with other indices (A[mask, :]) that does not fit its dimension is not a BoundsError, as before.

How it is built

  • mapreduce calls AK.mapreduce for scalar results and AK.mapreducedim! into an array of Base's element type otherwise (with init, or in overwrite mode). Base's empty results come from Base.mapreduce_empty_iter, or from Base.reducedim_init on an empty host stand-in along dims (with Base's errors); a one-element input is one scalar read and Base.mapreduce_first. Base.mapreducedim! folds into its destination.
  • findmin/findmax reduce (f(x), i) pairs without init, and give Base's indices through keys(A).
  • findall passes AK the items to select: keys(A), or LinearIndices(A) for the predicate form of a 0-dimensional array.
  • JLArrays runs AK's host algorithms (reductions) or Base (sorting, scans, 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 overrides GPUArrays.mapreducedim!, and with this GPUArrays fails to reduce, sort, scan or findall ("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):

  • CUDA.jl: GPUArrays.mapreducedim! and its kernels (CUDACore/src/mapreduce.jl), findall, to_index/to_indices, the _accumulate! methods, accumulate and accumulate_pairwise!, sort!/sort/sortperm(!)/partialsort(!), reverse(!).
  • Metal.jl: GPUArrays.mapreducedim! and its kernels (src/mapreduce.jl), findall, to_index/to_indices, the _accumulate! methods, accumulate and accumulate_pairwise!. Its MPS sorts can stay where they are faster, if they accept AK algorithms.
  • AMDGPU.jl: GPUArrays.mapreducedim! (src/kernels/mapreduce.jl), findall, to_index/to_indices, the sort!/sortperm(!) and accumulate(!)/cumsum/cumprod wrappers (which call AK 0.4's API and have to change anyway), and reverse(!).
  • oneAPI.jl: GPUArrays.mapreducedim! (src/mapreduce.jl), findall, to_index/to_indices, the sort!/sortperm(!) and accumulate(!)/cumsum/cumprod wrappers (which call AK 0.4's API too).
  • OpenCL.jl: 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 along dims=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):

Commit JLArray and Array
a66d284 sorting and reverse 25,232
86920aa reductions 25,581
57a33a3 scans 26,493
f6492fb findall and the predicates 26,702

On the last commit, also:

  • JLArray and Array on Julia 1.10: 24,795.
  • CLArray (OpenCL.jl on PoCL): 10,071.
  • CuArray (RTX 5080): 10,775, with CUDA.jl's methods for the families above overwritten in the test session by methods that call GPUArrays' (CUDA.jl's shadow GPUArrays' until it deletes them). With CUDA.jl as it is, 77 fail, all in its own sorting, reverse, findall and scan methods: Base's algorithm objects, empty and 0-dimensional arrays, init=nothing, mixed element types and an ambiguity with findall(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

  • AK 0.5 is registered: remove the [sources] entry and set the compat to AcceleratedKernels = "0.5".
  • Back-end PRs that delete the methods above and require GPUArrays 12 are ready. Until then the back-end jobs here fail to install (their compat excludes 12.0, and AMDGPU.jl and oneAPI.jl still require AK 0.4) and soft-fail.
  • The back-end bugs AK works around for now are fixed upstream and the workarounds removed (the checklist in Rework the host API: algorithms as values, primitive contracts, workspaces AcceleratedKernels.jl#133).

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

sum of Float32 with init=0f0 on an RTX 5080 (Julia 1.13, CUDA.jl 6.4.0), in µs, minimum of 30 synchronized runs: CUDA.jl's GPUArrays.mapreducedim! kernel, which GPUArrays used so far, against AK's mapreducedim!, and for "all" against AK.mapreduce (what GPUArrays now calls for a scalar result; CUDA.jl's result is copied to the host too).

shape reduction CUDA.jl AK AK / CUDA.jl
10000 all 16.3 14.9 0.92
1000×10 dims=1 9.5 18.0 1.89
1000×10 dims=2 4.5 17.9 4.00
1000000 all 17.1 18.3 1.07
1000×1000 dims=1 14.9 21.0 1.40
1000×1000 dims=2 20.9 19.1 0.91
100000000 all 469.2 459.8 0.98
1000×100000 dims=1 1410.5 485.2 0.34
1000×100000 dims=2 1490.1 1800.0 1.21

@SimonDanisch

Copy link
Copy Markdown
Member

Great! I've been wondering if we shouldn't do this :)

maleadt and others added 5 commits October 1, 2026 11:27
`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.
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