Skip to content
Merged
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
4 changes: 2 additions & 2 deletions 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.3.2"
KernelAbstractions = "0.9.38"
LLVM = "9.6"
LLVM = "10"
LinearAlgebra = "1"
OpenCL_jll = "=2024.10.24"
Preferences = "1"
Expand All @@ -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"
Expand Down
4 changes: 2 additions & 2 deletions lib/intrinsics/Project.toml
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
name = "SPIRVIntrinsics"
uuid = "71d1d633-e7e8-4a92-83a1-de8814b09ba8"
authors = ["Tim Besard <tim.besard@gmail.com>"]
version = "1.1.4"
version = "1.2.0"

[deps]
ExprTools = "e2ba6199-217a-4e67-a87a-7c52f15ade04"
Expand All @@ -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"
2 changes: 1 addition & 1 deletion lib/intrinsics/src/SPIRVIntrinsics.jl
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
module SPIRVIntrinsics

using LLVM, LLVM.Interop
using LLVM, LLVM.IR, LLVM.Build, LLVM.Interop
using Core: LLVMPtr

import ExprTools
Expand Down
68 changes: 26 additions & 42 deletions lib/intrinsics/src/memory.jl
Original file line number Diff line number Diff line change
@@ -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
114 changes: 51 additions & 63 deletions lib/intrinsics/src/printf.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down
44 changes: 20 additions & 24 deletions lib/intrinsics/src/synchronization.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
46 changes: 24 additions & 22 deletions lib/intrinsics/src/work_item.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand All @@ -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
2 changes: 1 addition & 1 deletion src/OpenCL.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading