Skip to content
Closed
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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
2 changes: 1 addition & 1 deletion lib/intrinsics/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
50 changes: 17 additions & 33 deletions lib/intrinsics/src/memory.jl
Original file line number Diff line number Diff line change
@@ -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
96 changes: 40 additions & 56 deletions lib/intrinsics/src/printf.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
48 changes: 16 additions & 32 deletions src/device/array.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down
24 changes: 5 additions & 19 deletions src/device/random.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
28 changes: 4 additions & 24 deletions src/device/runtime.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
Loading