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
10 changes: 10 additions & 0 deletions src/validation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,16 @@ function check_invocation(@nospecialize(job::CompilerJob))
This is a CPU-only object not supported by GPUCompiler."""))
end

# Before `Core.TypeEgal`, `Type{T}` is only a singleton when `T` has a unique
# representation, so e.g. `Type{Union{Missing, Int}}` or `Type{Vector}` would be
# passed as a boxed host pointer.
if Base.isType(dt)
throw(KernelError(job, "passing a non-singleton type argument",
"""Argument $arg_i to your kernel function is the type $(dt.parameters[1]), which
cannot be passed to a GPU kernel on this version of Julia.
Pass `Val($(dt.parameters[1]))` instead, or a value of that type."""))
end

# If an object doesn't have fields, it can only be used by identity, so we can allow
# them to be passed to the GPU (this also applies to e.g. Symbols).
if fieldcount(dt) == 0
Expand Down
24 changes: 24 additions & 0 deletions test/native.jl
Original file line number Diff line number Diff line change
Expand Up @@ -1728,6 +1728,30 @@ end
end
end

@testset "non-singleton type arguments" begin
mod = @eval module $(gensym())
import ..sink
foo(::Type{T}) where {T} = (sink(Int(Missing <: T)); return)
bar(::Val{T}) where {T} = (sink(Int(Missing <: T)); return)
end

for T in (Union{Missing, Int}, Vector)
if isdefined(Core, :TypeEgal)
# `Core.TypeEgal` makes every closed type argument a singleton
Native.code_execution(mod.foo, Tuple{Type{T}})
else
@test_throws_message(KernelError,
Native.code_execution(mod.foo, Tuple{Type{T}})) do msg
occursin("passing a non-singleton type argument", msg) &&
occursin(string(T), msg)
end
end

# the suggested alternative
Native.code_execution(mod.bar, Tuple{Val{T}})
end
end

@testset "invalid LLVM IR" begin
mod = @eval module $(gensym())
foobar(i) = println(i)
Expand Down
Loading