diff --git a/Project.toml b/Project.toml index 2a992247..0e15ba26 100644 --- a/Project.toml +++ b/Project.toml @@ -42,7 +42,7 @@ GPUArrays = "11.5.14" GPUCompiler = "2.9" GPUToolbox = "3.3.1" KernelAbstractions = "0.9.39" -LLVM = "6, 7, 8, 9" +LLVM = "10" NEO_jll = "=26.18.38308" PrecompileTools = "1" Preferences = "1" diff --git a/src/compiler/compilation.jl b/src/compiler/compilation.jl index 8f463103..b1ce7533 100644 --- a/src/compiler/compilation.jl +++ b/src/compiler/compilation.jl @@ -59,12 +59,14 @@ function GPUCompiler.finish_module!(job::oneAPICompilerJob, mod::LLVM.Module, job, mod, entry) # OpenCL 2.0 - push!(metadata(mod)["opencl.ocl.version"], + push!( + get!(mod.metadata, "opencl.ocl.version").operands, MDNode([ConstantInt(Int32(2)), ConstantInt(Int32(0))])) # SPIR-V 1.5 - push!(metadata(mod)["opencl.spirv.version"], + push!( + get!(mod.metadata, "opencl.spirv.version").operands, MDNode([ConstantInt(Int32(1)), ConstantInt(Int32(5))])) @@ -90,7 +92,7 @@ function GPUCompiler.finish_ir!(job::oneAPICompilerJob, mod::LLVM.Module, # (`supports_bfloat16 = false`, see `_compiler_config`), so this pass never # runs there. target = job.config.target - if @static(isdefined(Core, :BFloat16) && isdefined(LLVM, :BFloatType)) && + if isdefined(Core, :BFloat16) && target.supports_bfloat16 && !occursin("SPV_KHR_bfloat16", target.extensions) lower_bfloat_to_i16!(mod) end @@ -105,21 +107,21 @@ function add_heap!(mod::LLVM.Module, entry::LLVM.Function) heap_field = Base.fieldindex(KernelState, :heap) - 1 uses_heap(mod, T_state, heap_field) || return false - params = parameters(entry) - if isempty(params) || value_type(params[1]) != T_state - error("kernel `$(LLVM.name(entry))` allocates but has no kernel state to hold the heap") + params = entry.parameters + if isempty(params) || params[1].value_type != T_state + error("kernel `$(entry.name)` allocates but has no kernel state to hold the heap") end state = params[1] - users = LLVM.Value[user(use) for use in uses(state)] + # collect the users now, before the insertvalue below becomes one of them + users = collect(state.users) T_size = convert(LLVMType, Csize_t) T_heap = LLVM.StructType([T_size, T_size, LLVM.ArrayType(LLVM.Int8Type(), HEAP_SIZE)]) T_ptr = convert(LLVMType, fieldtype(KernelState, :heap)) @dispose builder = IRBuilder() begin - position!(builder, first(instructions(first(blocks(entry))))) + position!(builder, LLVM.at_begin(entry.entry)) - heap = alloca!(builder, T_heap, "heap") - alignment!(heap, HEAP_ALIGNMENT) + heap = alloca!(builder, T_heap, "heap"; align = HEAP_ALIGNMENT) store!(builder, ConstantInt(T_size, 0), struct_gep!(builder, T_heap, heap, 0)) store!(builder, ConstantInt(T_size, HEAP_SIZE), struct_gep!(builder, T_heap, heap, 1)) @@ -127,10 +129,7 @@ function add_heap!(mod::LLVM.Module, entry::LLVM.Function) ptr = pointercast!(builder, heap, T_ptr) new_state = insert_value!(builder, state, ptr, heap_field, "state") for u in users - ops = operands(u) - for i in 1:length(ops) - ops[i] == state && (ops[i] = new_state) - end + replace!(u.operands, state => new_state) end end @@ -139,10 +138,10 @@ end # Inspect field reads rather than calls to malloc, which may already have been inlined. function uses_heap(mod::LLVM.Module, T_state::LLVMType, heap_field::Integer) - for f in functions(mod), bb in blocks(f), inst in instructions(bb) + for f in mod.functions, bb in f.blocks, inst in bb.instructions inst isa LLVM.ExtractValueInst || continue - value_type(operands(inst)[1]) == T_state || continue - unsafe_load(LLVM.API.LLVMGetIndices(inst)) == heap_field && return true + inst.operands[1].value_type == T_state || continue + first(inst.indices) == heap_field && return true end return false end @@ -165,48 +164,42 @@ function lower_bfloat_to_i16!(mod::LLVM.Module) eliminate_bf16_bitcasts!(mod, T_bf16, T_i16) # Phase 2: Replace remaining bfloat GEPs, loads, and stores with i16 equivalents. - for f in functions(mod) - isempty(blocks(f)) && continue - for bb in blocks(f) - to_replace = LLVM.Instruction[] - for inst in instructions(bb) - opcode = LLVM.API.LLVMGetInstructionOpcode(inst) - if opcode == LLVM.API.LLVMGetElementPtr - src_ty = LLVMType(LLVM.API.LLVMGetGEPSourceElementType(inst)) - src_ty == T_bf16 && push!(to_replace, inst) - elseif opcode == LLVM.API.LLVMLoad - value_type(inst) == T_bf16 && push!(to_replace, inst) - elseif opcode == LLVM.API.LLVMStore - value_type(LLVM.operands(inst)[1]) == T_bf16 && push!(to_replace, inst) + @dispose builder = IRBuilder() begin + for f in mod.functions + isdeclaration(f) && continue + for bb in f.blocks + to_replace = LLVM.Instruction[] + for inst in bb.instructions + if inst isa LLVM.GetElementPtrInst + inst.source_element_type == T_bf16 && push!(to_replace, inst) + elseif inst isa LLVM.LoadInst + inst.value_type == T_bf16 && push!(to_replace, inst) + elseif inst isa LLVM.StoreInst + inst.value_operand.value_type == T_bf16 && push!(to_replace, inst) + end end - end - for inst in to_replace - opcode = LLVM.API.LLVMGetInstructionOpcode(inst) - builder = LLVM.IRBuilder() - LLVM.position!(builder, inst) - - if opcode == LLVM.API.LLVMGetElementPtr - ptr = LLVM.operands(inst)[1] - indices = LLVM.Value[LLVM.operands(inst)[i] for i in 2:length(LLVM.operands(inst))] - new_gep = if LLVM.API.LLVMIsInBounds(inst) != 0 - LLVM.inbounds_gep!(builder, T_i16, ptr, indices) - else - LLVM.gep!(builder, T_i16, ptr, indices) + for inst in to_replace + position!(builder, LLVM.before(inst)) + + if inst isa LLVM.GetElementPtrInst + ptr = inst.pointer_operand + new_gep = if inst.inbounds + LLVM.inbounds_gep!(builder, T_i16, ptr, inst.indices) + else + LLVM.gep!(builder, T_i16, ptr, inst.indices) + end + LLVM.replace_uses!(inst, new_gep) + elseif inst isa LLVM.LoadInst + ptr = inst.pointer_operand + new_load = LLVM.load!(builder, T_i16, ptr) + LLVM.replace_uses!(inst, new_load) + elseif inst isa LLVM.StoreInst + LLVM.store!(builder, inst.value_operand, inst.pointer_operand) end - LLVM.replace_uses!(inst, new_gep) - elseif opcode == LLVM.API.LLVMLoad - ptr = LLVM.operands(inst)[1] - new_load = LLVM.load!(builder, T_i16, ptr) - LLVM.replace_uses!(inst, new_load) - elseif opcode == LLVM.API.LLVMStore - val = LLVM.operands(inst)[1] - ptr = LLVM.operands(inst)[2] - LLVM.store!(builder, val, ptr) - end - LLVM.API.LLVMInstructionEraseFromParent(inst) - LLVM.dispose(builder) + erase!(inst) + end end end end @@ -219,15 +212,15 @@ function eliminate_bf16_bitcasts!(mod::LLVM.Module, T_bf16::LLVMType, T_i16::LLV changed = true while changed changed = false - for f in functions(mod) - isempty(blocks(f)) && continue - for bb in blocks(f) + for f in mod.functions + isdeclaration(f) && continue + for bb in f.blocks to_delete = LLVM.Instruction[] - for inst in instructions(bb) - if LLVM.API.LLVMGetInstructionOpcode(inst) == LLVM.API.LLVMBitCast - src = LLVM.operands(inst)[1] - src_ty = value_type(src) - dst_ty = value_type(inst) + for inst in bb.instructions + if inst isa LLVM.BitCastInst + src = inst.operands[1] + src_ty = src.value_type + dst_ty = inst.value_type if (src_ty == T_i16 && dst_ty == T_bf16) || (src_ty == T_bf16 && dst_ty == T_i16) || (src_ty == dst_ty) @@ -238,7 +231,7 @@ function eliminate_bf16_bitcasts!(mod::LLVM.Module, T_bf16::LLVMType, T_i16::LLV end end for inst in to_delete - LLVM.API.LLVMInstructionEraseFromParent(inst) + erase!(inst) end end end @@ -319,11 +312,13 @@ end # session-portable so they survive precompilation when stored on a cached `CodeInstance`. function compile_to_obj(@nospecialize(job::CompilerJob)) # TODO: on 1.9, this actually creates a context. cache those. - asm, meta = JuliaContext() do ctx - GPUCompiler.compile(:obj, job) + return JuliaContext() do ctx + asm, meta = GPUCompiler.compile(:obj, job) + # the IR belongs to us: dispose of it, or it leaks along with the context + @dispose ir = meta.ir begin + (image = asm, entry = meta.entry.name) + end end - - (image=asm, entry=LLVM.name(meta.entry)) end # link the SPIR-V bytes into a session-local `ZeKernel` on the given context and device. diff --git a/src/oneAPI.jl b/src/oneAPI.jl index 6b6db42c..d402308e 100644 --- a/src/oneAPI.jl +++ b/src/oneAPI.jl @@ -14,8 +14,7 @@ import Preferences import KernelAbstractions: KernelAbstractions -using LLVM -using LLVM.Interop +using LLVM, LLVM.IR, LLVM.Build, LLVM.Interop using Core: LLVMPtr import Libdl