diff --git a/Project.toml b/Project.toml index 0236d89d..10100a2d 100644 --- a/Project.toml +++ b/Project.toml @@ -35,7 +35,7 @@ GPUArrays = "11.2.1" GPUCompiler = "2.9" GPUToolbox = "3.1" KernelAbstractions = "0.9.38" -LLVM = "9.6" +LLVM = "9.14" LinearAlgebra = "1" OpenCL_jll = "=2024.10.24" Preferences = "1" diff --git a/lib/intrinsics/Project.toml b/lib/intrinsics/Project.toml index 653472ca..64b8066d 100644 --- a/lib/intrinsics/Project.toml +++ b/lib/intrinsics/Project.toml @@ -18,7 +18,7 @@ SPIRVIntrinsicsSIMDExt = "SIMD" [compat] ExprTools = "0.1" GPUToolbox = "0.2, 0.3, 1, 2, 3" -LLVM = "9.1" +LLVM = "9.14" SIMD = "3.6" SpecialFunctions = "1.3, 2" julia = "1.10" diff --git a/lib/intrinsics/src/memory.jl b/lib/intrinsics/src/memory.jl index da5c5b2b..0339d616 100644 --- a/lib/intrinsics/src/memory.jl +++ b/lib/intrinsics/src/memory.jl @@ -1,38 +1,22 @@ # local memory # get a pointer to local memory, with known (static) or zero length (dynamic) -@generated function emit_localmemory(::Type{T}, ::Val{len}=Val(0)) where {T,len} - Context() do ctx - # XXX: as long as LLVMPtr is emitted as i8*, it doesn't make sense to type the GV - eltyp = convert(LLVMType, LLVM.Int8Type()) - T_ptr = convert(LLVMType, LLVMPtr{T,AS.Workgroup}) - - # create a function - llvm_f, _ = create_function(T_ptr) - - # create the global variable - mod = LLVM.parent(llvm_f) - gv_typ = LLVM.ArrayType(eltyp, len * sizeof(T)) - gv = GlobalVariable(mod, gv_typ, "local_memory", AS.Workgroup) - if len > 0 - linkage!(gv, LLVM.API.LLVMInternalLinkage) - initializer!(gv, null(gv_typ)) - end - # TODO: Make the alignment configurable - alignment!(gv, Base.datatype_alignment(T)) - - # generate IR - IRBuilder() do builder - entry = BasicBlock(llvm_f, "entry") - position!(builder, entry) - - ptr = gep!(builder, gv_typ, gv, [ConstantInt(0), ConstantInt(0)]) - - untyped_ptr = bitcast!(builder, ptr, T_ptr) - - ret!(builder, untyped_ptr) - end - - call_function(llvm_f, LLVMPtr{T,AS.Workgroup}) +@llvmgenerated builder function emit_localmemory(::Type{T}, + ::Val{len}=Val(0))::LLVMPtr{T,AS.Workgroup} where {T,len} + # XXX: as long as LLVMPtr is emitted as i8*, it doesn't make sense to type the GV + eltyp = LLVM.Int8Type() + T_ptr = convert(LLVMType, LLVMPtr{T,AS.Workgroup}) + + # create the global variable + gv_typ = LLVM.ArrayType(eltyp, len * sizeof(T)) + gv = GlobalVariable(current_module(builder), gv_typ, "local_memory", AS.Workgroup) + if len > 0 + linkage!(gv, LLVM.API.LLVMInternalLinkage) + initializer!(gv, null(gv_typ)) end + # TODO: Make the alignment configurable + alignment!(gv, Base.datatype_alignment(T)) + + ptr = gep!(builder, gv_typ, gv, [ConstantInt(0), ConstantInt(0)]) + bitcast!(builder, ptr, T_ptr) end diff --git a/lib/intrinsics/src/printf.jl b/lib/intrinsics/src/printf.jl index 203a966d..3177c949 100644 --- a/lib/intrinsics/src/printf.jl +++ b/lib/intrinsics/src/printf.jl @@ -28,70 +28,54 @@ end arg_exprs = [:( argspec[$i] ) for i in 1:length(argspec)] arg_types = [argspec...] - Context() do ctx - T_void = LLVM.VoidType() + generate_llvmcall(Int32, Tuple{arg_types...}, arg_exprs...) do builder, args... T_int32 = LLVM.Int32Type() T_pint8 = LLVM.PointerType(LLVM.Int8Type(), AS.UniformConstant) - # create functions - param_types = LLVMType[convert(LLVMType, typ) for typ in arg_types] - llvm_f, _ = create_function(T_int32, param_types) - mod = LLVM.parent(llvm_f) - - IRBuilder() do builder - entry = BasicBlock(llvm_f, "entry") - position!(builder, entry) - - # `printf` needs to be invoked very specifically, e.g., the format string needs - # to be a pointer to a string, and arguments need to match exactly what is - # expected by the format string, so we cannot rely on how the arguments to this - # function have been passed in (by `llvmcall`). - T_actual_args = LLVMType[] - actual_args = LLVM.Value[] - for (_, (arg, argtyp)) in enumerate(zip(parameters(llvm_f), arg_types)) - if argtyp <: LLVMPtr - # passed as i8* - T,AS = argtyp.parameters - actual_typ = LLVM.PointerType(convert(LLVMType, T), AS) - actual_arg = bitcast!(builder, arg, actual_typ) - elseif argtyp <: Ptr - T = eltype(argtyp) - if T === Nothing - T = Int8 - end - actual_typ = LLVM.PointerType(convert(LLVMType, T)) - actual_arg = if value_type(arg) isa LLVM.PointerType - # passed as i8* or ptr - bitcast!(builder, arg, actual_typ) - else - # passed as i64 - inttoptr!(builder, arg, actual_typ) - end - elseif argtyp <: Bool - # passed as i8 - T = eltype(argtyp) - actual_typ = LLVM.Int1Type() - actual_arg = trunc!(builder, arg, actual_typ) + # `printf` needs to be invoked very specifically, e.g., the format string needs + # to be a pointer to a string, and arguments need to match exactly what is + # expected by the format string, so we cannot rely on how the arguments to this + # function have been passed in (by `llvmcall`). + T_actual_args = LLVMType[] + actual_args = LLVM.Value[] + for (arg, argtyp) in zip(args, arg_types) + if argtyp <: LLVMPtr + # passed as i8* + T,AS = argtyp.parameters + actual_typ = LLVM.PointerType(convert(LLVMType, T), AS) + actual_arg = bitcast!(builder, arg, actual_typ) + elseif argtyp <: Ptr + T = eltype(argtyp) + if T === Nothing + T = Int8 + end + actual_typ = LLVM.PointerType(convert(LLVMType, T)) + actual_arg = if value_type(arg) isa LLVM.PointerType + # passed as i8* or ptr + bitcast!(builder, arg, actual_typ) else - actual_typ = convert(LLVMType, argtyp) - actual_arg = arg + # passed as i64 + inttoptr!(builder, arg, actual_typ) end - push!(T_actual_args, actual_typ) - push!(actual_args, actual_arg) + elseif argtyp <: Bool + # passed as i8 + actual_typ = LLVM.Int1Type() + actual_arg = trunc!(builder, arg, actual_typ) + else + actual_typ = convert(LLVMType, argtyp) + actual_arg = arg end - - str = globalstring_ptr!(builder, String(fmt); addrspace=AS.UniformConstant) - - # invoke printf and return - printf_typ = LLVM.FunctionType(T_int32, [T_pint8]; vararg=true) - printf = LLVM.Function(mod, "printf", printf_typ) - push!(function_attributes(printf), EnumAttribute("nobuiltin")) - chars = call!(builder, printf_typ, printf, [str, actual_args...]) - - ret!(builder, chars) + push!(T_actual_args, actual_typ) + push!(actual_args, actual_arg) end - call_function(llvm_f, Int32, Tuple{arg_types...}, arg_exprs...) + str = globalstring_ptr!(builder, String(fmt); addrspace=AS.UniformConstant) + + # invoke printf and return + printf_typ = LLVM.FunctionType(T_int32, [T_pint8]; vararg=true) + printf = LLVM.Function(current_module(builder), "printf", printf_typ) + push!(function_attributes(printf), EnumAttribute("nobuiltin")) + call!(builder, printf_typ, printf, [str, actual_args...]) end end diff --git a/src/device/array.jl b/src/device/array.jl index 62466015..c961dc66 100644 --- a/src/device/array.jl +++ b/src/device/array.jl @@ -165,40 +165,24 @@ end # there is no SPIR-V equivalent of NVPTX's `ld.global.nc`, so instead of a dedicated # instruction we mark the load `!invariant.load`, which lets LLVM hoist it out of loops # and reorder it across stores to other objects. -@inline @generated function unsafe_invariant_load(ptr::LLVMPtr{T,AS}, i::I, - ::Val{align}) where {T,AS,I,align} +@inline function unsafe_invariant_load(ptr::LLVMPtr{T}, i::I, ::Val{align}) where {T,I,align} sizeof(T) == 0 && return T.instance - ispow2(align) || - return :(error("unsafe_invariant_load: alignment must be a power of 2, got $($align)")) - @dispose ctx=Context() begin - eltyp = convert(LLVMType, T) - T_idx = convert(LLVMType, I) - T_ptr = convert(LLVMType, ptr) - T_typed_ptr = LLVM.PointerType(eltyp, AS) - - llvm_f, _ = create_function(eltyp, LLVMType[T_ptr, T_idx]) - - @dispose builder=IRBuilder() begin - entry = BasicBlock(llvm_f, "entry") - position!(builder, entry) - base = if supports_typed_pointers(ctx) - bitcast!(builder, parameters(llvm_f)[1], T_typed_ptr) - else - parameters(llvm_f)[1] - end - gep = inbounds_gep!(builder, eltyp, base, [parameters(llvm_f)[2]]) - ld = load!(builder, eltyp, gep) - if AS != 0 - metadata(ld)[LLVM.MD_tbaa] = tbaa_addrspace(AS) - end - metadata(ld)[LLVM.MD_invariant_load] = MDNode(LLVM.Metadata[]) - alignment!(ld, align) - - ret!(builder, ld) - end - - call_function(llvm_f, T, Tuple{LLVMPtr{T,AS},I}, :ptr, :(i-one(I))) + ispow2(align) || error("unsafe_invariant_load: alignment must be a power of 2, got ", align) + return _unsafe_invariant_load(ptr, i - one(I), Val(align)) +end +@llvmgenerated builder function _unsafe_invariant_load(ptr::LLVMPtr{T,AS}, i::Integer, + ::Val{align})::T where {T,AS,align} + eltyp = convert(LLVMType, T) + if supports_typed_pointers(LLVM.context()) + ptr = bitcast!(builder, ptr, LLVM.PointerType(eltyp, AS)) + end + ld = load!(builder, eltyp, inbounds_gep!(builder, eltyp, ptr, [i])) + if AS != 0 + metadata(ld)[LLVM.MD_tbaa] = tbaa_addrspace(AS) end + metadata(ld)[LLVM.MD_invariant_load] = MDNode(LLVM.Metadata[]) + alignment!(ld, align) + ld end @device_function @inline function const_arrayref(A::CLDeviceArray{T}, index::Integer) where {T} diff --git a/src/device/random.jl b/src/device/random.jl index 88b7b225..53995242 100644 --- a/src/device/random.jl +++ b/src/device/random.jl @@ -150,36 +150,22 @@ end # a hacky method of exposing constant tables as constant GPU memory function emit_constant_array(name::Symbol, data::AbstractArray{T}) where {T} - @dispose ctx=Context() begin + generate_llvmcall(LLVMPtr{T,AS.UniformConstant}, Tuple{}) do builder T_val = convert(LLVMType, T) T_ptr = convert(LLVMType, LLVMPtr{T,AS.UniformConstant}) - # define function and get LLVM module - llvm_f, _ = create_function(T_ptr) - mod = LLVM.parent(llvm_f) - # create a global memory global variable # TODO: global_var alignment? T_global = LLVM.ArrayType(T_val, length(data)) # XXX: why can't we use a single name like emit_shmem - gv = GlobalVariable(mod, T_global, "gpu_$(name)_data", AS.UniformConstant) + gv = GlobalVariable(current_module(builder), T_global, "gpu_$(name)_data", + AS.UniformConstant) linkage!(gv, LLVM.API.LLVMInternalLinkage) initializer!(gv, ConstantArray(data)) alignment!(gv, 16) - # generate IR - @dispose builder=IRBuilder() begin - entry = BasicBlock(llvm_f, "entry") - position!(builder, entry) - - ptr = gep!(builder, T_global, gv, [ConstantInt(0), ConstantInt(0)]) - - untyped_ptr = bitcast!(builder, ptr, T_ptr) - - ret!(builder, untyped_ptr) - end - - call_function(llvm_f, LLVMPtr{T,AS.UniformConstant}) + ptr = gep!(builder, T_global, gv, [ConstantInt(0), ConstantInt(0)]) + bitcast!(builder, ptr, T_ptr) end end diff --git a/src/device/runtime.jl b/src/device/runtime.jl index c691b250..b997d31d 100644 --- a/src/device/runtime.jl +++ b/src/device/runtime.jl @@ -259,30 +259,10 @@ function additional_arg_intr(mod::LLVM.Module, T_state, name) end # run-time equivalent -function additional_arg_value(state, name) - @dispose ctx=Context() begin - T_state = convert(LLVMType, state) - - # create function - llvm_f, _ = create_function(T_state) - mod = LLVM.parent(llvm_f) - - # get intrinsic - state_intr = additional_arg_intr(mod, T_state, name) - state_intr_ft = function_type(state_intr) - - # generate IR - @dispose builder=IRBuilder() begin - entry = BasicBlock(llvm_f, "entry") - position!(builder, entry) - - val = call!(builder, state_intr_ft, state_intr, Value[], name) - - ret!(builder, val) - end - - call_function(llvm_f, state) - end +additional_arg_value(state, name) = generate_llvmcall(state, Tuple{}) do builder + T_state = convert(LLVMType, state) + state_intr = additional_arg_intr(current_module(builder), T_state, name) + call!(builder, function_type(state_intr), state_intr, Value[], name) end for name in [:random_keys, :random_counters]