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
2 changes: 2 additions & 0 deletions lib/intrinsics/src/SPIRVIntrinsics.jl
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,8 @@ include("math.jl")
include("integer.jl")
include("atomic.jl")
include("shuffle.jl")
include("vote.jl")
include("collective.jl")

# helper macro to import all names from this package, even non-exported ones.
macro import_all()
Expand Down
29 changes: 29 additions & 0 deletions lib/intrinsics/src/collective.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
# Sub-group collectives from `cl_khr_subgroups`, which the SPIR-V back-end lowers to the
# `OpGroup*` instructions. All work-items of the sub-group have to execute them together.

const collective_types = [Int32, UInt32, Int64, UInt64, Float16, Float32, Float64]

for op in (:add, :min, :max)
reduce = Symbol(:sub_group_reduce_, op)
scan_inclusive = Symbol(:sub_group_scan_inclusive_, op)
scan_exclusive = Symbol(:sub_group_scan_exclusive_, op)
@eval export $reduce, $scan_inclusive, $scan_exclusive
for T in collective_types
@eval begin
@device_function $reduce(x::$T) =
@builtin_ccall($(String(reduce)), $T, ($T,), x, convergent = true)
@device_function $scan_inclusive(x::$T) =
@builtin_ccall($(String(scan_inclusive)), $T, ($T,), x, convergent = true)
@device_function $scan_exclusive(x::$T) =
@builtin_ccall($(String(scan_exclusive)), $T, ($T,), x, convergent = true)
end
end
end

# `x` of the work-item with (1-based) sub-group local id `lane`, which has to be the same for
# all work-items of the sub-group
export sub_group_broadcast
for T in collective_types
@eval @device_function sub_group_broadcast(x::$T, lane::Integer) =
@builtin_ccall("sub_group_broadcast", $T, ($T, UInt32), x, (lane - 1) % UInt32, convergent = true)
end
8 changes: 6 additions & 2 deletions lib/intrinsics/src/shuffle.jl
Original file line number Diff line number Diff line change
@@ -1,16 +1,20 @@
export sub_group_shuffle, sub_group_shuffle_xor

# Shuffles from `cl_khr_subgroup_shuffle`. A lane that is out of range (e.g. `i < 1`) gives an
# undefined value, as with the OpenCL C built-ins, so the lane is passed modulo `UInt32`
# rather than converted with a check.

const gentypes = [Int8, UInt8, Int16, UInt16, Int32, UInt32, Int64, UInt64, Float16, Float32, Float64]

for gentype in gentypes
@eval begin
@device_function sub_group_shuffle(x::$gentype, i::Integer) =
@builtin_ccall("__spirv_GroupNonUniformShuffle", $gentype,
(UInt32, $gentype, UInt32),
UInt32(Scope.Subgroup), x, UInt32(i - 1))
UInt32(Scope.Subgroup), x, (i - 1) % UInt32, convergent = true)
@device_function sub_group_shuffle_xor(x::$gentype, mask::Integer) =
@builtin_ccall("__spirv_GroupNonUniformShuffleXor", $gentype,
(UInt32, $gentype, UInt32),
UInt32(Scope.Subgroup), x, UInt32(mask))
UInt32(Scope.Subgroup), x, mask % UInt32, convergent = true)
end
end
33 changes: 4 additions & 29 deletions lib/intrinsics/src/synchronization.jl
Original file line number Diff line number Diff line change
Expand Up @@ -30,38 +30,13 @@ module MemorySemantics
const Signal = 0x8000
end

# `@builtin_ccall` does not support additional attributes like `convergent`
# XXX: is this even needed? Doesn't LLVM reconstruct these?
# using the `@builtin_ccall` version causes validation issues.

# call a builtin that is `convergent`, i.e., that must not be made control-dependent on
# additional values
function convergent_call!(builder::IRBuilder, name::String, args::Vector{<:Value})
ft = LLVM.FunctionType(LLVM.VoidType(), [arg.value_type for arg in args])
f = LLVM.Function(current_module(builder), name, ft)
push!(f.function_attributes, EnumAttribute(:convergent))
call!(builder, ft, f, args)
return
end

#@device_function @inline memory_barrier(scope, semantics) =
# @builtin_ccall("__spirv_MemoryBarrier", Cvoid, (UInt32, UInt32), scope, semantics)
@device_function memory_barrier(scope, semantics) =
_memory_barrier(convert(UInt32, scope), convert(UInt32, semantics))
@llvmgenerated builder function _memory_barrier(scope::UInt32, semantics::UInt32)::Nothing
convergent_call!(builder, "_Z21__spirv_MemoryBarrierjj", [scope, semantics])
end
@builtin_ccall("__spirv_MemoryBarrier", Cvoid, (UInt32, UInt32), scope, semantics,
convergent = true)

