diff --git a/Project.toml b/Project.toml index cbaf7cdf..ec86c38c 100644 --- a/Project.toml +++ b/Project.toml @@ -35,7 +35,7 @@ GPUArrays = "11.2.1" GPUCompiler = "2.9" GPUToolbox = "3.3.2" KernelAbstractions = "0.9.38" -LLVM = "9.6" +LLVM = "10" LinearAlgebra = "1" OpenCL_jll = "=2024.10.24" Preferences = "1" @@ -44,7 +44,7 @@ Random = "1" Random123 = "1.7.1" RandomNumbers = "1.6.0" Reexport = "1" -SPIRVIntrinsics = "1.1" +SPIRVIntrinsics = "1.2" SPIRV_LLVM_Backend_jll = "23" SPIRV_Tools_jll = "2025.1" StaticArrays = "1" diff --git a/lib/intrinsics/Project.toml b/lib/intrinsics/Project.toml index 8a6a2fbd..57f56bd7 100644 --- a/lib/intrinsics/Project.toml +++ b/lib/intrinsics/Project.toml @@ -1,7 +1,7 @@ name = "SPIRVIntrinsics" uuid = "71d1d633-e7e8-4a92-83a1-de8814b09ba8" authors = ["Tim Besard "] -version = "1.1.4" +version = "1.2.0" [deps] ExprTools = "e2ba6199-217a-4e67-a87a-7c52f15ade04" @@ -18,7 +18,7 @@ SPIRVIntrinsicsSIMDExt = "SIMD" [compat] ExprTools = "0.1" GPUToolbox = "0.2, 0.3, 1, 2, 3" -LLVM = "9.1" +LLVM = "10" SIMD = "3.6" SpecialFunctions = "1.3, 2" julia = "1.10" diff --git a/lib/intrinsics/src/SPIRVIntrinsics.jl b/lib/intrinsics/src/SPIRVIntrinsics.jl index ee432ca5..cf694e3e 100644 --- a/lib/intrinsics/src/SPIRVIntrinsics.jl +++ b/lib/intrinsics/src/SPIRVIntrinsics.jl @@ -1,6 +1,6 @@ module SPIRVIntrinsics -using LLVM, LLVM.Interop +using LLVM, LLVM.IR, LLVM.Build, LLVM.Interop using Core: LLVMPtr import ExprTools diff --git a/lib/intrinsics/src/memory.jl b/lib/intrinsics/src/memory.jl index 9a17b447..7e796eb9 100644 --- a/lib/intrinsics/src/memory.jl +++ b/lib/intrinsics/src/memory.jl @@ -1,47 +1,31 @@ # 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) - - # determine the array size: an array of a bits union stores a selector byte per - # element after the values - sz = len * sizeof(T) - Base.isbitsunion(T) && (sz += len) - - # create the global variable - mod = LLVM.parent(llvm_f) - gv_typ = LLVM.ArrayType(eltyp, sz) - 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 - align = 1 - for typ in Base.uniontypes(T) - typ.layout != C_NULL && (align = max(align, Base.datatype_alignment(typ))) - end - alignment!(gv, align) - - # 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}) + + # determine the array size: an array of a bits union stores a selector byte per + # element after the values + sz = len * sizeof(T) + Base.isbitsunion(T) && (sz += len) + + # create the global variable + gv_typ = LLVM.ArrayType(eltyp, sz) + gv = GlobalVariable(current_module(builder), gv_typ, "local_memory", AS.Workgroup) + if len > 0 + gv.linkage = LLVM.Linkage.Internal + gv.initializer = null(gv_typ) end + # TODO: Make the alignment configurable + align = 1 + for typ in Base.uniontypes(T) + typ.layout != C_NULL && (align = max(align, Base.datatype_alignment(typ))) + end + gv.alignment = align + + 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..28aa8066 100644 --- a/lib/intrinsics/src/printf.jl +++ b/lib/intrinsics/src/printf.jl @@ -26,73 +26,61 @@ end @generated function emit_printf(::Val{fmt}, argspec...) where {fmt} arg_exprs = [:( argspec[$i] ) for i in 1:length(argspec)] - arg_types = [argspec...] - - Context() do ctx - T_void = LLVM.VoidType() - 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) - else - actual_typ = convert(LLVMType, argtyp) - actual_arg = arg - end - push!(T_actual_args, actual_typ) - push!(actual_args, actual_arg) - end - - str = globalstring_ptr!(builder, String(fmt); addrspace=AS.UniformConstant) + arg_types = Tuple{argspec...} - # 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...]) + # pass the format string and the argument types as statically-known arguments, so that + # the IR generator is compiled once instead of for every format string + generate_llvmcall(printf_ir, Int32, Tuple{Val{fmt}, Type{arg_types}, argspec...}, + Val(fmt), arg_types, arg_exprs...) +end - ret!(builder, chars) +function printf_ir(builder, fmt::Val, arg_types::Type{<:Tuple}, args...) + @nospecialize + T_int32 = LLVM.Int32Type() + T_pint8 = LLVM.PointerType(LLVM.Int8Type(), AS.UniformConstant) + + # `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`). + actual_args = LLVM.Value[] + for (arg, argtyp) in zip(args, fieldtypes(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 arg.value_type 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 + actual_typ = LLVM.Int1Type() + actual_arg = trunc!(builder, arg, actual_typ) + else + actual_arg = arg end - - call_function(llvm_f, Int32, Tuple{arg_types...}, arg_exprs...) + push!(actual_args, actual_arg) end + + fmt_str = String(typeof(fmt).parameters[1]::Symbol) + str = globalstring_ptr!(builder, fmt_str; 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!(printf.function_attributes, EnumAttribute(:nobuiltin)) + call!(builder, printf_typ, printf, [str, actual_args...]) end diff --git a/lib/intrinsics/src/synchronization.jl b/lib/intrinsics/src/synchronization.jl index e8cf6aef..c21af60b 100644 --- a/lib/intrinsics/src/synchronization.jl +++ b/lib/intrinsics/src/synchronization.jl @@ -34,38 +34,34 @@ end # XXX: is this even needed? Doesn't LLVM reconstruct these? # using the `@builtin_ccall` version causes validation issues. +# call a builtin that is `convergent`, i.e., that must not be made control-dependent on +# additional values +function convergent_call!(builder::IRBuilder, name::String, args::Vector{<:Value}) + ft = LLVM.FunctionType(LLVM.VoidType(), [arg.value_type for arg in args]) + f = LLVM.Function(current_module(builder), name, ft) + push!(f.function_attributes, EnumAttribute(:convergent)) + call!(builder, ft, f, args) + return +end + #@device_function @inline memory_barrier(scope, semantics) = # @builtin_ccall("__spirv_MemoryBarrier", Cvoid, (UInt32, UInt32), scope, semantics) @device_function memory_barrier(scope, semantics) = - Base.llvmcall((""" - declare void @_Z21__spirv_MemoryBarrierjj(i32, i32) #0 - define void @entry(i32 %scope, i32 %semantics) #1 { - call void @_Z21__spirv_MemoryBarrierjj(i32 %scope, i32 %semantics) - ret void - } - attributes #0 = { convergent } - attributes #1 = { alwaysinline } - """, "entry"), - Cvoid, Tuple{UInt32, UInt32}, convert(UInt32, scope), convert(UInt32, semantics)) + _memory_barrier(convert(UInt32, scope), convert(UInt32, semantics)) +@llvmgenerated builder function _memory_barrier(scope::UInt32, semantics::UInt32)::Nothing + convergent_call!(builder, "_Z21__spirv_MemoryBarrierjj", [scope, semantics]) +end #@device_function @inline control_barrier(execution_scope, memory_scope, memory_semantics) = # @builtin_ccall("__spirv_ControlBarrier", Cvoid, (UInt32, UInt32, UInt32), # execution_scope, memory_scope, memory_semantics) @device_function @inline control_barrier(execution_scope, memory_scope, memory_semantics) = - Base.llvmcall((""" - declare void @_Z22__spirv_ControlBarrierjjj(i32, i32, i32) #0 - define void @entry(i32 %execution, i32 %memory, i32 %semantics) #1 { - call void @_Z22__spirv_ControlBarrierjjj(i32 %execution, i32 %memory, i32 %semantics) - ret void - } - attributes #0 = { convergent } - attributes #1 = { alwaysinline } - """, "entry"), - Cvoid, - Tuple{UInt32, UInt32, UInt32}, - convert(UInt32, execution_scope), - convert(UInt32, memory_scope), - convert(UInt32, memory_semantics)) + _control_barrier(convert(UInt32, execution_scope), convert(UInt32, memory_scope), + convert(UInt32, memory_semantics)) +@llvmgenerated builder function _control_barrier(execution::UInt32, memory::UInt32, + semantics::UInt32)::Nothing + convergent_call!(builder, "_Z22__spirv_ControlBarrierjjj", [execution, memory, semantics]) +end ## OpenCL-compatible fence API diff --git a/lib/intrinsics/src/work_item.jl b/lib/intrinsics/src/work_item.jl index 5cf49bbf..9005ed0d 100644 --- a/lib/intrinsics/src/work_item.jl +++ b/lib/intrinsics/src/work_item.jl @@ -5,6 +5,25 @@ # NOTE: these functions now unsafely truncate to Int to avoid top bit checks. # we should probably use range metadata instead. +# load a built-in variable, which the SPIR-V back-end expects as an external global in the +# Input storage class +@llvmgenerated builder function builtin_variable(::Val{name}, ::Type{T})::T where {name,T} + T_val = convert(LLVMType, T) + gv = GlobalVariable(current_module(builder), T_val, String(name), AS.Input) + load!(builder, T_val, gv) +end + +# load a component of a built-in vector variable, by calling the function that the SPIR-V +# back-end lowers to a load and an extract of the component +@llvmgenerated builder function builtin_vector_variable(::Val{name}, idx::Int32)::UInt where {name} + ft = LLVM.FunctionType(convert(LLVMType, UInt), [idx.value_type]) + f = LLVM.Function(current_module(builder), String(name), ft) + push!(f.function_attributes, EnumAttribute(:nounwind)) + push!(f.function_attributes, EnumAttribute(:willreturn)) + f.memory_effects = MemoryEffects(:none) + call!(builder, ft, f, [idx]) +end + # 1D values for (julia_name, (spirv_name, julia_type, offset)) in [ # indices @@ -18,19 +37,11 @@ for (julia_name, (spirv_name, julia_type, offset)) in [ :get_max_sub_group_size => (:BuiltInSubgroupMaxSize, UInt32, 0), :get_num_sub_groups => (:BuiltInNumSubgroups, UInt32, 0), :get_enqueued_num_sub_groups => (:BuiltInNumEnqueuedSubgroups, UInt32, 0)] - gvar_name = Symbol("@__spirv_$(spirv_name)") - width = sizeof(julia_type) * 8 + gvar_name = Symbol("__spirv_$(spirv_name)") @eval begin export $julia_name @device_function $julia_name() = - Base.llvmcall( - $("""$gvar_name = external addrspace($(AS.Input)) global i$(width) - define i$(width) @entry() #0 { - %val = load i$(width), i$(width) addrspace($(AS.Input))* $gvar_name - ret i$(width) %val - } - attributes #0 = { alwaysinline } - """, "entry"), $julia_type, Tuple{}) % Int + $offset + builtin_variable(Val($(QuoteNode(gvar_name))), $julia_type) % Int + $offset end end @@ -52,20 +63,11 @@ for (julia_name, (spirv_name, offset)) in [ :get_enqueued_local_size => (:BuiltInEnqueuedWorkgroupSize, 0), :get_num_groups => (:BuiltInNumWorkgroups, 0)] fname = "__spirv_$(spirv_name)" - mangled = "_Z$(length(fname))$(fname)i" - push!(known_intrinsics, mangled) - width = Int === Int64 ? 64 : 32 + mangled = Symbol("_Z$(length(fname))$(fname)i") + push!(known_intrinsics, String(mangled)) @eval begin export $julia_name @device_function $julia_name(dimindx::Integer=1u32) = - Base.llvmcall( - $("""declare i$(width) @$(mangled)(i32) #0 - define i$(width) @entry(i32 %idx) #1 { - %val = call i$(width) @$(mangled)(i32 %idx) - ret i$(width) %val - } - attributes #0 = { nounwind readnone willreturn } - attributes #1 = { alwaysinline } - """, "entry"), UInt, Tuple{Int32}, (dimindx - 1u32) % Int32) % Int + $offset + builtin_vector_variable(Val($(QuoteNode(mangled))), (dimindx - 1u32) % Int32) % Int + $offset end end diff --git a/src/OpenCL.jl b/src/OpenCL.jl index 8756602f..3e3ef109 100644 --- a/src/OpenCL.jl +++ b/src/OpenCL.jl @@ -2,7 +2,7 @@ module OpenCL using GPUCompiler import GPUToolbox -using LLVM, LLVM.Interop +using LLVM, LLVM.IR, LLVM.Build, LLVM.Interop using SPIRV_LLVM_Backend_jll, SPIRV_Tools_jll, spirv2clc_jll using Adapt using Reexport diff --git a/src/compiler/capabilities.jl b/src/compiler/capabilities.jl index 52562bad..631dcfd6 100644 --- a/src/compiler/capabilities.jl +++ b/src/compiler/capabilities.jl @@ -156,15 +156,12 @@ feature_supported(dev::cl.Device, name::Symbol) = feature_supported(device_featu # Load the feature bitset that `finish_module!` materializes as a module-scope constant. Once the # constant is in place the load folds away, so `has_feature` branches resolve at compile time. The # global uses the UniformConstant (2) storage class to stay valid SPIR-V if it ever survives. -@device_function @inline function feature_bitset() - Base.llvmcall( - ("""@__opencl_feature_bitset = external addrspace(2) global i64 - define i64 @entry() #0 { - %v = load i64, i64 addrspace(2)* @__opencl_feature_bitset - ret i64 %v - } - attributes #0 = { alwaysinline } - """, "entry"), UInt64, Tuple{}) +@device_function @inline feature_bitset() = _feature_bitset() +@llvmgenerated builder function _feature_bitset()::UInt64 + T = LLVM.Int64Type() + gv = GlobalVariable(current_module(builder), T, "__opencl_feature_bitset", + AS.UniformConstant) + load!(builder, T, gv) end export has_feature diff --git a/src/compiler/compilation.jl b/src/compiler/compilation.jl index 413fd787..4cbe085c 100644 --- a/src/compiler/compilation.jl +++ b/src/compiler/compilation.jl @@ -73,22 +73,22 @@ function GPUCompiler.finish_module!(@nospecialize(job::OpenCLCompilerJob), sg_size = job.config.params.sub_group_size if sg_size !== nothing - metadata(entry)["intel_reqd_sub_group_size"] = MDNode([ConstantInt(Int32(sg_size))]) + entry.metadata["intel_reqd_sub_group_size"] = MDNode([ConstantInt(Int32(sg_size))]) end # materialize the feature bitset for `has_feature`, if the kernel referenced it. A constant # initializer plus private linkage lets the optimizer fold the loads and drop the global, so it # never reaches SPIR-V. - if haskey(globals(mod), "__opencl_feature_bitset") - gv = globals(mod)["__opencl_feature_bitset"] - initializer!(gv, ConstantInt(LLVM.Int64Type(), job.config.params.features)) - linkage!(gv, LLVM.API.LLVMPrivateLinkage) - constant!(gv, true) + gv = get(mod.globals, "__opencl_feature_bitset", nothing) + if gv !== nothing + gv.initializer = ConstantInt(LLVM.Int64Type(), job.config.params.features) + gv.linkage = LLVM.Linkage.Private + gv.constant = true end # if this kernel uses our RNG, we should prime the shared state. # XXX: these transformations should really happen at the Julia IR level... - if haskey(functions(mod), "julia.opencl.random_keys") && job.config.kernel + if haskey(mod.functions, "julia.opencl.random_keys") && job.config.kernel # insert call to `initialize_rng_state` f = initialize_rng_state ft = typeof(f) @@ -102,30 +102,22 @@ function GPUCompiler.finish_module!(@nospecialize(job::OpenCLCompilerJob), GPUCompiler.deferred_codegen_jobs[id] = job # generate IR for calls to `deferred_codegen` and the resulting function pointer - top_bb = first(blocks(entry)) - bb = BasicBlock(top_bb, "initialize_rng") + top_bb = entry.entry + bb = BasicBlock(LLVM.before(top_bb), "initialize_rng") @dispose builder=IRBuilder() begin - position!(builder, bb) - subprogram = LLVM.subprogram(entry) + position!(builder, LLVM.at_end(bb)) + subprogram = entry.subprogram if subprogram !== nothing loc = DILocation(0, 0, subprogram) - debuglocation!(builder, loc) + builder.debug_location = loc end - debuglocation!(builder, first(instructions(top_bb))) # call the `deferred_codegen` marker function - T_ptr = if LLVM.version() >= v"17" - LLVM.PointerType() - elseif VERSION >= v"1.12.0-DEV.225" - LLVM.PointerType(LLVM.Int8Type()) - else - LLVM.Int64Type() - end + # (declared like GPUCompiler's `ccall("extern deferred_codegen", llvmcall, Ptr{Cvoid}, ...)`) + T_ptr = convert(LLVMType, Ptr{Cvoid}) T_id = convert(LLVMType, Int) deferred_codegen_ft = LLVM.FunctionType(T_ptr, [T_id]) - deferred_codegen = if haskey(functions(mod), "deferred_codegen") - functions(mod)["deferred_codegen"] - else + deferred_codegen = get!(mod.functions, "deferred_codegen") do LLVM.Function(mod, "deferred_codegen", deferred_codegen_ft) end fptr = call!(builder, deferred_codegen_ft, deferred_codegen, [ConstantInt(id)]) @@ -139,7 +131,7 @@ function GPUCompiler.finish_module!(@nospecialize(job::OpenCLCompilerJob), br!(builder, top_bb) # note the use of the device-side RNG in this kernel - push!(function_attributes(entry), StringAttribute("julia.opencl.rng", "")) + push!(entry.function_attributes, StringAttribute("julia.opencl.rng", "")) end # XXX: put some of the above behind GPUCompiler abstractions @@ -247,9 +239,13 @@ function compile_to_obj(@nospecialize(job::CompilerJob)) JuliaContext() do ctx obj, meta = GPUCompiler.compile(:obj, job) - entry = LLVM.name(meta.entry) - device_rng = StringAttribute("julia.opencl.rng", "") in collect(function_attributes(meta.entry)) - (; obj, entry, device_rng) + + # we own the IR: inspect it, then dispose of it + @dispose ir=meta.ir begin + entry = meta.entry.name + device_rng = haskey(meta.entry.function_attributes, "julia.opencl.rng") + (; obj, entry, device_rng) + end end end diff --git a/src/device/array.jl b/src/device/array.jl index 62466015..85a7851c 100644 --- a/src/device/array.jl +++ b/src/device/array.jl @@ -165,40 +165,25 @@ 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} +# +# like `unsafe_load`, the index is widened to `Int` in Julia, where its signedness is known, +# because `getelementptr` sign-extends narrower indices. +@inline function unsafe_invariant_load(ptr::LLVMPtr{T}, i::Integer, ::Val{align}) where {T,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, Int(i) - 1, Val(align)) +end +@llvmgenerated builder function _unsafe_invariant_load(ptr::LLVMPtr{T,AS}, i::Int, + ::Val{align})::T where {T,AS,align} + eltyp = convert(LLVMType, T) + # `LLVMPtr` is an `i8*` with typed pointers (with opaque pointers, this cast folds away) + ptr = bitcast!(builder, ptr, LLVM.PointerType(eltyp, AS)) + ld = load!(builder, eltyp, inbounds_gep!(builder, eltyp, ptr, [i]); align) + if AS != 0 + ld.metadata[MD_tbaa] = tbaa_addrspace(AS) end + ld.metadata[MD_invariant_load] = MDNode(LLVM.Metadata[]) + 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..3ce97e66 100644 --- a/src/device/random.jl +++ b/src/device/random.jl @@ -150,37 +150,34 @@ 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 - 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) - 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) + generate_llvmcall(ConstantArrayIR(name, data), LLVMPtr{T,AS.UniformConstant}, Tuple{}) +end - ret!(builder, untyped_ptr) - end +# the IR generator of `emit_constant_array`: a callable object with abstractly-typed fields +# instead of a closure, so that it is compiled once instead of for every element type +struct ConstantArrayIR + name::Symbol + data::AbstractArray +end - call_function(llvm_f, LLVMPtr{T,AS.UniformConstant}) - end +function (gen::ConstantArrayIR)(builder) + (; name, data) = gen + T = eltype(data) + T_val = convert(LLVMType, T) + T_ptr = convert(LLVMType, LLVMPtr{T,AS.UniformConstant}) + + # 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(current_module(builder), T_global, "gpu_$(name)_data", + AS.UniformConstant) + gv.linkage = LLVM.Linkage.Internal + gv.initializer = ConstantArray(data) + gv.alignment = 16 + + ptr = gep!(builder, T_global, gv, [ConstantInt(0), ConstantInt(0)]) + bitcast!(builder, ptr, T_ptr) end for var in [:ki, :wi, :fi, :ke, :we, :fe] diff --git a/src/device/runtime.jl b/src/device/runtime.jl index c691b250..dee4ec57 100644 --- a/src/device/runtime.jl +++ b/src/device/runtime.jl @@ -248,44 +248,19 @@ end # then get propagated across function calls to the caller. function additional_arg_intr(mod::LLVM.Module, T_state, name) - state_intr = if haskey(functions(mod), "julia.opencl.$name") - functions(mod)["julia.opencl.$name"] - else - LLVM.Function(mod, "julia.opencl.$name", LLVM.FunctionType(T_state)) + get!(mod.functions, "julia.opencl.$name") do + state_intr = LLVM.Function(mod, "julia.opencl.$name", LLVM.FunctionType(T_state)) + state_intr.memory_effects = MemoryEffects(:none) + state_intr end - push!(function_attributes(state_intr), EnumAttribute("readnone", 0)) - - return state_intr 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 +@llvmgenerated builder function additional_arg_value(::Type{T}, ::Val{name})::T where {T,name} + state_intr = additional_arg_intr(current_module(builder), convert(LLVMType, T), name) + call!(builder, state_intr.function_type, state_intr, Value[], String(name)) end for name in [:random_keys, :random_counters] - @eval @inline @generated $name() = - additional_arg_value(LLVMPtr{UInt32, AS.Workgroup}, $(String(name))) + @eval @inline $name() = additional_arg_value(LLVMPtr{UInt32, AS.Workgroup}, Val($(QuoteNode(name)))) end diff --git a/test/intrinsics.jl b/test/intrinsics.jl index 590443de..4c1b41d8 100644 --- a/test/intrinsics.jl +++ b/test/intrinsics.jl @@ -391,6 +391,11 @@ end (inv, ComplexF32, (ComplexF32,))) tt = Tuple{CLDeviceArray{T,0,AS.CrossWorkgroup}, typeof(f), args...} job = GPUCompiler.CompilerJob(GPUCompiler.methodinstance(typeof(kernel), tt), config) - @test GPUCompiler.JuliaContext(_ -> GPUCompiler.compile(:llvm, job)) !== nothing + # the IR belongs to the caller of `compile` + @test GPUCompiler.JuliaContext() do _ + ir, meta = GPUCompiler.compile(:llvm, job) + GPUCompiler.LLVM.dispose(ir) + meta !== nothing + end end end