diff --git a/docs/src/api.md b/docs/src/api.md index 0d1d77696..f5f3e3845 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -15,6 +15,13 @@ @ndrange ``` +### Reductions and scans + +```@docs +@groupreduce +@groupscan +``` + ## Host language !!! note diff --git a/docs/src/kernelinterface.md b/docs/src/kernelinterface.md index f804aa5ea..8e9c3c885 100644 --- a/docs/src/kernelinterface.md +++ b/docs/src/kernelinterface.md @@ -210,7 +210,9 @@ used anywhere. kernel languages built on it. A kernel is either a KernelInterface kernel or a KernelAbstractions `@kernel`, never a mix: `@kernel` adds padding work-items to partial work-groups, which skip the kernel's body, so KernelInterface's sub-group functions called - in a `@kernel` aren't reached by all work-items. Host code may use the queries of both. + in a `@kernel` aren't reached by all work-items. A `@kernel` uses KernelAbstractions' + collectives instead, e.g. [`@groupreduce`](@ref KernelAbstractions.@groupreduce). Host + code may use the queries of both. ```@docs get_sub_group_size diff --git a/docs/src/kernels.md b/docs/src/kernels.md index 5d6df509f..f6d56052a 100644 --- a/docs/src/kernels.md +++ b/docs/src/kernels.md @@ -213,6 +213,22 @@ 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 and scans + +[`@groupreduce`](@ref) and [`@groupscan`](@ref) reduce and scan values over the workgroup. +Like [`@synchronize`](@ref), they are collectives: all work-items of the workgroup have to +reach them, not in a branch or loop that only some of them take, and they have to be used as +statements of their own, e.g. `res = @groupreduce(+, val, zero(T))`. Work-items that pad a +partial workgroup take part as well, contributing the neutral element. On backends with +sub-groups, `@groupreduce` uses them where it can. + +A kernel is either a `@kernel` or a kernel written against +[KernelInterface](@ref kernelinterface), never a mix: don't call KernelInterface's device +functions, such as its sub-group shuffles and votes, in a `@kernel`. They have to be executed +by all work-items of a sub-group, but in a `@kernel` with the default bounds checking, every +statement other than the collectives above only runs on the work-items that are part of the +`ndrange`. Host code can use KernelInterface's queries, e.g. to pick a kernel. + ## 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 06329b3f3..fa568c7e9 100644 --- a/src/KernelAbstractions.jl +++ b/src/KernelAbstractions.jl @@ -42,6 +42,8 @@ and then invoked on the arguments. - [`@uniform`](@ref) - [`@synchronize`](@ref) - [`@print`](@ref) +- [`@groupreduce`](@ref) +- [`@groupscan`](@ref) # Kernel constructor @@ -670,6 +672,7 @@ function __workitems_iterspace end end include("macros.jl") +include("groupreduction.jl") include("spawn.jl") include("foreach_index.jl") diff --git a/src/backend_launch.jl b/src/backend_launch.jl index 65bb040c5..353799581 100644 --- a/src/backend_launch.jl +++ b/src/backend_launch.jl @@ -8,15 +8,27 @@ ### """ - mkcontext(kernel::Kernel, ndrange, iterspace, [launch]) + mkcontext(kernel::Kernel, ndrange, iterspace, [launch]; subgroups = nothing) The hidden context argument for launching `kernel` over `ndrange`, partitioned as -`iterspace`, with the launch configuration `launch` (see [`select_launch`](@ref)). +`iterspace`, with the launch configuration `launch` (see [`select_launch`](@ref)) and the +backend's sub-group capabilities `subgroups` (see [`kernel_subgroups`](@ref)). """ mkcontext(kernel::Kernel, _ndrange, iterspace) = CompilerMetadata{ndrange(kernel), DynamicCheck}(_ndrange, iterspace) -mkcontext(kernel::Kernel, _ndrange, iterspace, launch) = - CompilerMetadata{ndrange(kernel), DynamicCheck}(_ndrange, iterspace; launch) +mkcontext(kernel::Kernel, _ndrange, iterspace, launch; subgroups = nothing) = + CompilerMetadata{ndrange(kernel), DynamicCheck}(_ndrange, iterspace; launch, subgroups) + +""" + kernel_subgroups(kernel::Kernel) + +The sub-group capabilities of the backend of `kernel` that its work-group collectives can +use, a [`SubgroupCapabilities`](@ref), or `nothing` if the backend has none, or if the kernel +doesn't use collectives (which doesn't depend on the backend, so that launching other +kernels doesn't query it). +""" +kernel_subgroups(kernel::Kernel) = + uses_collectives(kernel.f) ? subgroup_capabilities(backend(kernel)) : nothing mkcontext(kernel::Kernel, I, _ndrange, iterspace, ::Dynamic) where {Dynamic} = CompilerMetadata{ndrange(kernel), Dynamic}(I, _ndrange, iterspace) @@ -103,10 +115,18 @@ function launch_tuple(obj::Kernel, args::Tuple; ndrange = nothing, workgroupsize end function launch_kernel(obj::Kernel, launch, ndrange, _workgroupsize, iterspace, args::Tuple) + # a constant `nothing` for kernels without collectives; for kernels with, the call below + # is a function barrier that makes the rest of the launch type stable + subgroups = kernel_subgroups(obj) + launch_kernel(obj, launch, subgroups, ndrange, _workgroupsize, iterspace, args) + return nothing +end + +function launch_kernel(obj::Kernel, launch, subgroups, ndrange, _workgroupsize, iterspace, args::Tuple) b = backend(obj) # this might not be the final context, since we may tune the workgroupsize - ctx = mkcontext(obj, ndrange, iterspace, launch) + ctx = mkcontext(obj, ndrange, iterspace, launch; subgroups) kernel = compile(obj, ctx, args) # tune the workgroup size, keeping the context type (and thus the kernel) the same @@ -114,7 +134,7 @@ function launch_kernel(obj::Kernel, launch, ndrange, _workgroupsize, iterspace, range = something(ndrange, static_ndrange(obj)) threads = KI.launch_configuration(kernel; nitems = saturated_prod(extents(range))).workgroupsize iterspace, _ = partition(obj, ndrange, launch_workgroupsize(b, launch, threads, range)) - ctx = mkcontext(obj, ndrange, iterspace, launch) + ctx = mkcontext(obj, ndrange, iterspace, launch; subgroups) end # the geometry is valid by construction, except for the kernel's limits diff --git a/src/compiler.jl b/src/compiler.jl index c61f7d471..cbc22721b 100644 --- a/src/compiler.jl +++ b/src/compiler.jl @@ -1,26 +1,29 @@ """ - CompilerMetadata{StaticNDRange, CheckBounds, I, NDRange, Iterspace, Launch} + CompilerMetadata{StaticNDRange, CheckBounds, I, NDRange, Iterspace, Launch, Subgroups} The hidden context argument of kernels written with [`@kernel`](@ref). The `launch` field tells the index functions how the backend launched the kernel: `nothing` for a 1-D launch -indexed in `Int`, or a [`LinearLaunch`](@ref) or [`NDLaunch`](@ref). +indexed in `Int`, or a [`LinearLaunch`](@ref) or [`NDLaunch`](@ref). The `subgroups` field +holds the backend's [`SubgroupCapabilities`](@ref) for kernels that use work-group +collectives, or `nothing`. """ -struct CompilerMetadata{StaticNDRange, CheckBounds, I, NDRange, Iterspace, Launch} +struct CompilerMetadata{StaticNDRange, CheckBounds, I, NDRange, Iterspace, Launch, Subgroups} groupindex::I ndrange::NDRange iterspace::Iterspace launch::Launch + subgroups::Subgroups # CPU variant function CompilerMetadata{NDRange, CB}(idx, ndrange, iterspace) where {NDRange, CB} ndrange = cartesian(ndrange) - return new{NDRange, CB, typeof(idx), typeof(ndrange), typeof(iterspace), Nothing}(idx, ndrange, iterspace, nothing) + return new{NDRange, CB, typeof(idx), typeof(ndrange), typeof(iterspace), Nothing, Nothing}(idx, ndrange, iterspace, nothing, nothing) end # GPU variante: index is given implicit - function CompilerMetadata{NDRange, CB}(ndrange, iterspace; launch = nothing) where {NDRange, CB} + function CompilerMetadata{NDRange, CB}(ndrange, iterspace; launch = nothing, subgroups = nothing) where {NDRange, CB} ndrange = cartesian(ndrange) - return new{NDRange, CB, Nothing, typeof(ndrange), typeof(iterspace), typeof(launch)}(nothing, ndrange, iterspace, launch) + return new{NDRange, CB, Nothing, typeof(ndrange), typeof(iterspace), typeof(launch), typeof(subgroups)}(nothing, ndrange, iterspace, launch, subgroups) end end @@ -34,6 +37,7 @@ cartesian(t::Tuple) = CartesianIndices(t) @inline __iterspace(cm::CompilerMetadata) = cm.iterspace @inline __groupindex(cm::CompilerMetadata) = cm.groupindex @inline __launch(cm::CompilerMetadata) = cm.launch +@inline __subgroups(cm::CompilerMetadata) = cm.subgroups @inline __groupsize(cm::CompilerMetadata) = size(workitems(__iterspace(cm))) @inline __dynamic_checkbounds(::CompilerMetadata{NDRange, CB}) where {NDRange, CB} = CB <: DynamicCheck @inline __ndrange(::CompilerMetadata{NDRange}) where {NDRange <: StaticSize} = CartesianIndices(get(NDRange)) @@ -48,7 +52,7 @@ cartesian(t::Tuple) = CartesianIndices(t) function Adapt.adapt_structure(to, cm::CompilerMetadata{NDRange, CB, I}) where {NDRange, CB, I} iterspace = Adapt.adapt(to, cm.iterspace) if I === Nothing - return CompilerMetadata{NDRange, CB}(cm.ndrange, iterspace; cm.launch) + return CompilerMetadata{NDRange, CB}(cm.ndrange, iterspace; cm.launch, cm.subgroups) else return CompilerMetadata{NDRange, CB}(cm.groupindex, cm.ndrange, iterspace) end diff --git a/src/groupreduction.jl b/src/groupreduction.jl new file mode 100644 index 000000000..14decc286 --- /dev/null +++ b/src/groupreduction.jl @@ -0,0 +1,348 @@ +### +# Work-group reductions and scans +# - @groupreduce +# - @groupscan +### + +export @groupreduce, @groupscan + +""" + @groupreduce(op, val, neutral; groupsize) + +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: the values are +combined in an unspecified order. `neutral` has to be its neutral element +(`op(neutral, x) == op(x, neutral) == x`). For example, to find the smallest value and its +index, reduce `(value, index)` pairs with an operator that breaks ties by the index, such +as `min` on tuples, with the neutral element `(typemax(T), typemax(Int))`. + +The result has the type of `neutral`, which `val` is converted to: `op(x, y)` has to return +that type `T = typeof(neutral)` for arguments of type `T`. + +`@groupreduce` is a collective, like [`@synchronize`](@ref): it has to be used as a +statement on its own (`res = @groupreduce(op, val, neutral)`), and reached by all +work-items of the workgroup, not in a branch or loop that only some of them execute, or +after an early `return`. Work-items that pad a partial workgroup (outside of the `ndrange`) +take part as well: they contribute `neutral`, without evaluating `val`. So `op` and +`neutral` have to be the same for the whole workgroup, and computable on those work-items. + +The reduction stores values in local memory, so its size has to be known at compile time: +either the kernel has a static workgroup size, or the keyword `groupsize` gives an upper +bound of the workgroup size (a constant, e.g. a literal or a type parameter of the kernel). + +On backends with sub-groups, KernelAbstractions reduces the values of every sub-group with +[`KernelInterface.sub_group_reduce`](@ref) first, if the backend can shuffle values of the +type of `neutral`. + +```julia +@kernel function sum_kernel!(out, @Const(x)) + i = @index(Global) + res = @groupreduce(+, x[i], zero(eltype(out)); groupsize = 1024) + if @index(Local, Linear) == 1 + out[@index(Group, Linear)] = res + end +end +``` +""" +macro groupreduce(args...) + op, val, neutral, options = parse_collective("@groupreduce", args, (:groupsize,)) + bound = collective_bound(Base.get(options, :groupsize, nothing)) + return __collective_call(neutral, val) do neutral, val + :($__groupreduce($(esc(:__ctx__)), $(esc(op)), $val, $neutral, $bound)) + 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, +but needn't be commutative, and `neutral` its neutral element +(`op(neutral, x) == op(x, neutral) == x`). The result has the type of `neutral`, which `val` +is converted to, and `op` has to return that type. + +Like [`@groupreduce`](@ref), `@groupscan` is a collective that has to be used as a statement +on its own and reached by all work-items of the workgroup; padding work-items contribute +`neutral`. The scan stores two values per work-item in local memory, sized like for +`@groupreduce`: either the kernel has a static workgroup size, or the keyword `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...) + op, val, neutral, options = parse_collective("@groupscan", args, (:groupsize, :inclusive)) + bound = collective_bound(Base.get(options, :groupsize, nothing)) + 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 + +collective_bound(::Nothing) = :($__static_groupsize($(esc(:__ctx__)))) +collective_bound(groupsize) = :(Val($(esc(groupsize)))) + +# 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 + +# The `op`, `val` and `neutral` of a collective's arguments, and its `key = value` options +# (also after a `;`), which have to be among `keys`. +function parse_collective(name, args, keys) + positional = Any[] + options = Dict{Symbol, Any}() + function option!(ex) + key = ex.args[1] + (key isa Symbol && key in keys) || + error("$name: unknown option `$key`, expected one of $(join(keys, ", "))") + haskey(options, key) && error("$name: option `$key` given more than once") + options[key] = 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 + length(positional) == 3 || + error("$name expects `op`, `val` and `neutral`, and the options $(join(keys, ", ")) as keywords") + return positional..., options +end + +const COLLECTIVES = (Symbol("@groupreduce"), Symbol("@groupscan")) + +# 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. `neutral` is evaluated once, before the collective. +function mask_collective(stmt) + if isexpr(stmt, :(=)) + binding, call = mask_collective(stmt.args[2]).args + return Expr(:block, binding, Expr(:(=), stmt.args[1], call)) + 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 `op`, `val` and `neutral`") + val, neutral = args[i + 1], args[i + 2] + n = gensym(:neutral) + args[i + 1] = :(__active_lane__ ? $val : $n) + args[i + 2] = n + return Expr(:block, :($n = $neutral), 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}) where {T, N} + n = prod(groupsize(ctx)) + n <= N || throw(ArgumentError("@groupreduce: the workgroup size exceeds the given upper bound")) + storage = KI.localmemory(T, Val(N)) + if __shuffleable(__subgroups(ctx), T) + res = __groupreduce_subgroups(op, convert(T, val), neutral, storage) + else + res = __groupreduce_tree(op, convert(T, val), storage, __index_Local_Linear(ctx), 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 `KI.sub_group_reduce`, which backends can implement natively, +# then reduce the results of the sub-groups. This only uses the sub-group ids, not how the +# work-items form sub-groups, and all sub-groups execute the same sub-group operations, so it +# needs neither linear nor independent sub-groups. There are at most as many sub-groups as +# work-items, so `storage` has room for a value per sub-group. +@inline function __groupreduce_subgroups(op, val, neutral, storage) + sg = KI.get_sub_group_id() + lane = KI.get_sub_group_local_id() + val = KI.sub_group_reduce(op, val) + if lane == 1 + @inbounds storage[sg] = val + end + KI.barrier() + + # every sub-group reduces the results of all sub-groups, which avoids running sub-group + # operations 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 = KI.sub_group_reduce(op, acc) + KI.barrier() + # the sub-groups can combine the values in different orders, so broadcast one result + if sg == 1 && lane == 1 + @inbounds storage[1] = acc + end + KI.barrier() + return @inbounds storage[1] +end + +@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 + + +## sub-group capabilities + +""" + SubgroupCapabilities{F16, F32, F64}() + +The sub-group capabilities of a backend that the work-group collectives use, passed to the +kernel in its [`CompilerMetadata`](@ref): the backend supports sub-groups and shuffles of +`UInt32` values (and thus of other integer-like values, as words), and of `Float16`, +`Float32` and `Float64` values if `F16`, `F32` and `F64`. Kernels on backends without these +capabilities get `nothing` instead. +""" +struct SubgroupCapabilities{F16, F32, F64} end + +function subgroup_capabilities(backend::KI.Backend) + (KI.supports_subgroups(backend) && KI.supports_shuffle(backend, UInt32)) || return nothing + return SubgroupCapabilities{ + KI.supports_shuffle(backend, Float16), + KI.supports_shuffle(backend, Float32), + KI.supports_shuffle(backend, Float64), + }() +end + +# Whether values of type `T` can be shuffled. Floating-point types need the backend's support: +# a backend that shuffles a float type natively may do so without regard for the device, e.g. +# `Float64` on a device without it. Other primitive types are shuffled as words, and structs +# packed into `UInt32` words, whatever the types of their fields. +__shuffleable(::Nothing, ::Type) = false +@generated function __shuffleable(::SubgroupCapabilities{F16, F32, F64}, ::Type{T}) where {F16, F32, F64, T} + function shuffleable(S) + if isprimitivetype(S) + S === Float16 && return F16 + S === Float32 && return F32 + S === Float64 && return F64 + return !(S <: AbstractFloat) && sizeof(S) in (1, 2, 4, 8, 16) + end + isbitstype(S) || return false + fields = KI.primitive_fields!(Any[], S, :val) + return all(((F, _),) -> sizeof(F) in (1, 2, 4, 8, 16), fields) + end + return shuffleable(T) +end + +# Whether a kernel's body contains a work-group collective, which needs the backend's +# sub-group capabilities: the `@kernel` macro adds a method taking a `CollectivesQuery` to the +# kernel's function, so that other kernels don't query them at every launch. (`applicable` is +# constant-folded.) It only does so for the first method of a kernel function: other methods +# with collectives use local memory. A callable that forwards its arguments, e.g. a wrapper of +# a kernel function, is taken to use collectives, which only costs the query of the +# capabilities at its launches. +struct CollectivesQuery end +uses_collectives(f) = applicable(f, CollectivesQuery()) diff --git a/src/macros.jl b/src/macros.jl index 3d19333c8..bce7dc8f7 100644 --- a/src/macros.jl +++ b/src/macros.jl @@ -70,6 +70,16 @@ function __kernel(expr, __source__::LineNumberNode, __module__::Module, force_in end gpu_function = combinedef(def_gpu) + # Properties of the kernel are methods of its function, rather than of a function of + # KernelAbstractions, so that a kernel can also be defined in a local scope. Like the + # constructors, they are defined with the first method of the kernel, since a method can't + # be overwritten during precompilation. + queries = Any[] + if find_collective(def[:body]) + # see `uses_collectives` + push!(queries, :($gpu_name(::$CollectivesQuery) = true)) + end + # create constructor functions _name = Symbol(:_, name) constructors = quote @@ -81,6 +91,7 @@ function __kernel(expr, __source__::LineNumberNode, __module__::Module, force_in $name(dev, size) = $_name(dev, $StaticSize(size), $DynamicSize()) $name(dev, size, range) = $_name(dev, $StaticSize(size), $StaticSize(range)) $name(dev, size::$_Size, range::$_Size) = $_name(dev, size, range) + $(queries...) end end constructors = relocate_lines(constructors, __source__) @@ -173,7 +184,8 @@ function transform_gpu!(def, constargs, force_inbounds, unsafe_indices) push!(let_constargs, :($arg = $constify($arg))) end end - pushfirst!(def[:args], :__ctx__) + # typed, so that the queries of `@kernel` (e.g. `uses_collectives`) never match the kernel + pushfirst!(def[:args], :(__ctx__::$CompilerMetadata)) # `Any[]`, since `split` hands back `LineNumberNode`s alongside `Expr`s new_stmts = Any[] body = MacroTools.flatten(def[:body]) @@ -244,10 +256,21 @@ function is_scope_construct(expr::Expr) # expr.head === :let end +function find_collective(stmt) + result = Ref(false) + postwalk(stmt) do expr + result[] |= is_collective(expr) + expr + end + return result[] +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[] @@ -279,6 +302,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) @@ -298,6 +332,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..e18e2ab3f --- /dev/null +++ b/test/groupreduce.jl @@ -0,0 +1,324 @@ +# 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) + i = @index(Global, Linear) + res = @groupreduce(op, x[i], neutral) + out[i] = res +end + +@kernel function groupreduce_bound!(out, @Const(x), op, neutral, ::Val{N}) where {N} + i = @index(Global, Linear) + res = @groupreduce(op, x[i], neutral; groupsize = N) + 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)) + i = @index(Global, Linear) + acc = zero(eltype(out)) + for k in 1:3 + res = @groupreduce(+, k * x[i], zero(eltype(out))) + acc += res + end + if true + m = @groupreduce max x[i] typemin(eltype(out)) + end + out[i] = acc + m +end + +@kernel function groupreduce_cartesian!(out, @Const(x)) + I = @index(Global, Cartesian) + res = @groupreduce(+, x[I], zero(eltype(out))) + out[I] = res +end + +@kernel unsafe_indices = true function groupreduce_unsafe!(out, @Const(x)) + i = @index(Global, Linear) + val = i <= length(x) ? x[i] : zero(eltype(out)) + res = @groupreduce(+, val, zero(eltype(out))) + if i <= length(out) + out[i] = res + end +end + +@kernel function groupreduce_noargs!() +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)) + i = @index(Global, Linear) + res = @groupreduce(+, x[i], Int64(0)) + out[i] = res +end + +@kernel function groupscan_mixed!(out, @Const(x)) + i = @index(Global, Linear) + res = @groupscan(+, x[i], 0) + out[i] = res +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; groupsize = 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 + +# 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 + +# A 32-bit float type that no backend shuffles, so that `@groupreduce` uses local memory +# only, also on backends with sub-groups. Its `+` adds the bits as integers. +primitive type TreeFloat <: AbstractFloat 32 end +TreeFloat(x::Integer) = reinterpret(TreeFloat, Int32(x)) +Base.:+(a::TreeFloat, b::TreeFloat) = reinterpret(TreeFloat, reinterpret(Int32, a) + reinterpret(Int32, b)) +Base.zero(::Type{TreeFloat}) = TreeFloat(0) +Base.typemin(::Type{TreeFloat}) = TreeFloat(typemin(Int32)) +Base.:(==)(a::TreeFloat, b::TreeFloat) = reinterpret(Int32, a) == reinterpret(Int32, b) + +# Run `f`, and return whether it failed to compile with GPUCompiler.jl#1004 (Metal: a phi of +# a by-reference argument and a device pointer, as in `x[i]` or the `neutral` argument), +# which the local-memory path runs into. +function gpucompiler_1004(f) + try + f() + return false + catch err + occursin("Invalid phi record", sprint(showerror, err)) || rethrow() + return true + end +end + +# `neutral` is evaluated once per work-item, also by the padding work-items +counted_neutral() = 0 +count_symbol(ex::Expr, sym) = sum(arg -> count_symbol(arg, sym), ex.args; init = 0) +count_symbol(ex, sym) = Int(ex === sym) + +function groupreduce_testsuite(backend, AT) + b = backend() + caps = KernelAbstractions.subgroup_capabilities(b) + + @testset "sub-group capabilities" begin + sub_groups = KI.supports_subgroups(b) && KI.supports_shuffle(b, UInt32) + @test (caps !== nothing) == sub_groups + shuffleable(T) = KernelAbstractions.__shuffleable(caps, T) + @test shuffleable(Int32) == sub_groups + @test shuffleable(Tuple{Int64, Bool}) == sub_groups + @test shuffleable(Float32) == (sub_groups && KI.supports_shuffle(b, Float32)) + @test shuffleable(Tuple{Float32, Int32}) == (sub_groups && KI.supports_shuffle(b, Float32)) + @test !shuffleable(TreeFloat) + @test !shuffleable(Ref{Int}) + # kernels without collectives don't query the capabilities + @test KernelAbstractions.kernel_subgroups(groupreduce_static!(b, 64)) == caps + @test KernelAbstractions.kernel_subgroups(groupscan!(b, 64)) == caps + @test KernelAbstractions.uses_collectives(groupreduce_static!(b, 64).f) + # also for a kernel without arguments, whose function takes a single argument + @test !KernelAbstractions.uses_collectives(groupreduce_noargs!(b).f) + end + + types = (Int32, Int64, Float32, TreeFloat) + @testset "$T, $(nameof(typeof(op)))" for T in types, (op, neutral) in ((+, zero(T)), (max, typemin(T))) + T === TreeFloat && op === max && 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(fill(zero(T), n)) + if gpucompiler_1004(() -> groupreduce_static!(b, groupsize)(out, AT(x), op, neutral; ndrange = n)) + @test_broken false + continue + end + @test Array(out) == groupwise(op, x, groupsize) + + fill!(out, zero(T)) + if gpucompiler_1004(() -> groupreduce_bound!(b)(out, AT(x), op, neutral, Val(256); ndrange = n, workgroupsize = groupsize)) + @test_broken false + continue + end + @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; 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); 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); 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); ndrange = 128) + @test Array(out) == groupwise(+, x, 64) + end + + @testset "mixed types" begin + x = Int32.(rand(1:100, 100)) + out = AT(zeros(Int64, 100)) + groupreduce_mixed!(b, 64)(out, AT(x); ndrange = 100) + @test Array(out) == groupwise(+, Int64.(x), 64) + 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 + + @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 + @test_throws "given more than once" @macroexpand @kernel function f(y, x) + res = @groupreduce(+, x[1], 0; groupsize = 32, groupsize = 64) + end + # the upper bound of the workgroup size is a keyword + @test_throws "expects `op`, `val` and `neutral`" @macroexpand @kernel function f(y, x) + res = @groupreduce(+, x[1], 0, 64) + end + @test_throws "unknown option" @macroexpand @kernel function f(y, x) + res = @groupreduce(+, x[1], 0; subgroups = true) + end + end + + @testset "kernels in a local scope" begin + # the kernel's properties are methods of its function, which works in a local scope too + function local_groupreduce(b, AT, x) + @kernel function local_groupreduce!(out, @Const(x)) + i = @index(Global, Linear) + res = @groupreduce(+, x[i], zero(eltype(out))) + out[i] = res + end + out = AT(zeros(eltype(x), length(x))) + kernel = local_groupreduce!(b, 32) + @test KernelAbstractions.uses_collectives(kernel.f) + kernel(out, AT(x); ndrange = length(x)) + return Array(out) + end + x = Int32.(rand(1:100, 100)) + @test local_groupreduce(b, AT, x) == groupwise(+, x, 32) + end + + @testset "neutral evaluated once" begin + ex = @macroexpand @kernel function f(y, x) + i = @index(Global, Linear) + res = @groupreduce(+, x[i], counted_neutral()) + y[i] = res + end + @test count_symbol(ex, :counted_neutral) == 1 + end + return +end diff --git a/test/runtests.jl b/test/runtests.jl index ffb17131d..e2def3965 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -300,6 +300,22 @@ end end end +# `@groupreduce` reduces sub-groups with PoCL's native collectives +@kernel function groupreduce_codegen!(out, @Const(x)) + i = @index(Global, Linear) + res = @groupreduce(+, x[i], zero(eltype(out))) + out[i] = res +end +@testset "POCL @groupreduce with sub-groups" begin + x = ones(Int32, 100) + out = zeros(Int32, 100) + ir = sprint() do io + @device_code_llvm io = io debuginfo = :none groupreduce_codegen!(POCLBackend(), 64)(out, x; ndrange = 100) + end + @test out == [fill(64, 64); fill(36, 36)] + @test occursin("sub_group_reduce_add", ir) +end + # Julia doesn't turn a splat of more than 32 elements into a direct call, so a launch with # many arguments allocates unless every layer passes them on as a tuple @testset "POCL launch with many arguments" begin diff --git a/test/testsuite.jl b/test/testsuite.jl index 541028ad8..375407df2 100644 --- a/test/testsuite.jl +++ b/test/testsuite.jl @@ -45,6 +45,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 @@ -123,6 +124,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