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
63 changes: 63 additions & 0 deletions src/spirv.jl
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,11 @@ Base.@kwdef struct SPIRVCompilerTarget <: AbstractCompilerTarget
supports_bfloat16::Bool = false

backend::Symbol = isavailable(SPIRV_LLVM_Backend_jll) ? :llvm : :khronos

# the driver that will consume the SPIR-V, used to work around its bugs. `:generic` if
# unknown, `:intel` for Intel's GPU driver (NEO, with IGC), `:pocl`, `:nvidia`, ...
driver::Symbol = :generic

# XXX: these don't really belong in the _target_ struct
validate::Bool = false
optimize::Bool = false
Expand Down Expand Up @@ -117,6 +122,11 @@ function finish_ir!(job::CompilerJob{SPIRVCompilerTarget}, mod::LLVM.Module,
# the SPIR-V back-ends lower `llvm.minimum`/`llvm.maximum` to NaN-ignoring `fmin`/`fmax`
lower_minimum_maximum!(mod)

# IGC drops fields when legalizing aggregates built by nested `insertvalue`s
if job.config.target.driver === :intel
flatten_nested_insertvalue!(mod)
end

# convert the kernel state argument to a byval reference
if job.config.kernel
state = kernel_state_type(job)
Expand Down Expand Up @@ -314,6 +324,59 @@ function rm_freeze!(@nospecialize(job::CompilerJob), mod::LLVM.Module)
return changed
end

# flatten `insertvalue`s with multiple indices into single-index ones, extracting and
# re-inserting the intermediate aggregates: `insertvalue %agg, %val, 1, 0` becomes
# %sub = extractvalue %agg, 1
# %new = insertvalue %sub, %val, 0
# insertvalue %agg, %new, 1
#
# this works around a bug in Intel's graphics compiler, whose `TypesLegalizationPass` splits
# aggregate stores (and phis, which it lowers to stores) into per-field stores by looking up
# each field in the `insertvalue` chain. that lookup gives up on an `insertvalue` with more
# indices than the field it is looking for, silently dropping the store of that field.
# e.g., storing `(flag::Bool, (a, b))` built by inserting `flag`, `a` and `b`, loses `flag`
# (intel/intel-graphics-compiler#378, JuliaGPU/OpenCL.jl#502, JuliaGPU/oneAPI.jl#259).
# this needs to run after optimization, as InstCombine folds these sequences back together.
function flatten_nested_insertvalue!(mod::LLVM.Module)
changed = false
@tracepoint "flatten nested insertvalue" begin

for f in functions(mod), bb in blocks(f)
worklist = filter(collect(instructions(bb))) do inst
opcode(inst) == LLVM.API.LLVMInsertValue && LLVM.API.LLVMGetNumIndices(inst) > 1
end
isempty(worklist) && continue

@dispose builder=IRBuilder() begin
for inst in worklist
agg, val = operands(inst)
n = LLVM.API.LLVMGetNumIndices(inst)
idxptr = LLVM.API.LLVMGetIndices(inst)
indices = [unsafe_load(idxptr, i) for i in 1:n]

position!(builder, inst)
new = flatten_insertvalue!(builder, agg, val, indices)
replace_uses!(inst, new)
erase!(inst)
changed = true
end
end
end

end
return changed
end

function flatten_insertvalue!(builder::IRBuilder, agg::LLVM.Value, val::LLVM.Value,
indices::AbstractVector)
idx = first(indices)
if length(indices) > 1
sub = extract_value!(builder, agg, idx)
val = flatten_insertvalue!(builder, sub, val, @view indices[2:end])
end
return insert_value!(builder, agg, val, idx)
end

# expand `llvm.minimum` and `llvm.maximum`, which Julia uses for `min` and `max` of
# floating-point numbers. these return NaN when either operand is NaN, and order -0.0 before
# +0.0, but both SPIR-V back-ends translate them to OpenCL's `fmin` and `fmax`, which return
Expand Down
4 changes: 2 additions & 2 deletions test/helpers/spirv.jl
Original file line number Diff line number Diff line change
Expand Up @@ -8,11 +8,11 @@ GPUCompiler.runtime_module(::CompilerJob{<:Any,CompilerParams}) = TestRuntime

function create_job(@nospecialize(func), @nospecialize(types);
supports_fp16=true, supports_fp64=true, supports_bfloat16=false,
backend::Symbol, kwargs...)
backend::Symbol, driver::Symbol=:generic, kwargs...)
config_kwargs, kwargs = split_kwargs(kwargs, GPUCompiler.CONFIG_KWARGS)
source = methodinstance(typeof(func), Base.to_tuple_type(types), Base.get_world_counter())
target = SPIRVCompilerTarget(; backend, validate=true, optimize=true,
supports_fp16, supports_fp64, supports_bfloat16)
supports_fp16, supports_fp64, supports_bfloat16, driver)
params = CompilerParams()
config = CompilerConfig(target, params; kernel=false, config_kwargs...)
CompilerJob(source, config), kwargs
Expand Down
29 changes: 29 additions & 0 deletions test/spirv.jl
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,35 @@ end
end
end

@testset "nested insertvalue" begin
# Intel's graphics compiler drops fields when legalizing aggregates built by nested
# `insertvalue`s, so these are flattened for that driver (JuliaGPU/OpenCL.jl#502)
mod = @eval module $(gensym())
struct FlagFirst
valid::Bool
value::Tuple{Float32, Int32}
end
kernel(p::Core.LLVMPtr{FlagFirst,1}, x::Float32, i::Int32) =
(unsafe_store!(p, FlagFirst(x >= 0, (x, i))); return)
end
tt = Tuple{Core.LLVMPtr{mod.FlagFirst,1}, Float32, Int32}

@test @filecheck begin
@check_label "define {{.*}} @{{(julia|j)_kernel_[0-9]+}}"
@check "insertvalue {{.*}}, 1, 0"
SPIRV.code_llvm(mod.kernel, tt; backend)
end

@test @filecheck begin
@check_label "define {{.*}} @{{(julia|j)_kernel_[0-9]+}}"
@check_not "insertvalue {{.*}}, {{[0-9]+}}, {{[0-9]+}}"
@check "extractvalue {{.*}}, 1"
@check "insertvalue {{.*}}, 0"
@check "insertvalue {{.*}}, 1"
SPIRV.code_llvm(mod.kernel, tt; backend, driver=:intel)
end
end

############################################################################################

@testset "asm" begin
Expand Down
Loading