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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
131 changes: 63 additions & 68 deletions src/compiler/compilation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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))]))

Expand All @@ -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
Expand All @@ -105,32 +107,29 @@ 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))

# Replace the original uses, excluding the insertvalue that constructs the state.
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

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