Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 18 additions & 4 deletions docs/src/kernelinterface.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
29 changes: 19 additions & 10 deletions lib/KernelInterface/src/device.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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`.

Expand Down Expand Up @@ -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`.

Expand All @@ -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`.

Expand All @@ -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`.

Expand Down
4 changes: 3 additions & 1 deletion lib/KernelInterface/src/host.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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).

Expand Down
133 changes: 82 additions & 51 deletions lib/KernelInterface/test/interface.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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)
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down
Loading