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
7 changes: 6 additions & 1 deletion lib/level-zero/module.jl
Original file line number Diff line number Diff line change
Expand Up @@ -79,14 +79,15 @@ export ZeKernel
mutable struct ZeKernel
mod::ZeModule
handle::ze_kernel_handle_t
lock::ReentrantLock

function ZeKernel(mod, name)
GC.@preserve name begin
desc_ref = Ref(ze_kernel_desc_t(; pKernelName=pointer(name)))
handle_ref = Ref{ze_kernel_handle_t}()
zeKernelCreate(mod, desc_ref, handle_ref)
end
obj = new(mod, handle_ref[])
obj = new(mod, handle_ref[], ReentrantLock())

finalizer(obj) do obj
zeKernelDestroy(obj)
Expand All @@ -95,6 +96,10 @@ mutable struct ZeKernel
end
end

Base.lock(kernel::ZeKernel) = lock(getfield(kernel, :lock))
Base.lock(f::Function, kernel::ZeKernel) = lock(f, getfield(kernel, :lock))
Base.unlock(kernel::ZeKernel) = unlock(getfield(kernel, :lock))

Base.unsafe_convert(::Type{ze_kernel_handle_t}, kernel::ZeKernel) = kernel.handle

Base.:(==)(a::ZeKernel, b::ZeKernel) = a.handle == b.handle
Expand Down
10 changes: 5 additions & 5 deletions src/compiler/compilation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -302,14 +302,14 @@ end
# cache of compiler configurations, per device (but additionally configurable via kwargs)
const _toolchain = Ref{Any}()
const _compiler_configs = Dict{UInt, oneAPICompilerConfig}()
const compiler_config_lock = ReentrantLock()
function compiler_config(dev; kwargs...)
h = hash(dev.driver, hash(dev, hash(kwargs)))
config = get(_compiler_configs, h, nothing)
if config === nothing
config = _compiler_config(dev; kwargs...)
_compiler_configs[h] = config
Base.@lock compiler_config_lock begin
get!(_compiler_configs, h) do
_compiler_config(dev; kwargs...)
end
end
return config
end
# Whether the driver's SPIR-V runtime accepts the SPV_KHR_bfloat16 extension.
function _driver_supports_bfloat16_spirv(dev=device())
Expand Down
14 changes: 8 additions & 6 deletions src/compiler/execution.jl
Original file line number Diff line number Diff line change
Expand Up @@ -323,13 +323,15 @@ const _kernel_instances = Dict{UInt, Any}()

@inline function onecall(kernel::ZeKernel, tt, args...; groups::ZeDim=1, items::ZeDim=1,
queue::ZeCommandQueue=global_queue(context(), device()))
for (i, arg) in enumerate(args)
oneL0.arguments(kernel)[i] = arg
end
Base.@lock kernel begin
for (i, arg) in enumerate(args)
oneL0.arguments(kernel)[i] = arg
end

groupsize!(kernel, items)
execute!(queue) do list
append_launch!(list, kernel, groups)
groupsize!(kernel, items)
execute!(queue) do list
append_launch!(list, kernel, groups)
end
end
end

Expand Down
8 changes: 6 additions & 2 deletions src/context.jl
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,7 @@ function is_integrated(dev::ZeDevice=device())
end

const global_contexts = Dict{ZeDriver,ZeContext}()
const global_contexts_lock = ReentrantLock()

"""
context() -> ZeContext
Expand All @@ -166,8 +167,11 @@ See also: [`context!`](@ref), [`driver`](@ref)
"""
function context()
get!(task_local_storage(), :ZeContext) do
get!(global_contexts, driver()) do
ZeContext(driver())
drv = driver()
Base.@lock global_contexts_lock begin
get!(global_contexts, drv) do
ZeContext(drv)
end
end
end
end
Expand Down
6 changes: 4 additions & 2 deletions src/gpuarrays.jl
Original file line number Diff line number Diff line change
@@ -1,9 +1,11 @@
# GPUArrays.jl interface

const GLOBAL_RNGs = Dict{ZeDevice,GPUArrays.RNG}()
function GPUArrays.default_rng(::Type{<:oneArray})
dev = device()
get!(GLOBAL_RNGs, dev) do
rngs = get!(task_local_storage(), :oneAPI_GLOBAL_RNGs) do
Dict{ZeDevice,GPUArrays.RNG}()
end
get!(rngs, dev) do
N = oneL0.compute_properties(dev).maxTotalGroupSize
state = oneArray{NTuple{4, UInt32}}(undef, N)
rng = GPUArrays.RNG(state)
Expand Down
Loading