diff --git a/docs/src/api.md b/docs/src/api.md index 826cd7455..f7c8e9e28 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -15,6 +15,15 @@ @ndrange ``` +### Reductions and scans + +```@docs +@groupreduce +@subgroupreduce +@groupscan +@subgroupscan +``` + ## Host language !!! note diff --git a/docs/src/kernels.md b/docs/src/kernels.md index e1bf57a8b..17cc84ea6 100644 --- a/docs/src/kernels.md +++ b/docs/src/kernels.md @@ -213,6 +213,36 @@ statements, and [`@uniform`](@ref) evaluates an expression outside the work-item can be reused across `@synchronize` statements. For scratch storage that does not need to survive across `@synchronize`, an `MArray` can be used instead. +## Reductions, scans, and sub-groups + +[`@groupreduce`](@ref) and [`@groupscan`](@ref) reduce and scan values over the workgroup, +and [`@subgroupreduce`](@ref) and [`@subgroupscan`](@ref) over a sub-group. Like +[`@synchronize`](@ref), they are collectives: all work-items take part, including the ones +that pad a partial workgroup, which contribute the neutral element. They have to be used as +statements of their own, e.g. `res = @groupreduce(+, val, zero(T))`. + +For other sub-group operations, kernels can call the functions of +[KernelInterface](@ref kernelinterface) directly, e.g. `KernelInterface.shfl` or +`KernelInterface.sub_group_ballot`. These have to be executed by all work-items of the +sub-group as well, but `@kernel` only knows about its own collectives: in a kernel with +the default bounds checking, every other statement only runs on the work-items that are part +of the `ndrange`, so a partial workgroup leaves out the padding work-items. Use +`@kernel unsafe_indices=true` for kernels that call KernelInterface's sub-group functions, +and derive and check the indices yourself (without `@index(Global)`): + +```julia +# the sums of the sub-groups of every workgroup, in `out[sub-group, workgroup]` +@kernel unsafe_indices=true function sub_group_sums!(out, @Const(x)) + N = @uniform prod(@groupsize()) + i = (@index(Group, Linear) - 1) * N + @index(Local, Linear) + val = i <= length(x) ? x[i] : zero(eltype(x)) + total = KernelInterface.sub_group_reduce(+, val) + if KernelInterface.get_sub_group_local_id() == 1 + out[KernelInterface.get_sub_group_id(), @index(Group, Linear)] = total + end +end +``` + ## Launching kernels Construct a kernel by calling the kernel function on a backend and optional static sizes, then diff --git a/src/KernelAbstractions.jl b/src/KernelAbstractions.jl index f549aec38..e63459f37 100644 --- a/src/KernelAbstractions.jl +++ b/src/KernelAbstractions.jl @@ -37,6 +37,10 @@ and then invoked on the arguments. - [`@uniform`](@ref) - [`@synchronize`](@ref) - [`@print`](@ref) +- [`@groupreduce`](@ref) +- [`@subgroupreduce`](@ref) +- [`@groupscan`](@ref) +- [`@subgroupscan`](@ref) # Kernel constructor @@ -633,6 +637,7 @@ function __workitems_iterspace end end include("macros.jl") +include("groupreduction.jl") include("spawn.jl") ### diff --git a/src/groupreduction.jl b/src/groupreduction.jl new file mode 100644 index 000000000..d98489c6b --- /dev/null +++ b/src/groupreduction.jl @@ -0,0 +1,374 @@ +### +# Group and sub-group reductions and scans +# - @groupreduce +# - @subgroupreduce +# - @groupscan +# - @subgroupscan +### + +export @groupreduce, @subgroupreduce, @groupscan, @subgroupscan + +""" + @groupreduce(op, val, neutral[, groupsize]; subgroups = false) + +Reduce `val` over all work-items of the workgroup with the binary operator `op`, and return +the result on every work-item. `op` has to be associative and commutative, and `neutral` +its neutral element (`op(neutral, x) == x`). The result has the type of `neutral`, which +`val` is converted to. + +Work-items that are not part of the `ndrange` (padding of a partial workgroup) contribute +`neutral`, and don't evaluate `val`. Like [`@synchronize`](@ref), `@groupreduce` must be +reached by all work-items of the workgroup, and has to be used as a statement on its own: +`res = @groupreduce(op, val, neutral)`. + +The reduction uses local memory for one value per work-item, so its size has to be known at +compile time: either the kernel has a static workgroup size, or `groupsize` gives an upper +bound of the workgroup size (a constant, e.g. a literal or a type parameter of the kernel). + +With `subgroups = true` each sub-group first reduces its values with +[`KernelInterface.shfl_down`](@ref), which needs fewer barriers. It must only be used on +backends that support shuffles of the type of `neutral`, see +[`KernelInterface.supports_shuffle`](@ref). `subgroups` has to be a constant as well, e.g. +a type parameter set from the host: + +```julia +@kernel function sum_kernel!(out, @Const(x), ::Val{S}) where {S} + i = @index(Global) + res = @groupreduce(+, x[i], zero(eltype(out)), 1024; subgroups = S) + if @index(Local, Linear) == 1 + out[@index(Group, Linear)] = res + end +end + +subgroups = KernelInterface.supports_shuffle(backend, eltype(out)) +sum_kernel!(backend)(out, x, Val(subgroups); ndrange = length(x)) +``` +""" +macro groupreduce(args...) + positional, options = split_options(args, (:groupsize, :subgroups)) + 3 <= length(positional) <= 4 || + error("@groupreduce expects `op`, `val`, `neutral` and optionally `groupsize`") + op, val, neutral = positional + groupsize = length(positional) == 4 ? positional[4] : Base.get(options, :groupsize, nothing) + bound = groupsize === nothing ? :($__static_groupsize($(esc(:__ctx__)))) : + :(Val($(esc(groupsize)))) + subgroups = Base.get(options, :subgroups, false) + return __collective_call(neutral, val) do neutral, val + :( + $__groupreduce( + $(esc(:__ctx__)), $(esc(op)), $val, $neutral, $bound, Val($(esc(subgroups))), + ) + ) + end +end + +""" + @subgroupreduce(op, val, neutral) + +Reduce `val` over the work-items of the sub-group with the binary operator `op`, in the order +of the lanes, with [`KernelInterface.sub_group_reduce`](@ref), and return the result on every +work-item of the sub-group. `op` has to be associative, and `neutral` its neutral element. The +result has the type of `neutral`. + +Work-items that are not part of the `ndrange` contribute `neutral`. `@subgroupreduce` must be +reached by all work-items of the sub-group, and has to be used as a statement on its own: +`res = @subgroupreduce(op, val, neutral)`. + +It must only be used on backends that support shuffles of the type of `neutral`, see +[`KernelInterface.supports_shuffle`](@ref). + +!!! note + Unlike `@subgroupreduce`, calls of KernelInterface's sub-group functions (e.g. + `KernelInterface.sub_group_reduce`) in a `@kernel` aren't executed by the work-items + that pad a partial workgroup, see [Reductions, scans, and sub-groups](@ref). Use + `@kernel unsafe_indices=true` for such kernels. +""" +macro subgroupreduce(op, val, neutral) + return __collective_call(neutral, val) do neutral, val + :($__subgroupreduce($(esc(op)), $val, $neutral)) + end +end + +""" + @groupscan(op, val, neutral[, groupsize]; inclusive = true) + +Scan `val` over the work-items of the workgroup with the binary operator `op`, in the order +of `@index(Local, Linear)`: the work-item with local index `i` gets +`op(...op(op(val₁, val₂), val₃)..., valᵢ)` (inclusive), or the same up to `valᵢ₋₁` and +`neutral` for the first work-item (with `inclusive = false`). `op` has to be associative, +and `neutral` its neutral element (`op(neutral, x) == op(x, neutral) == x`). The result has +the type of `neutral`, which `val` is converted to. + +Work-items that are not part of the `ndrange` (padding of a partial workgroup) contribute +`neutral`, and don't evaluate `val`. Like [`@synchronize`](@ref), `@groupscan` must be +reached by all work-items of the workgroup, and has to be used as a statement on its own: +`res = @groupscan(op, val, neutral)`. + +The scan uses local memory for two values per work-item, sized like for +[`@groupreduce`](@ref): either the kernel has a static workgroup size, or `groupsize` gives +an upper bound of the workgroup size. `inclusive` has to be a constant as well. + +For example, to compact the elements of `x` that satisfy `pred` within each workgroup: + +```julia +@kernel function compact!(out, counts, @Const(x), pred) + i = @index(Global, Linear) + keep = pred(x[i]) + offset = @groupscan(+, Int32(keep), Int32(0); inclusive = false) + total = @groupreduce(+, Int32(keep), Int32(0)) + base = (@index(Group, Linear) - 1) * prod(@groupsize()) + if keep + out[base + offset + 1] = x[i] + end + if @index(Local, Linear) == 1 + counts[@index(Group, Linear)] = total + end +end +``` +""" +macro groupscan(args...) + positional, options = split_options(args, (:groupsize, :inclusive)) + 3 <= length(positional) <= 4 || + error("@groupscan expects `op`, `val`, `neutral` and optionally `groupsize`") + op, val, neutral = positional + groupsize = length(positional) == 4 ? positional[4] : Base.get(options, :groupsize, nothing) + bound = groupsize === nothing ? :($__static_groupsize($(esc(:__ctx__)))) : + :(Val($(esc(groupsize)))) + inclusive = Base.get(options, :inclusive, true) + return __collective_call(neutral, val) do neutral, val + :( + $__groupscan( + $(esc(:__ctx__)), $(esc(op)), $val, $neutral, $bound, Val($(esc(inclusive))), + ) + ) + end +end + +""" + @subgroupscan(op, val, neutral; inclusive = true) + +Scan `val` over the work-items of the sub-group with the binary operator `op`, in the order +of `KernelInterface.get_sub_group_local_id()`, with [`KernelInterface.sub_group_scan`](@ref). +Like [`@groupscan`](@ref), the scan is inclusive by default, and with `inclusive = false` +exclusive. `op` has to be associative, and `neutral` its neutral element. The result has the +type of `neutral`. + +How the work-items of a workgroup form sub-groups is unspecified, so the order of a +sub-group's lanes doesn't need to match `@index(Local, Linear)`. + +Work-items that are not part of the `ndrange` contribute `neutral`. `@subgroupscan` must be +reached by all work-items of the sub-group, and has to be used as a statement on its own: +`res = @subgroupscan(op, val, neutral)`. + +It must only be used on backends that support shuffles of the type of `neutral`, see +[`KernelInterface.supports_shuffle`](@ref). +""" +macro subgroupscan(args...) + positional, options = split_options(args, (:inclusive,)) + length(positional) == 3 || error("@subgroupscan expects `op`, `val` and `neutral`") + op, val, neutral = positional + inclusive = Base.get(options, :inclusive, true) + return __collective_call(neutral, val) do neutral, val + :($__subgroupscan($(esc(op)), $val, $neutral, Val($(esc(inclusive))))) + end +end + +# Convert `val` to the type of `neutral` *before* the call of the collective. The type of +# `val` may differ between the work-items, e.g. `Union{Float32, Float64}` for an accumulator +# that only some work-items added a `Float64` to, or because padding work-items contribute +# `neutral` instead of `val` (see `mask_collective`). Julia union-splits a call with such an +# argument into one call per type, so the work-items would execute different copies of the +# collective, and its barriers and shuffles. Only the `convert` may be split this way. +function __collective_call(f, neutral, val) + n, v = gensym(:neutral), gensym(:val) + return quote + let $v = $(esc(val)), $n = $(esc(neutral)) + $(f(n, :($convert($typeof($n), $v)))) + end + end +end + +# Separate `key = value` options (also after a `;`) from the positional macro arguments. +function split_options(args, keys) + positional = Any[] + options = Dict{Symbol, Any}() + function option!(ex) + (ex.args[1] isa Symbol && ex.args[1] in keys) || + error("unknown option `$(ex.args[1])`, expected one of $(join(keys, ", "))") + options[ex.args[1]] = ex.args[2] + return + end + for arg in args + if isexpr(arg, :parameters) + foreach(option!, arg.args) + elseif isexpr(arg, :(=)) || isexpr(arg, :kw) + option!(arg) + else + push!(positional, arg) + end + end + return positional, options +end + +const COLLECTIVES = ( + Symbol("@groupreduce"), Symbol("@subgroupreduce"), + Symbol("@groupscan"), Symbol("@subgroupscan"), +) + +# Whether `expr` is a collective that all work-items of a workgroup take part in. +is_collective(expr) = any(name -> is_macrocall(expr, name), COLLECTIVES) + +# A collective that is used as a statement: `@groupreduce(...)` or `lhs = @groupreduce(...)`. +function is_collective_stmt(stmt) + is_collective(stmt) && return true + isexpr(stmt, :(=)) && is_collective(stmt.args[2]) || return false + lhs = stmt.args[1] + lhs isa Symbol && return true + isexpr(lhs, :(::)) && lhs.args[1] isa Symbol && return true + isexpr(lhs, :tuple) && all(x -> x isa Symbol, lhs.args) && return true + return false +end + +collective_error(expr) = error( + "`$(expr.args[1])` must be used as a statement of its own, " * + "e.g. `res = $(expr.args[1])(op, val, neutral)`, found `$(expr)`" +) + +# Rewrite a collective statement, so that padding work-items contribute the neutral element +# instead of evaluating the value. +function mask_collective(stmt) + if isexpr(stmt, :(=)) + return Expr(:(=), stmt.args[1], mask_collective(stmt.args[2])) + end + # `args[2]` is the macro's `LineNumberNode`, options may come first after a `;` + args = copy(stmt.args) + i = 3 + while i <= length(args) && (isexpr(args[i], :parameters) || args[i] isa LineNumberNode) + i += 1 + end + length(args) >= i + 2 || error("`$(args[1])` expects at least `op`, `val` and `neutral`") + val, neutral = args[i + 1], args[i + 2] + args[i + 1] = :(__active_lane__ ? $val : $neutral) + return Expr(:macrocall, args...) +end + +# The workgroup size as a `Val`, if it is static. +@inline __static_groupsize(ctx::CompilerMetadata) = __static_groupsize(__iterspace(ctx)) +@inline __static_groupsize(::NDRange{N, B, W}) where {N, B, W <: StaticSize} = Val(prod(get(W))) +@inline __static_groupsize(::NDRange) = throw( + ArgumentError( + "group reductions and scans require a static workgroup size or an upper bound of it" + ) +) + +# The largest power of two smaller than `n`, or 0. +@inline function __prevpow2(n::T) where {T <: Integer} + n <= one(T) && return zero(T) + return one(T) << (8 * sizeof(T) - 1 - leading_zeros(n - one(T))) +end + +@inline function __groupreduce(ctx, op, val, neutral::T, ::Val{N}, ::Val{subgroups}) where {T, N, subgroups} + n = prod(groupsize(ctx)) + n <= N || throw(ArgumentError("@groupreduce: the workgroup size exceeds the given upper bound")) + storage = KI.localmemory(T, Val(N)) + lid = __index_Local_Linear(ctx) + if subgroups + res = __groupreduce_subgroups(op, convert(T, val), neutral, storage) + else + res = __groupreduce_tree(op, convert(T, val), storage, lid, n) + end + # all work-items have to read the result before the storage can be reused + KI.barrier() + return res +end + +# Tree reduction in local memory, folding the upper half of the values onto the lower half. +@inline function __groupreduce_tree(op, val, storage, lid, n) + @inbounds storage[lid] = val + KI.barrier() + s = __prevpow2(n) + while s > 0 + if lid <= s && lid + s <= n + @inbounds storage[lid] = op(storage[lid], storage[lid + s]) + end + KI.barrier() + s >>= 1 + end + return @inbounds storage[1] +end + +# Reduce every sub-group with shuffles, then reduce the results of the sub-groups in the first one. +@inline function __groupreduce_subgroups(op, val, neutral, storage) + sg = KI.get_sub_group_id() + lane = KI.get_sub_group_local_id() + val = __subgroupreduce(op, val, neutral) + if lane == 1 + @inbounds storage[sg] = val + end + KI.barrier() + + # every sub-group reduces the results of all sub-groups, which avoids running shuffles + # in a branch that only some sub-groups take + width = KI.get_sub_group_size() + acc = neutral + i = lane + while i <= KI.get_num_sub_groups() + acc = op(acc, @inbounds storage[i]) + i += width + end + acc = __subgroupreduce(op, acc, neutral) + KI.barrier() + if sg == 1 && lane == 1 + @inbounds storage[1] = acc + end + KI.barrier() + return @inbounds storage[1] +end + +# backends can implement `KI.sub_group_reduce` with native operations +@inline __subgroupreduce(op, val, neutral::T) where {T} = KI.sub_group_reduce(op, convert(T, val)) + +@inline function __groupscan(ctx, op, val, neutral::T, ::Val{N}, ::Val{inclusive}) where {T, N, inclusive} + n = prod(groupsize(ctx)) + n <= N || throw(ArgumentError("@groupscan: the workgroup size exceeds the given upper bound")) + # two buffers, the scan reads from one and writes to the other + storage = KI.localmemory(T, Val(2 * N)) + lid = __index_Local_Linear(ctx) + + # Hillis-Steele: after the step with distance `d`, every work-item holds the scan of the + # (up to) `2d` values ending at its own + src = 0 + @inbounds storage[lid] = convert(T, val) + KI.barrier() + d = 1 + while d < n + x = @inbounds storage[src + lid] + if lid > d + x = op(@inbounds(storage[src + lid - d]), x) + end + @inbounds storage[(N - src) + lid] = x + KI.barrier() + src = N - src + d <<= 1 + end + + if inclusive + res = @inbounds storage[src + lid] + else + res = lid == 1 ? neutral : @inbounds storage[src + lid - 1] + end + # all work-items have to read their result before the storage can be reused + KI.barrier() + return res +end + +# backends can implement `KI.sub_group_scan` with native operations +@inline function __subgroupscan(op, val, neutral::T, ::Val{inclusive}) where {T, inclusive} + val = KI.sub_group_scan(op, convert(T, val)) + if !inclusive + lane = KI.get_sub_group_local_id() + prev = KI.shfl_up(val, 1) + val = lane == 1 ? neutral : prev + end + return val +end diff --git a/src/macros.jl b/src/macros.jl index 678c64696..a4dc8c266 100644 --- a/src/macros.jl +++ b/src/macros.jl @@ -146,10 +146,12 @@ function is_scope_construct(expr::Expr) # expr.head === :let end +# Whether `stmt` contains a `@synchronize`, or a collective like `@groupreduce` that all +# work-items of the workgroup have to reach as well. function find_sync(stmt) result = Ref(false) postwalk(stmt) do expr - result[] |= is_sync(expr) + result[] |= is_sync(expr) || is_collective(expr) expr end return result[] @@ -181,6 +183,17 @@ function split(stmts) continue end + if is_collective_stmt(stmt) + # executed by all work-items, the padding ones contribute the neutral element + loop = WorkgroupLoop(current, allocations, false, nothing) + push!(new_stmts, emit(loop)) + allocations = Any[] + current = Any[] + take_line!(new_stmts) + push!(new_stmts, mask_collective(stmt)) + continue + end + has_sync = find_sync(stmt) if has_sync loop = WorkgroupLoop(current, allocations, is_sync(stmt), line) @@ -200,6 +213,7 @@ function split(stmts) recurse(x) = x function recurse(expr::Expr) expr = unblock_lines(expr) + is_collective(expr) && collective_error(expr) if expr.head in (:if, :elseif) && find_sync(expr) return split_branches(expr, recurse) elseif is_scope_construct(expr) && any(find_sync, expr.args) diff --git a/test/groupreduce.jl b/test/groupreduce.jl new file mode 100644 index 000000000..ff9e8f893 --- /dev/null +++ b/test/groupreduce.jl @@ -0,0 +1,335 @@ +# one result per workgroup, written by every work-item, to check that all of them get it +@kernel function groupreduce_static!(out, @Const(x), op, neutral, ::Val{S}) where {S} + i = @index(Global, Linear) + res = @groupreduce(op, x[i], neutral; subgroups = S) + out[i] = res +end + +@kernel function groupreduce_bound!(out, @Const(x), op, neutral, ::Val{N}, ::Val{S}) where {N, S} + i = @index(Global, Linear) + res = @groupreduce(op, x[i], neutral, N; subgroups = S) + out[i] = res +end + +# the same call site, and thus local memory, reused in a loop and in a branch +@kernel function groupreduce_loop!(out, @Const(x), ::Val{S}) where {S} + i = @index(Global, Linear) + acc = zero(eltype(out)) + for k in 1:3 + res = @groupreduce(+, k * x[i], zero(eltype(out)); subgroups = S) + acc += res + end + if true + m = @groupreduce max x[i] typemin(eltype(out)) subgroups = S + end + out[i] = acc + m +end + +@kernel function groupreduce_cartesian!(out, @Const(x), ::Val{S}) where {S} + I = @index(Global, Cartesian) + res = @groupreduce(+, x[I], zero(eltype(out)); subgroups = S) + out[I] = res +end + +@kernel unsafe_indices = true function groupreduce_unsafe!(out, @Const(x), ::Val{S}) where {S} + i = @index(Global, Linear) + val = i <= length(x) ? x[i] : zero(eltype(out)) + res = @groupreduce(+, val, zero(eltype(out)); subgroups = S) + if i <= length(out) + out[i] = res + end +end + +# every work-item gets the reduction of its sub-group; `sgs` records the sub-group +@kernel function subgroupreduce!(out, sgs, @Const(x)) + i = @index(Global, Linear) + res = @subgroupreduce(+, x[i], zero(eltype(out))) + out[i] = res + sgs[i] = (@index(Group, Linear), KernelInterface.get_sub_group_id()) +end + +# `val` of another type than `neutral`: the padding work-items contribute `neutral`, so the +# value is a `Union` of both types, and the call of the collective must not be union-split +@kernel function groupreduce_mixed!(out, @Const(x), ::Val{S}) where {S} + i = @index(Global, Linear) + res = @groupreduce(+, x[i], 0.0; subgroups = S) + out[i] = res +end + +# an accumulator that stays `Float32` on some work-items, and becomes `Float64` on others +@kernel function subgroupreduce_union!(out, sgs, @Const(x)) + i = @index(Global, Linear) + acc = 0.0f0 + for k in 1:2 + if x[i] > 50 + acc += Float64(x[i]) + end + end + res = @subgroupreduce(+, acc, 0.0f0) + out[i] = res + sgs[i] = (@index(Group, Linear), KernelInterface.get_sub_group_id()) +end + +@kernel function groupscan_mixed!(out, @Const(x)) + i = @index(Global, Linear) + res = @groupscan(+, x[i], 0) + out[i] = res +end + +@kernel function subgroupscan_mixed!(out, lanes, @Const(x)) + i = @index(Global, Linear) + res = @subgroupscan(+, x[i], 0) + out[i] = res + lanes[i] = KernelInterface.get_sub_group_local_id() +end + +# the composition of affine maps `x -> a * x + b`, first `f` then `g`: associative, but not +# commutative, so that the scans have to combine the values in order +compose(f, g) = (g[1] * f[1], g[1] * f[2] + g[2]) +const affine_identity = (1, 0) + +@kernel function groupscan!(out, @Const(x), op, neutral, ::Val{I}) where {I} + i = @index(Global, Linear) + res = @groupscan(op, x[i], neutral; inclusive = I) + out[i] = res +end + +@kernel function groupscan_bound!(out, @Const(x), op, neutral, ::Val{N}, ::Val{I}) where {N, I} + i = @index(Global, Linear) + res = @groupscan(op, x[i], neutral, N; inclusive = I) + out[i] = res +end + +# the scan in the order of the local linear index, with padding in the middle of a workgroup +@kernel function groupscan_cartesian!(out, @Const(x)) + I = @index(Global, Cartesian) + res = @groupscan(compose, x[I], affine_identity) + out[I] = res +end + +# the same call site in a loop, and an exclusive scan of the counts +@kernel function groupscan_loop!(out, @Const(x)) + i = @index(Global, Linear) + acc = 0 + for k in 1:3 + res = @groupscan(+, k * x[i], 0) + acc += res + end + excl = @groupscan (+) x[i] 0 inclusive = false + out[i] = acc + excl +end + +@kernel function subgroupscan!(out, lanes, @Const(x), op, neutral, ::Val{I}) where {I} + i = @index(Global, Linear) + res = @subgroupscan(op, x[i], neutral; inclusive = I) + out[i] = res + lanes[i] = KernelInterface.get_sub_group_local_id() +end + +# reference: the reduction of each workgroup of `groupsize` consecutive elements +function groupwise(op, x, groupsize) + return [reduce(op, x[((cld(i, groupsize) - 1) * groupsize + 1):min(cld(i, groupsize) * groupsize, end)]) for i in eachindex(x)] +end + +# reference: the scan of each workgroup of `groupsize` consecutive elements +function groupwise_scan(op, x, groupsize, neutral, inclusive) + out = similar(x, typeof(neutral)) + for first in 1:groupsize:length(x) + group = first:min(first + groupsize - 1, length(x)) + acc = neutral + for i in group + inclusive || (out[i] = acc) + acc = op(acc, x[i]) + inclusive && (out[i] = acc) + end + end + return out +end + +function groupreduce_testsuite(backend, AT) + b = backend() + algorithms = KI.supports_subgroups(b) ? (false, true) : (false,) + + @testset "subgroups = $S" for S in algorithms + types = (Int32, Int64, Float32) + @testset "$T, $(nameof(typeof(op)))" for T in types, (op, neutral) in ((+, zero(T)), (max, typemin(T))) + S && !KI.supports_shuffle(b, T) && continue + for (groupsize, n) in ((64, 64), (64, 256), (32, 100), (256, 1000), (7, 23), (1, 3)) + x = T.(rand(1:100, n)) + out = AT(zeros(T, n)) + groupreduce_static!(b, groupsize)(out, AT(x), op, neutral, Val(S); ndrange = n) + @test Array(out) == groupwise(op, x, groupsize) + + fill!(out, zero(T)) + groupreduce_bound!(b)(out, AT(x), op, neutral, Val(256), Val(S); ndrange = n, workgroupsize = groupsize) + @test Array(out) == groupwise(op, x, groupsize) + end + end + + @testset "argmin" begin + # (value, index) pairs, the smallest value with the smallest index + for (groupsize, n) in ((64, 256), (32, 100), (7, 23)) + x = [(Float32(rand(1:20)), Int32(i)) for i in 1:n] + neutral = (Inf32, typemax(Int32)) + out = AT(fill(neutral, n)) + groupreduce_static!(b, groupsize)(out, AT(x), min, neutral, Val(S); ndrange = n) + @test Array(out) == groupwise(min, x, groupsize) + end + end + + @testset "loop" begin + x = rand(1:100, 100) + out = AT(zeros(Int, 100)) + groupreduce_loop!(b, 64)(out, AT(x), Val(S); ndrange = 100) + @test Array(out) == 6 .* groupwise(+, x, 64) .+ groupwise(max, x, 64) + end + + @testset "cartesian" begin + x = rand(1:100, 10, 12) + out = AT(zeros(Int, 10, 12)) + groupreduce_cartesian!(b, (4, 8))(out, AT(x), Val(S); ndrange = size(x)) + ref = similar(x) + for I in CartesianIndices(x) + g = (cld(I[1], 4) - 1) * 4 .+ (1:4), (cld(I[2], 8) - 1) * 8 .+ (1:8) + ref[I] = sum(x[intersect(g[1], axes(x, 1)), intersect(g[2], axes(x, 2))]) + end + @test Array(out) == ref + end + + @testset "unsafe_indices" begin + x = rand(1:100, 100) + out = AT(zeros(Int, 100)) + groupreduce_unsafe!(b, 64)(out, AT(x), Val(S); ndrange = 128) + @test Array(out) == groupwise(+, x, 64) + end + + @testset "mixed types" begin + x = Float32.(rand(1:100, 100)) + out = AT(zeros(Float64, 100)) + groupreduce_mixed!(b, 64)(out, AT(x), Val(S); ndrange = 100) + @test Array(out) == groupwise(+, Float64.(x), 64) + end + end + + if KI.supports_subgroups(b) && KI.supports_shuffle(b, Float32) + @testset "@subgroupreduce" begin + width = KI.sub_group_size(b) + for (groupsize, n) in ((width, 4width), (2width, 2width + 5)) + x = Float32.(rand(1:100, n)) + out = AT(fill(-1.0f0, n)) + sgs = AT(fill((0, 0), n)) + subgroupreduce!(b, groupsize)(out, sgs, AT(x); ndrange = n) + out, sgs = Array(out), Array(sgs) + # padding work-items contribute zero + @test all(i -> out[i] == sum(x[j] for j in 1:n if sgs[j] == sgs[i]), 1:n) + + out = AT(fill(-1.0f0, n)) + sgs = AT(fill((0, 0), n)) + subgroupreduce_union!(b, groupsize)(out, sgs, AT(x); ndrange = n) + out, sgs = Array(out), Array(sgs) + @test all(i -> out[i] == sum(2x[j] for j in 1:n if sgs[j] == sgs[i] && x[j] > 50; init = 0.0f0), 1:n) + end + end + end + + @testset "@groupscan" begin + @testset "inclusive = $I" for I in (true, false) + for (groupsize, n) in ((64, 64), (64, 256), (32, 100), (256, 1000), (7, 23), (1, 3)) + x = rand(1:100, n) + out = AT(zeros(Int, n)) + groupscan!(b, groupsize)(out, AT(x), +, 0, Val(I); ndrange = n) + @test Array(out) == groupwise_scan(+, x, groupsize, 0, I) + + y = [(rand((-1, 1, 2)), rand(-5:5)) for _ in 1:n] + out = AT(fill((0, 0), n)) + groupscan_bound!(b)(out, AT(y), compose, affine_identity, Val(256), Val(I); ndrange = n, workgroupsize = groupsize) + @test Array(out) == groupwise_scan(compose, y, groupsize, affine_identity, I) + end + end + + @testset "cartesian" begin + x = [(rand((-1, 1, 2)), rand(-5:5)) for _ in 1:10, _ in 1:12] + out = AT(fill((0, 0), 10, 12)) + groupscan_cartesian!(b, (4, 8))(out, AT(x); ndrange = size(x)) + ref = similar(x) + for gi in 1:4:10, gj in 1:8:12 + acc = affine_identity + # local linear order is column-major within the workgroup + for j in gj:(gj + 7), i in gi:(gi + 3) + (i <= 10 && j <= 12) || continue + acc = compose(acc, x[i, j]) + ref[i, j] = acc + end + end + @test Array(out) == ref + end + + @testset "loop" begin + x = rand(1:100, 100) + out = AT(zeros(Int, 100)) + groupscan_loop!(b, 64)(out, AT(x); ndrange = 100) + @test Array(out) == 6 .* groupwise_scan(+, x, 64, 0, true) .+ groupwise_scan(+, x, 64, 0, false) + end + + @testset "mixed types" begin + x = Int32.(rand(1:100, 100)) + out = AT(zeros(Int, 100)) + groupscan_mixed!(b, 64)(out, AT(x); ndrange = 100) + @test Array(out) == groupwise_scan(+, Int.(x), 64, 0, true) + end + end + + if KI.supports_subgroups(b) && KI.supports_shuffle(b, Int) + @testset "@subgroupscan, inclusive = $I" for I in (true, false) + width = KI.sub_group_size(b) + # a 1-D workgroup of at most the sub-group width is a single sub-group + for n in (4width, 2width + 5) + x = [(rand((-1, 1, 2)), rand(-5:5)) for _ in 1:n] + out = AT(fill((0, 0), n)) + lanes = AT(zeros(Int, n)) + subgroupscan!(b, width)(out, lanes, AT(x), compose, affine_identity, Val(I); ndrange = n) + out, lanes = Array(out), Array(lanes) + # the scan in the order of the lanes; padding work-items contribute the identity + ref = similar(out) + for first in 1:width:n + group = first:min(first + width - 1, n) + for i in group + before = filter(j -> I ? lanes[j] <= lanes[i] : lanes[j] < lanes[i], group) + sort!(before; by = j -> lanes[j]) + ref[i] = foldl(compose, x[before]; init = affine_identity) + end + end + @test out == ref + end + end + + @testset "@subgroupscan, mixed types" begin + width = KI.sub_group_size(b) + n = 2width + 5 + x = Int32.(rand(1:100, n)) + out = AT(zeros(Int, n)) + lanes = AT(zeros(Int, n)) + subgroupscan_mixed!(b, width)(out, lanes, AT(x); ndrange = n) + out, lanes = Array(out), Array(lanes) + ref = similar(out) + for first in 1:width:n + group = first:min(first + width - 1, n) + for i in group + ref[i] = sum(Int(x[j]) for j in group if lanes[j] <= lanes[i]) + end + end + @test out == ref + end + end + + @testset "errors" begin + @test_throws "must be used as a statement" @macroexpand @kernel function f(y, x) + i = @index(Global) + y[i] = @groupreduce(+, x[i], 0) + end + @test_throws "unknown option" @macroexpand @kernel function f(y, x) + res = @groupreduce(+, x[1], 0; foo = 1) + end + end + return +end diff --git a/test/testsuite.jl b/test/testsuite.jl index cd75adbc3..8c0dc7b4d 100644 --- a/test/testsuite.jl +++ b/test/testsuite.jl @@ -44,6 +44,7 @@ include("convert.jl") include("specialfunctions.jl") include("random.jl") include("spawn.jl") +include("groupreduce.jl") function testsuite(backend, backend_str, backend_mod, AT, DAT; skip_tests = Set{String}()) @conditional_testset "Unittests" skip_tests begin @@ -118,6 +119,10 @@ function testsuite(backend, backend_str, backend_mod, AT, DAT; skip_tests = Set{ spawn_testsuite(backend, AT) end + @conditional_testset "Group reductions" skip_tests begin + groupreduce_testsuite(backend, AT) + end + return end