#@device_function @inline control_barrier(execution_scope, memory_scope, memory_semantics) =
# @builtin_ccall("__spirv_ControlBarrier", Cvoid, (UInt32, UInt32, UInt32),
# execution_scope, memory_scope, memory_semantics)
@device_function @inline control_barrier(execution_scope, memory_scope, memory_semantics) =
_control_barrier(convert(UInt32, execution_scope), convert(UInt32, memory_scope),
convert(UInt32, memory_semantics))
@llvmgenerated builder function _control_barrier(execution::UInt32, memory::UInt32,
semantics::UInt32)::Nothing
convergent_call!(builder, "_Z22__spirv_ControlBarrierjjj", [execution, memory, semantics])
end
@builtin_ccall("__spirv_ControlBarrier", Cvoid, (UInt32, UInt32, UInt32),
execution_scope, memory_scope, memory_semantics, convergent = true)

## OpenCL-compatible fence API

Expand Down
45 changes: 42 additions & 3 deletions lib/intrinsics/src/utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,16 @@ const known_intrinsics = String["printf"]
#
# This macro also keeps track of called builtins, generating `ccall("extern...", llvmcall)`
# expressions for them (so that we can exclude them during IR verification).
#
# Builtins that all work-items of a (sub-)group have to execute together (barriers, sub-group
# shuffles, votes, reductions, ...) have to be called with `convergent = true`, see
# `convergent_builtin`.
macro builtin_ccall(name, ret, argtypes, args...)
convergent = false
if !isempty(args) && Meta.isexpr(args[end], :(=), 2) && args[end].args[1] == :convergent
convergent = args[end].args[2]::Bool
args = args[1:end-1]
end
@assert Meta.isexpr(argtypes, :tuple)
argtypes = argtypes.args

Expand Down Expand Up @@ -47,13 +56,43 @@ macro builtin_ccall(name, ret, argtypes, args...)
for t in argtypes
# with `@eval @builtin_ccall`, we get actual types in the ast, otherwise symbols
t = (isa(t, Symbol) || isa(t, Expr)) ? __module__.eval(t) : t
# `convergent_builtin` passes arguments as `llvmcall` lowers them, without the
# `Bool` and pointer conversions of `@typed_ccall`
convergent && (t <: Union{Bool, Ptr, LLVMPtr}) &&
error("@builtin_ccall with `convergent = true` does not support arguments of type $t")
mangled *= mangle(t)
end

push!(__module__.known_intrinsics, mangled)
esc(quote
@typed_ccall($mangled, llvmcall, $ret, ($(argtypes...),), $(args...))
end)
if convergent
rettyp = (isa(ret, Symbol) || isa(ret, Expr)) ? __module__.eval(ret) : ret
rettyp <: Union{Bool, Ptr, LLVMPtr} &&
error("@builtin_ccall with `convergent = true` does not support return type $rettyp")
@assert length(argtypes) == length(args)
converted = [:(convert($T, $arg)) for (T, arg) in zip(argtypes, args)]
esc(quote
$convergent_builtin($(Val(Symbol(mangled))), $ret, $(converted...))
end)
else
esc(quote
@typed_ccall($mangled, llvmcall, $ret, ($(argtypes...),), $(args...))
end)
end
end

# call a builtin that all work-items of the (sub-)group have to execute together. its
# declaration is marked `convergent`, so that the optimizer doesn't make the call
# control-dependent on additional values, e.g., by duplicating it into the arms of a branch
# that computes its argument, which makes the work-items execute different calls.
@llvmgenerated builder function convergent_builtin(::Val{name}, ::Type{T},
args...)::T where {name, T}
rt = T === Nothing ? LLVM.VoidType() : convert(LLVMType, T)
ft = LLVM.FunctionType(rt, LLVMType[arg.value_type for arg in args])
f = LLVM.Function(current_module(builder), String(name), ft)
push!(f.function_attributes, EnumAttribute(:convergent))
push!(f.function_attributes, EnumAttribute(:nounwind))
rv = call!(builder, ft, f, collect(Value, args))
T === Nothing ? nothing : rv
end


Expand Down
17 changes: 17 additions & 0 deletions lib/intrinsics/src/vote.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
export sub_group_any, sub_group_all, sub_group_ballot

