From 55ec9881b1ed734f871748822411f2955f9d0001 Mon Sep 17 00:00:00 2001 From: Valentin Churavy Date: Sat, 3 Oct 2026 21:25:29 +0200 Subject: [PATCH 1/5] Add @groupreduce and @subgroupreduce Rework of #559 on top of KernelInterface: - `@groupreduce(op, val, neutral[, groupsize]; subgroups=false)` reduces over the workgroup and returns the result on every work-item. It uses a local-memory tree by default, or a two-level reduction based on `KI.shfl_down` with `subgroups=true` (gated on the host by `KI.supports_shuffle`). The local memory is sized by the static workgroup size or an explicit upper bound. - `@subgroupreduce(op, val, neutral)` reduces over the sub-group with shuffles; the result is defined on the first lane. - Both are collectives in `@kernel`: the split treats them like `@synchronize`, and padding work-items contribute `neutral` without evaluating `val`, so ndranges that are not a multiple of the workgroup size work. Co-authored-by: Anton Smirnov Assisted-by: Claude Code (Opus 5.5) --- docs/src/api.md | 7 ++ src/KernelAbstractions.jl | 3 + src/groupreduction.jl | 231 ++++++++++++++++++++++++++++++++++++++ src/macros.jl | 16 ++- test/groupreduce.jl | 129 +++++++++++++++++++++ test/testsuite.jl | 5 + 6 files changed, 390 insertions(+), 1 deletion(-) create mode 100644 src/groupreduction.jl create mode 100644 test/groupreduce.jl diff --git a/docs/src/api.md b/docs/src/api.md index 826cd7455..aeb1fcacb 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -15,6 +15,13 @@ @ndrange ``` +### Reductions + +```@docs +@groupreduce +@subgroupreduce +``` + ## Host language !!! note diff --git a/src/KernelAbstractions.jl b/src/KernelAbstractions.jl index f549aec38..378911c83 100644 --- a/src/KernelAbstractions.jl +++ b/src/KernelAbstractions.jl @@ -37,6 +37,8 @@ and then invoked on the arguments. - [`@uniform`](@ref) - [`@synchronize`](@ref) - [`@print`](@ref) +- [`@groupreduce`](@ref) +- [`@subgroupreduce`](@ref) # Kernel constructor @@ -633,6 +635,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..250240018 --- /dev/null +++ b/src/groupreduction.jl @@ -0,0 +1,231 @@ +### +# Group and sub-group reductions +# - @groupreduce +# - @subgroupreduce +### + +export @groupreduce, @subgroupreduce + +""" + @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 quote + $__groupreduce( + $(esc(:__ctx__)), $(esc(op)), $(esc(val)), $(esc(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`, using +[`KernelInterface.shfl_down`](@ref). The result is only defined on the first work-item of +the sub-group (`KernelInterface.get_sub_group_local_id() == 1`). `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). +""" +macro subgroupreduce(op, val, neutral) + return :($__subgroupreduce($(esc(op)), $(esc(val)), $(esc(neutral)))) +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")) + +# 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( + "@groupreduce requires 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 + +# Combine contiguous ranges of lanes of doubling length, so that the first lane ends up with +# the reduction of the sub-group. Only requires `op` to be associative. +@inline function __subgroupreduce(op, val, neutral::T) where {T} + val = convert(T, val) + lane = KI.get_sub_group_local_id() + sgsize = KI.get_sub_group_size() + offset = 1 + while offset < sgsize + other = KI.shfl_down(val, offset) + # the result of shuffling from past the end of the sub-group is unspecified + if lane + offset <= sgsize + val = op(val, other) + end + offset <<= 1 + 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..97fd3293b --- /dev/null +++ b/test/groupreduce.jl @@ -0,0 +1,129 @@ +# 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 + +@kernel function subgroupreduce!(out, @Const(x)) + i = @index(Global, Linear) + res = @subgroupreduce(+, x[i], zero(eltype(out))) + if KernelInterface.get_sub_group_local_id() == 1 + out[i] = res + end +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 + +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 "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 + 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)) + subgroupreduce!(b, groupsize)(out, AT(x); ndrange = n) + ref = fill(-1.0f0, n) + for i in 1:width:n + ref[i] = sum(x[i:min(i + width - 1, n)]) + end + @test Array(out) == ref + end + 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 From 8f640ecdd2a6a91381f006a45ce8109198cd8db4 Mon Sep 17 00:00:00 2001 From: Valentin Churavy Date: Sun, 4 Oct 2026 00:49:22 +0200 Subject: [PATCH 2/5] Add @groupscan and @subgroupscan - `@groupscan(op, val, neutral[, groupsize]; inclusive = true)` scans over the workgroup in the order of the local linear index, with a Hillis-Steele scan in double-buffered local memory, sized like `@groupreduce`. - `@subgroupscan(op, val, neutral; inclusive = true)` scans over the lanes of a sub-group with `KI.shfl_up`. Both only need `op` to be associative, and are collectives in `@kernel` like the reductions: padding work-items contribute `neutral`. Assisted-by: Claude Code (Opus 5.5) --- docs/src/api.md | 4 +- src/KernelAbstractions.jl | 2 + src/groupreduction.jl | 150 +++++++++++++++++++++++++++++++++++++- test/groupreduce.jl | 123 +++++++++++++++++++++++++++++++ 4 files changed, 274 insertions(+), 5 deletions(-) diff --git a/docs/src/api.md b/docs/src/api.md index aeb1fcacb..f7c8e9e28 100644 --- a/docs/src/api.md +++ b/docs/src/api.md @@ -15,11 +15,13 @@ @ndrange ``` -### Reductions +### Reductions and scans ```@docs @groupreduce @subgroupreduce +@groupscan +@subgroupscan ``` ## Host language diff --git a/src/KernelAbstractions.jl b/src/KernelAbstractions.jl index 378911c83..e63459f37 100644 --- a/src/KernelAbstractions.jl +++ b/src/KernelAbstractions.jl @@ -39,6 +39,8 @@ and then invoked on the arguments. - [`@print`](@ref) - [`@groupreduce`](@ref) - [`@subgroupreduce`](@ref) +- [`@groupscan`](@ref) +- [`@subgroupscan`](@ref) # Kernel constructor diff --git a/src/groupreduction.jl b/src/groupreduction.jl index 250240018..5cc734532 100644 --- a/src/groupreduction.jl +++ b/src/groupreduction.jl @@ -1,10 +1,12 @@ ### -# Group and sub-group reductions +# Group and sub-group reductions and scans # - @groupreduce # - @subgroupreduce +# - @groupscan +# - @subgroupscan ### -export @groupreduce, @subgroupreduce +export @groupreduce, @subgroupreduce, @groupscan, @subgroupscan """ @groupreduce(op, val, neutral[, groupsize]; subgroups = false) @@ -78,6 +80,87 @@ macro subgroupreduce(op, val, neutral) return :($__subgroupreduce($(esc(op)), $(esc(val)), $(esc(neutral)))) 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 quote + $__groupscan( + $(esc(:__ctx__)), $(esc(op)), $(esc(val)), $(esc(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()`, using [`KernelInterface.shfl_up`](@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 :($__subgroupscan($(esc(op)), $(esc(val)), $(esc(neutral)), Val($(esc(inclusive))))) +end + # Separate `key = value` options (also after a `;`) from the positional macro arguments. function split_options(args, keys) positional = Any[] @@ -100,7 +183,10 @@ function split_options(args, keys) return positional, options end -const COLLECTIVES = (Symbol("@groupreduce"), Symbol("@subgroupreduce")) +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) @@ -144,7 +230,7 @@ end @inline __static_groupsize(::NDRange{N, B, W}) where {N, B, W <: StaticSize} = Val(prod(get(W))) @inline __static_groupsize(::NDRange) = throw( ArgumentError( - "@groupreduce requires a static workgroup size or an upper bound of it" + "group reductions and scans require a static workgroup size or an upper bound of it" ) ) @@ -229,3 +315,59 @@ end end return val 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 + +# Hillis-Steele with shuffles: after the step with `offset`, every lane holds the scan of the +# (up to) `2offset` lanes ending at its own. +@inline function __subgroupscan(op, val, neutral::T, ::Val{inclusive}) where {T, inclusive} + val = convert(T, val) + lane = KI.get_sub_group_local_id() + sgsize = KI.get_sub_group_size() + offset = 1 + while offset < sgsize + other = KI.shfl_up(val, offset) + # the result of shuffling from before the start of the sub-group is unspecified + if lane > offset + val = op(other, val) + end + offset <<= 1 + end + if !inclusive + prev = KI.shfl_up(val, 1) + val = lane == 1 ? neutral : prev + end + return val +end diff --git a/test/groupreduce.jl b/test/groupreduce.jl index 97fd3293b..4f73117ce 100644 --- a/test/groupreduce.jl +++ b/test/groupreduce.jl @@ -48,11 +48,69 @@ end end 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,) @@ -116,6 +174,71 @@ function groupreduce_testsuite(backend, AT) 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 + 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 + end + @testset "errors" begin @test_throws "must be used as a statement" @macroexpand @kernel function f(y, x) i = @index(Global) From 7de50d7c3499bce97a278935b7387e7f367e5332 Mon Sep 17 00:00:00 2001 From: Valentin Churavy Date: Sun, 4 Oct 2026 10:34:35 +0200 Subject: [PATCH 3/5] Build the sub-group reductions and scans on KernelInterface's `@subgroupreduce` and `@subgroupscan`, and the sub-group stage of `@groupreduce`, now use `KI.sub_group_reduce` and `KI.sub_group_scan`, which backends can implement with native operations. `@subgroupreduce` returns the result on every work-item of the sub-group. Test reductions of (value, index) pairs, and don't assume that sub-groups are formed from consecutive work-items. Assisted-by: Claude Code (Opus 5.5) --- src/groupreduction.jl | 46 ++++++++++--------------------------------- test/groupreduce.jl | 30 ++++++++++++++++++---------- 2 files changed, 30 insertions(+), 46 deletions(-) diff --git a/src/groupreduction.jl b/src/groupreduction.jl index 5cc734532..6197e6289 100644 --- a/src/groupreduction.jl +++ b/src/groupreduction.jl @@ -64,10 +64,10 @@ end """ @subgroupreduce(op, val, neutral) -Reduce `val` over the work-items of the sub-group with the binary operator `op`, using -[`KernelInterface.shfl_down`](@ref). The result is only defined on the first work-item of -the sub-group (`KernelInterface.get_sub_group_local_id() == 1`). `op` has to be associative, -and `neutral` its neutral element. The result has the type of `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: @@ -138,7 +138,7 @@ 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()`, using [`KernelInterface.shfl_up`](@ref). +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`. @@ -298,23 +298,8 @@ end return @inbounds storage[1] end -# Combine contiguous ranges of lanes of doubling length, so that the first lane ends up with -# the reduction of the sub-group. Only requires `op` to be associative. -@inline function __subgroupreduce(op, val, neutral::T) where {T} - val = convert(T, val) - lane = KI.get_sub_group_local_id() - sgsize = KI.get_sub_group_size() - offset = 1 - while offset < sgsize - other = KI.shfl_down(val, offset) - # the result of shuffling from past the end of the sub-group is unspecified - if lane + offset <= sgsize - val = op(val, other) - end - offset <<= 1 - end - return val -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)) @@ -350,22 +335,11 @@ end return res end -# Hillis-Steele with shuffles: after the step with `offset`, every lane holds the scan of the -# (up to) `2offset` lanes ending at its own. +# backends can implement `KI.sub_group_scan` with native operations @inline function __subgroupscan(op, val, neutral::T, ::Val{inclusive}) where {T, inclusive} - val = convert(T, val) - lane = KI.get_sub_group_local_id() - sgsize = KI.get_sub_group_size() - offset = 1 - while offset < sgsize - other = KI.shfl_up(val, offset) - # the result of shuffling from before the start of the sub-group is unspecified - if lane > offset - val = op(other, val) - end - offset <<= 1 - end + 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 diff --git a/test/groupreduce.jl b/test/groupreduce.jl index 4f73117ce..627a09f7c 100644 --- a/test/groupreduce.jl +++ b/test/groupreduce.jl @@ -40,12 +40,12 @@ end end end -@kernel function subgroupreduce!(out, @Const(x)) +# 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))) - if KernelInterface.get_sub_group_local_id() == 1 - out[i] = res - end + out[i] = res + sgs[i] = (@index(Group, Linear), KernelInterface.get_sub_group_id()) end # the composition of affine maps `x -> a * x + b`, first `f` then `g`: associative, but not @@ -131,6 +131,17 @@ function groupreduce_testsuite(backend, AT) 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)) @@ -164,12 +175,11 @@ function groupreduce_testsuite(backend, AT) for (groupsize, n) in ((width, 4width), (2width, 2width + 5)) x = Float32.(rand(1:100, n)) out = AT(fill(-1.0f0, n)) - subgroupreduce!(b, groupsize)(out, AT(x); ndrange = n) - ref = fill(-1.0f0, n) - for i in 1:width:n - ref[i] = sum(x[i:min(i + width - 1, n)]) - end - @test Array(out) == ref + 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) end end end From 4440c72f39d57f6d32c1903b66681ece9ad4c313 Mon Sep 17 00:00:00 2001 From: Valentin Churavy Date: Sun, 4 Oct 2026 12:19:27 +0200 Subject: [PATCH 4/5] Document KernelInterface's sub-group functions in @kernel KernelAbstractions' collectives are executed by the padding work-items of a partial workgroup, but direct calls of KernelInterface's sub-group functions aren't: kernels that use them need `unsafe_indices=true`. Assisted-by: Claude Code (Opus 5.5) --- docs/src/kernels.md | 30 ++++++++++++++++++++++++++++++ src/groupreduction.jl | 6 ++++++ 2 files changed, 36 insertions(+) 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/groupreduction.jl b/src/groupreduction.jl index 6197e6289..0cb287690 100644 --- a/src/groupreduction.jl +++ b/src/groupreduction.jl @@ -75,6 +75,12 @@ reached by all work-items of the sub-group, and has to be used as a statement on 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 :($__subgroupreduce($(esc(op)), $(esc(val)), $(esc(neutral)))) From 9b2032063a56aef56f1f4a0fd77b26564979e665 Mon Sep 17 00:00:00 2001 From: Valentin Churavy Date: Sun, 4 Oct 2026 14:05:10 +0200 Subject: [PATCH 5/5] Convert the value of a collective before calling it The type of the value passed to `@groupreduce`, `@subgroupreduce`, `@groupscan` or `@subgroupscan` may differ between the work-items: in a `@kernel` the padding work-items contribute `neutral` instead of `val`, so `@groupreduce(+, x[i]::Float32, 0.0)` reduces a `Union{Float32, Float64}`, and an accumulator that only some work-items add a `Float64` to is a `Union` as well. Julia union-splits the call of the collective with such an argument into one call per type, so the work-items of a workgroup executed different copies of its barriers and shuffles. On PoCL this silently gave wrong results (0 for the reduction of a padded workgroup, NaN for the energy of a Float32 system with a Float64 Coulomb constant in Molly). Convert the value to the type of `neutral` at the call site instead, so that only the conversion is union-split, and test mixed types for all four collectives. Assisted-by: Claude Code (Opus 5.5) --- src/groupreduction.jl | 41 ++++++++++++++++++------ test/groupreduce.jl | 73 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 104 insertions(+), 10 deletions(-) diff --git a/src/groupreduction.jl b/src/groupreduction.jl index 0cb287690..d98489c6b 100644 --- a/src/groupreduction.jl +++ b/src/groupreduction.jl @@ -53,10 +53,11 @@ macro groupreduce(args...) bound = groupsize === nothing ? :($__static_groupsize($(esc(:__ctx__)))) : :(Val($(esc(groupsize)))) subgroups = Base.get(options, :subgroups, false) - return quote - $__groupreduce( - $(esc(:__ctx__)), $(esc(op)), $(esc(val)), $(esc(neutral)), - $bound, Val($(esc(subgroups))), + return __collective_call(neutral, val) do neutral, val + :( + $__groupreduce( + $(esc(:__ctx__)), $(esc(op)), $val, $neutral, $bound, Val($(esc(subgroups))), + ) ) end end @@ -83,7 +84,9 @@ It must only be used on backends that support shuffles of the type of `neutral`, `@kernel unsafe_indices=true` for such kernels. """ macro subgroupreduce(op, val, neutral) - return :($__subgroupreduce($(esc(op)), $(esc(val)), $(esc(neutral)))) + return __collective_call(neutral, val) do neutral, val + :($__subgroupreduce($(esc(op)), $val, $neutral)) + end end """ @@ -132,10 +135,11 @@ macro groupscan(args...) bound = groupsize === nothing ? :($__static_groupsize($(esc(:__ctx__)))) : :(Val($(esc(groupsize)))) inclusive = Base.get(options, :inclusive, true) - return quote - $__groupscan( - $(esc(:__ctx__)), $(esc(op)), $(esc(val)), $(esc(neutral)), - $bound, Val($(esc(inclusive))), + return __collective_call(neutral, val) do neutral, val + :( + $__groupscan( + $(esc(:__ctx__)), $(esc(op)), $val, $neutral, $bound, Val($(esc(inclusive))), + ) ) end end @@ -164,7 +168,24 @@ macro subgroupscan(args...) length(positional) == 3 || error("@subgroupscan expects `op`, `val` and `neutral`") op, val, neutral = positional inclusive = Base.get(options, :inclusive, true) - return :($__subgroupscan($(esc(op)), $(esc(val)), $(esc(neutral)), Val($(esc(inclusive))))) + 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. diff --git a/test/groupreduce.jl b/test/groupreduce.jl index 627a09f7c..ff9e8f893 100644 --- a/test/groupreduce.jl +++ b/test/groupreduce.jl @@ -48,6 +48,41 @@ end 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]) @@ -167,6 +202,13 @@ function groupreduce_testsuite(backend, AT) 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) @@ -180,6 +222,12 @@ function groupreduce_testsuite(backend, AT) 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 @@ -222,6 +270,13 @@ function groupreduce_testsuite(backend, AT) 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) @@ -247,6 +302,24 @@ function groupreduce_testsuite(backend, AT) @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