From 1e46c0a138d9634c671dee519b82b9b2657ea114 Mon Sep 17 00:00:00 2001 From: Tim Besard Date: Wed, 30 Sep 2026 21:34:47 +0200 Subject: [PATCH] KernelInterface: don't promise how work-items form sub-groups KernelInterface promised that a work-group is divided into sub-groups of the sub-group width with at most one partial sub-group, i.e. `cld(prod(get_local_size()), get_max_sub_group_size())` of them. Only CUDA, HSA (AMD) and Metal specify linear packing of multi-dimensional work-groups. OpenCL requires that only the highest-numbered sub-group is smaller, but leaves the mapping implementation-defined, and SYCL, Vulkan and WebGPU don't specify either. Implementations differ: Intel's CPU OpenCL runtime forms sub-groups per row of a work-group by default, so a 33x2 work-group has four sub-groups of 32 and 1 work-items, and Intel's GPU runtime picks a hardware walk order that isn't x-fastest for some shapes. Require only what backends can ensure everywhere: unique and invariant (sub-group, lane) pairs, dense ids and lanes, and that a 1-D work-group of at most the sub-group width is a single sub-group. `get_num_sub_groups` is at least the count full sub-groups would need, but can be larger, so storage for a value per sub-group has to allow for one per work-item. The testsuite checks those invariants over 1-D, 2-D and 3-D work-groups whose size is or isn't a multiple of the width, and combines a value per sub-group through local memory, as reductions do. The shuffle test no longer assumes that a work-group of twice the width consists of two full sub-groups. --- docs/src/kernelinterface.md | 22 ++++- lib/KernelInterface/src/device.jl | 29 ++++-- lib/KernelInterface/src/host.jl | 4 +- lib/KernelInterface/test/interface.jl | 133 ++++++++++++++++---------- 4 files changed, 122 insertions(+), 66 deletions(-) diff --git a/docs/src/kernelinterface.md b/docs/src/kernelinterface.md index ca72776c2..eab4c4ab8 100644 --- a/docs/src/kernelinterface.md +++ b/docs/src/kernelinterface.md @@ -149,10 +149,24 @@ get_global_size ### Sub-groups Sub-groups are optional ([`supports_subgroups`](@ref)). A work-group is divided into -sub-groups of [`sub_group_size(backend)`](@ref sub_group_size) work-items, the last of which -can be partial. How work-items are assigned to sub-groups is unspecified, but every -work-item has a unique `(get_sub_group_id(), get_sub_group_local_id())` pair in its -work-group, which doesn't change during the kernel. +sub-groups of at most [`sub_group_size(backend)`](@ref sub_group_size) work-items. Which +work-items form a sub-group, how many sub-groups there are, and which of them are partial +is unspecified, and differs between devices and work-group shapes. For example, CUDA forms +warps from consecutive linear work-item indices, while Intel's CPU OpenCL runtime forms +sub-groups per row of a multi-dimensional work-group, so that a 33×2 work-group consists of +four sub-groups of 32 and 1 work-items. What KernelInterface guarantees, and backends that +report sub-group support have to ensure: + +- every work-item has a unique `(get_sub_group_id(), get_sub_group_local_id())` pair in its + work-group, which doesn't change during the kernel; +- the sub-group ids are `1:get_num_sub_groups()`, and the lanes of a sub-group are + `1:get_sub_group_size()`; +- a 1-D work-group of at most `sub_group_size(backend)` work-items is a single sub-group. + +In particular, [`get_num_sub_groups`](@ref) can be larger than +`cld(prod(get_local_size()), get_max_sub_group_size())`. Storage for a value per sub-group +has to be sized for up to one sub-group per work-item, and code combining those values has +to use `get_num_sub_groups()` rather than compute the count. ```@docs get_sub_group_size diff --git a/lib/KernelInterface/src/device.jl b/lib/KernelInterface/src/device.jl index 40081923f..5b0090914 100644 --- a/lib/KernelInterface/src/device.jl +++ b/lib/KernelInterface/src/device.jl @@ -126,16 +126,20 @@ end ## sub-groups # Sub-group support is optional, see `supports_subgroups`. A work-group is divided into -# sub-groups of `sub_group_size(backend)` work-items. How work-items are assigned to -# sub-groups is unspecified, except that `(get_sub_group_id(), get_sub_group_local_id())` -# is unique within a work-group and doesn't change during the kernel's execution. +# sub-groups of at most `sub_group_size(backend)` work-items. Which work-items form a +# sub-group, how many sub-groups there are and which are partial is unspecified, except +# that `(get_sub_group_id(), get_sub_group_local_id())` is unique within a work-group and +# doesn't change during the kernel's execution, and that a 1-D work-group of at most +# `sub_group_size(backend)` work-items is a single sub-group. Backends that can't ensure +# that don't report sub-group support. See the manual. """ get_sub_group_size([::Type{T}=Int])::T -The number of work-items in the sub-group: the sub-group width -([`get_max_sub_group_size`](@ref)), or fewer for the last sub-group of a work-group whose -size isn't a multiple of the width. +The number of work-items in the sub-group, at most the sub-group width +([`get_max_sub_group_size`](@ref)). Which sub-groups have fewer work-items than the width +is unspecified: when the work-group size isn't a multiple of the width, there can be more +than one, e.g. one per row of a multi-dimensional work-group. See [`get_local_id`](@ref) for the supported types `T`. @@ -167,7 +171,10 @@ See [`get_local_id`](@ref) for the supported types `T`. """ get_num_sub_groups([::Type{T}=Int])::T -The number of sub-groups in the work-group: `cld(prod(get_local_size()), get_max_sub_group_size())`. +The number of sub-groups in the work-group. It is at least +`cld(prod(get_local_size()), get_max_sub_group_size())`, but can be larger, since more than +one sub-group can be partial. Size storage for a value per sub-group for up to one +sub-group per work-item. See [`get_local_id`](@ref) for the supported types `T`. @@ -183,7 +190,8 @@ See [`get_local_id`](@ref) for the supported types `T`. """ get_sub_group_id([::Type{T}=Int])::T -The 1-based index of the sub-group within the work-group. +The 1-based index of the sub-group within the work-group, between 1 and +[`get_num_sub_groups`](@ref). How it relates to [`get_local_id`](@ref) is unspecified. See [`get_local_id`](@ref) for the supported types `T`. @@ -199,8 +207,9 @@ See [`get_local_id`](@ref) for the supported types `T`. """ get_sub_group_local_id([::Type{T}=Int])::T -The 1-based index of the work-item within its sub-group (its lane). It doesn't depend on -which work-items of the sub-group are active, e.g. in a divergent branch. +The 1-based index of the work-item within its sub-group (its lane), between 1 and +[`get_sub_group_size`](@ref). It doesn't depend on which work-items of the sub-group are +active, e.g. in a divergent branch. See [`get_local_id`](@ref) for the supported types `T`. diff --git a/lib/KernelInterface/src/host.jl b/lib/KernelInterface/src/host.jl index 09f55238e..b769783cd 100644 --- a/lib/KernelInterface/src/host.jl +++ b/lib/KernelInterface/src/host.jl @@ -263,7 +263,9 @@ supports_float64(::Backend) = false Whether kernels on the active device support sub-groups: the sub-group queries ([`get_sub_group_size`](@ref) etc.), [`sub_group_barrier`](@ref), and a fixed sub-group -width [`sub_group_size`](@ref). +width [`sub_group_size`](@ref). See the manual for what KernelInterface guarantees about +how work-groups are divided into sub-groups; a backend that can't ensure that reports +`false`. Which types [`shfl_down`](@ref) supports is queried separately with [`supports_shuffle`](@ref). diff --git a/lib/KernelInterface/test/interface.jl b/lib/KernelInterface/test/interface.jl index 2e5366689..210b77123 100644 --- a/lib/KernelInterface/test/interface.jl +++ b/lib/KernelInterface/test/interface.jl @@ -152,7 +152,7 @@ end function test_subgroup_kernel(results) l = KI.get_local_id() s = KI.get_local_size() - i = (l.y - 1) * s.x + l.x + (KI.get_group_id().x - 1) * s.x * s.y + i = ((l.z - 1) * s.y + (l.y - 1)) * s.x + l.x + (KI.get_group_id().x - 1) * s.x * s.y * s.z if i <= length(results) @inbounds results[i] = SubgroupData( @@ -166,6 +166,26 @@ function test_subgroup_kernel(results) return end +# Combine a value per sub-group through local memory, as reductions do: the first lane of +# every sub-group stores its size, and the first work-item adds up `get_num_sub_groups()` +# of them. `N` is the work-group size, which bounds the number of sub-groups. +function subgroup_combine_kernel(out, ::Val{N}) where {N} + partial = KI.localmemory(Int32, N) + if KI.get_sub_group_local_id() == 1 + @inbounds partial[KI.get_sub_group_id()] = KI.get_sub_group_size() + end + KI.barrier() + l = KI.get_local_id() + if l.x == 1 && l.y == 1 && l.z == 1 + total = Int32(0) + for i in 1:KI.get_num_sub_groups() + @inbounds total += partial[i] + end + @inbounds out[KI.get_group_id().x] = total + end + return +end + function subgroup_typecheck_kernel(results, val::T) where {T} # uniformly executed by the whole sub-group, as `shfl_down` requires shuffled = KI.shfl_down(val, 1) @@ -592,67 +612,73 @@ function interface_testsuite(backend::KI.Backend, AT) end end - # checks the sub-groups of a work-group of `items` work-items - function check_subgroups(data, items) + # checks the sub-groups of a work-group of shape `dims`. Which work-items form a + # sub-group, and how many sub-groups there are, is unspecified. + function check_subgroups(data, dims) + items = prod(dims) @test all(d -> d.max_sub_group_size == sg_size, data) - @test all(d -> d.num_sub_groups == cld(items, sg_size), data) - @test all(d -> 1 <= d.sub_group_id <= cld(items, sg_size), data) - @test all(d -> 1 <= d.sub_group_local_id <= d.sub_group_size, data) - # every work-item has its own (sub-group, lane) pair - @test allunique(map(d -> (d.sub_group_id, d.sub_group_local_id), data)) - # each sub-group has as many members as its size says - for id in unique(map(d -> d.sub_group_id, data)) + # all work-items agree on the number of sub-groups, which is at least what full + # sub-groups would need + n = first(data).num_sub_groups + @test all(d -> d.num_sub_groups == n, data) + @test cld(items, sg_size) <= n <= items + # the sub-group ids are 1:n + @test sort(unique(map(d -> d.sub_group_id, data))) == 1:n + # every sub-group has as many members as its size says, and they are its lanes + for id in 1:n members = filter(d -> d.sub_group_id == id, data) - @test all(d -> d.sub_group_size == length(members), members) + size = length(members) + @test 1 <= size <= sg_size + @test all(d -> d.sub_group_size == size, members) + @test sort(map(d -> d.sub_group_local_id, members)) == 1:size + end + # a 1-D work-group that fits a sub-group is one + if length(dims) == 1 && items <= sg_size + @test n == 1 end return end - @testset "Sub-groups" begin - sg_n = 2 - workgroupsize = sg_size * sg_n - numgroups = 2 - N = workgroupsize * numgroups - - results = AT(Vector{SubgroupData}(undef, N)) - kernel = KI.@launch backend launch = false test_subgroup_kernel(results) - if fits(kernel, (workgroupsize,)) - kernel(results; workgroupsize, numgroups) - KI.synchronize(backend) + # work-group shapes to check: 1-D ones around the width, and multi-dimensional ones + # whose first dimension is or isn't a multiple of the width + subgroup_shapes = unique( + [ + (sg_size - 1,), (sg_size,), (sg_size + 1,), (2 * sg_size,), (2 * sg_size + 1,), + (33, 2), (sg_size, 2), (sg_size + 1, 2), (7, 5), (5, 3, 2), (sg_size, 2, 2), + ] + ) + filter!(dims -> all(>(0), dims), subgroup_shapes) - host_results = Array(results) - @test all(d -> d.sub_group_size == sg_size, host_results) - for group in Iterators.partition(host_results, workgroupsize) - check_subgroups(collect(group), workgroupsize) + @testset "Sub-group formation" begin + numgroups = 2 + @testset "$dims" for dims in subgroup_shapes + items = prod(dims) + results = AT(Vector{SubgroupData}(undef, items * numgroups)) + kernel = KI.@launch backend launch = false test_subgroup_kernel(results) + if fits(kernel, dims) + kernel(results; workgroupsize = dims, numgroups) + KI.synchronize(backend) + for group in Iterators.partition(Array(results), items) + check_subgroups(collect(group), dims) + end + else + @test_skip "work-groups of $dims work-items" end - else - @test_skip "work-groups of $workgroupsize work-items" end end - @testset "Partial sub-groups" begin - # a 2-D work-group whose size isn't a multiple of the sub-group size, or else a - # 1-D one with one work-item more than a sub-group - numgroups = 2 - results = AT(Vector{SubgroupData}(undef, max(66, sg_size + 1) * numgroups)) - kernel = KI.@launch backend launch = false test_subgroup_kernel(results) - workgroupsize = fits(kernel, (33, 2)) ? (33, 2) : (sg_size + 1,) - items = prod(workgroupsize) - if fits(kernel, workgroupsize) && items % sg_size != 0 - kernel(results; workgroupsize, numgroups) - KI.synchronize(backend) - - host_results = Array(results)[1:(items * numgroups)] - for group in Iterators.partition(host_results, items) - group = collect(group) - check_subgroups(group, items) - # the sizes of the sub-groups add up to the work-group, with one partial one - sizes = Dict(d.sub_group_id => d.sub_group_size for d in group) - @test sum(values(sizes)) == items - @test count(<(sg_size), values(sizes)) == (items % sg_size == 0 ? 0 : 1) + @testset "Combining a value per sub-group" begin + numgroups = 3 + @testset "$dims" for dims in subgroup_shapes + out = KI.zeros(backend, Int32, numgroups) + kernel = KI.@launch backend launch = false subgroup_combine_kernel(out, Val(prod(dims))) + if fits(kernel, dims) + kernel(out, Val(prod(dims)); workgroupsize = dims, numgroups) + KI.synchronize(backend) + @test all(==(prod(dims)), Array(out)) + else + @test_skip "work-groups of $dims work-items" end - else - @test_skip "work-groups of $workgroupsize work-items" end end @@ -697,7 +723,12 @@ function interface_testsuite(backend::KI.Backend, AT) KI.synchronize(backend) out = Array(out) in_range = findall(i -> out[i, 1] + offset <= out[i, 2], 1:N) - @test length(in_range) == N - (N ÷ sg_size) * offset + # a work-group of one sub-group's width is a single sub-group; with more + # work-items, how many are in range depends on the sub-groups' sizes + if N == sg_size + @test length(in_range) == N - offset + end + @test !isempty(in_range) @test out[in_range, 3] == out[in_range, 1] .+ offset end end