From 1fe0ea9ca1ac3cef745be71398f427d4ac4e2e0e Mon Sep 17 00:00:00 2001 From: Tim Besard Date: Fri, 25 Sep 2026 11:20:47 +0200 Subject: [PATCH] Flatten nested insertvalues for Intel's SPIR-V driver. Intel's graphics compiler silently drops the store of a field when legalizing an aggregate that was built by a chain of insertvalues, if a later one in the chain has more indices. Add a `driver` field to SPIRVCompilerTarget so that consumers can identify the driver, and flatten nested insertvalues when it is `:intel`. --- src/spirv.jl | 63 +++++++++++++++++++++++++++++++++++++++++++ test/helpers/spirv.jl | 4 +-- test/spirv.jl | 29 ++++++++++++++++++++ 3 files changed, 94 insertions(+), 2 deletions(-) diff --git a/src/spirv.jl b/src/spirv.jl index b33bcdea..3a4515ba 100644 --- a/src/spirv.jl +++ b/src/spirv.jl @@ -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 @@ -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) @@ -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 diff --git a/test/helpers/spirv.jl b/test/helpers/spirv.jl index 5e765bb7..bfdcbba8 100644 --- a/test/helpers/spirv.jl +++ b/test/helpers/spirv.jl @@ -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 diff --git a/test/spirv.jl b/test/spirv.jl index b7e4a1f4..8fd930ac 100644 --- a/test/spirv.jl +++ b/test/spirv.jl @@ -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