# Sub-group votes. These call the OpenCL C built-ins, which the SPIR-V back-end lowers to
# `OpGroupAny`, `OpGroupAll` and `OpGroupNonUniformBallot`, since `@builtin_ccall` can't
# mangle the `bool` arguments of the corresponding SPIR-V wrapper builtins.

# `cl_khr_subgroups`
@device_function sub_group_any(predicate::Bool) =
@builtin_ccall("sub_group_any", Int32, (Int32,), Int32(predicate), convergent = true) != Int32(0)

@device_function sub_group_all(predicate::Bool) =
@builtin_ccall("sub_group_all", Int32, (Int32,), Int32(predicate), convergent = true) != Int32(0)

# `cl_khr_subgroup_ballot`: bit `i` (counting from the least significant bit of the first
# component) is set for the work-item with (0-based) sub-group local id `i`
@device_function sub_group_ballot(predicate::Bool) =
@builtin_ccall("sub_group_ballot", NTuple{4, VecElement{UInt32}}, (Int32,), Int32(predicate), convergent = true)
182 changes: 182 additions & 0 deletions test/intrinsics.jl
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,188 @@ cl.sub_groups_supported(cl.device()) && @testset "Sub-groups" begin
@test Array(d_in) == in[idxs]
end
end
@testset "shuffle out of range" begin
# an out-of-range lane gives an undefined value rather than an error
function shfl_up_kernel(d)
i = get_sub_group_local_id()
val = sub_group_shuffle(d[i], i - 1)
if i > 1
d[i] = val
end
return
end

@testset for T in cl.sub_group_shuffle_supported_types(cl.device())
a = rand(T, sg_size)
d_a = CLArray(a)
@opencl local_size = sg_size global_size = sg_size shfl_up_kernel(d_a)
@test Array(d_a) == [a[1]; a[1:(end - 1)]]
end
end
@testset "shuffle of a divergent value" begin
# without `convergent`, the optimizer duplicates the shuffle into both arms of the
# branch that computes its argument, so the work-items call it separately
function divergent_shfl_kernel(out, in, m)
i = get_sub_group_local_id()
x = i <= m ? (@inbounds in[i]) : Int32(0)
r = sub_group_shuffle(x, 1)
if i <= m
@inbounds out[i] = r
end
return
end
m = sg_size ÷ 2
a = Int32.(rand(1:20, sg_size))
d_out = CLArray(zeros(Int32, sg_size))
@opencl local_size = sg_size global_size = sg_size divergent_shfl_kernel(d_out, CLArray(a), m)
@test Array(d_out)[1:m] == fill(a[1], m)
end
# spirv2clc, which translates SPIR-V to OpenCL C for the OpenCL C program backend, doesn't
# implement the instructions of the votes, ballots and collectives (OpGroupAll, OpGroupAny,
# OpGroupBroadcast, OpGroupIAdd etc., and the GroupNonUniformBallot capability)
uses_spirv2clc = OpenCL.resolve_program_backend(cl.device(), OpenCL.program_backend()) !== :spirv
if uses_spirv2clc
@test_skip "sub-group votes and collectives through spirv2clc"
else
@testset "any/all" begin
function vote_kernel(out, pred)
i = get_sub_group_local_id()
out[i, 1] = sub_group_any(pred[i])
out[i, 2] = sub_group_all(pred[i])
return
end

@testset "$name" for (name, pred) in (
"none" => falses(sg_size), "all" => trues(sg_size),
"some" => [i % 3 == 1 for i in 1:sg_size],
)
d_out = CLArray(zeros(Bool, sg_size, 2))
@opencl local_size = sg_size global_size = sg_size vote_kernel(d_out, CLArray(collect(pred)))
out = Array(d_out)
@test all(==(any(pred)), out[:, 1])
@test all(==(all(pred)), out[:, 2])
end

# the predicate comes from a short-circuiting `&&` and is used again after the vote:
# without `convergent`, jump threading duplicates the vote into both arms of the
# `&&` (one of them with a constant `false` predicate), so the work-items of the
# sub-group call it separately (PoCL then returns `false` to all of them)
function divergent_vote_kernel(out, a, b, n)
i = get_sub_group_local_id() % Int32
n_any = Int32(0)
n_all = Int32(0)
m = Int32(0)
@inbounds while m < n
j = ((i - Int32(1) + m) % n) + Int32(1)
pred = a[i] > 0.0f0 && b[j] > 0.0f0
if sub_group_any(pred)
n_any += pred ? Int32(2) : Int32(1)
end
if sub_group_all(pred)
n_all += pred ? Int32(2) : Int32(1)
end
m += Int32(1)
end
@inbounds out[i, 1] = n_any
@inbounds out[i, 2] = n_all
return
end
a = Float32[isodd(i) for i in 1:sg_size]
b = ones(Float32, sg_size)
d_out = CLArray(zeros(Int32, sg_size, 2))
@opencl local_size = sg_size global_size = sg_size divergent_vote_kernel(d_out, CLArray(a), CLArray(b), Int32(sg_size))
out = Array(d_out)
@test out[:, 1] == [isodd(i) ? 2sg_size : sg_size for i in 1:sg_size]
@test all(==(0), out[:, 2])
end
"cl_khr_subgroup_ballot" in cl.device().extensions && @testset "ballot" begin
function ballot_kernel(out, pred)
i = get_sub_group_local_id()
mask = sub_group_ballot(pred[i])
for j in 1:4
out[i, j] = mask[j].value
end
return
end

pred = [i % 3 == 1 || i == sg_size for i in 1:sg_size]
d_out = CLArray(zeros(UInt32, sg_size, 4))
@opencl local_size = sg_size global_size = sg_size ballot_kernel(d_out, CLArray(pred))
expected = zeros(UInt32, 4)
for i in 1:sg_size
pred[i] && (expected[(i - 1) ÷ 32 + 1] |= UInt32(1) << ((i - 1) % 32))
end
out = Array(d_out)
@test all(i -> out[i, :] == expected, 1:sg_size)
end
@testset "collectives" begin
function collective_kernel(out, in, lane)
i = get_sub_group_local_id()
x = in[i]
out[i, 1] = sub_group_reduce_add(x)
out[i, 2] = sub_group_reduce_min(x)
out[i, 3] = sub_group_reduce_max(x)
out[i, 4] = sub_group_scan_inclusive_add(x)
out[i, 5] = sub_group_scan_exclusive_add(x)
out[i, 6] = sub_group_scan_inclusive_max(x)
out[i, 7] = sub_group_scan_exclusive_min(x)
out[i, 8] = sub_group_broadcast(x, lane)
return
end

lane = min(3, sg_size)
types = [Int32, UInt32, Int64, UInt64, Float32]
"cl_khr_fp16" in cl.device().extensions && push!(types, Float16)
"cl_khr_fp64" in cl.device().extensions && push!(types, Float64)
@testset for T in types
# small integers, so that the sums are exact
a = T.(rand(1:20, sg_size))
d_out = CLArray(zeros(T, sg_size, 8))
@opencl local_size = sg_size global_size = sg_size collective_kernel(d_out, CLArray(a), lane)
out = Array(d_out)
@test all(==(sum(a)), out[:, 1])
@test all(==(minimum(a)), out[:, 2])
@test all(==(maximum(a)), out[:, 3])
@test out[:, 4] == cumsum(a)
@test out[:, 5] == [zero(T); cumsum(a)[1:(end - 1)]]
@test out[:, 6] == accumulate(max, a)
# the exclusive scan of the first work-item is the identity, `typemax` for `min`
@test out[2:end, 7] == accumulate(min, a)[1:(end - 1)]
@test out[1, 7] == (T <: AbstractFloat ? T(Inf) : typemax(T))
@test all(==(a[lane]), out[:, 8])
end

# the values come from a divergent branch: without `convergent`, the optimizer
# duplicates the collective into both arms of the branch, which call it separately
function divergent_reduce_kernel(out, in, m)
i = get_sub_group_local_id()
x = i <= m ? (@inbounds in[i]) : Int32(0)
r = sub_group_reduce_add(x)
if i <= m
@inbounds out[i] = r
end
return
end
m = sg_size ÷ 2
a = Int32.(rand(1:20, sg_size))
d_out = CLArray(zeros(Int32, sg_size))
@opencl local_size = sg_size global_size = sg_size divergent_reduce_kernel(d_out, CLArray(a), m)
@test all(==(sum(a[1:m])), Array(d_out)[1:m])

# the same with a bounds check, whose (never taken) early exit made PoCL 7.2 peel
# the first work-item of the region, which then read a separate copy of the
# collective's scratch memory (pocl/pocl#2239, in pocl_jll 7.2.1+1)
function divergent_checked_reduce_kernel(out, in, m)
i = get_sub_group_local_id()
x = i <= m ? in[i] : Int32(0)
out[i] = sub_group_reduce_add(x)
return
end
d_out = CLArray(zeros(Int32, sg_size))
@opencl local_size = sg_size global_size = sg_size divergent_checked_reduce_kernel(d_out, CLArray(a), m)
@test all(==(sum(a[1:m])), Array(d_out))
end
end
end
end # if cl.sub_groups_supported(cl.device())

Expand Down
Loading