diff --git a/lib/level-zero/module.jl b/lib/level-zero/module.jl index 0c9ddbb1..68a5af61 100644 --- a/lib/level-zero/module.jl +++ b/lib/level-zero/module.jl @@ -79,6 +79,7 @@ export ZeKernel mutable struct ZeKernel mod::ZeModule handle::ze_kernel_handle_t + lock::ReentrantLock function ZeKernel(mod, name) GC.@preserve name begin @@ -86,7 +87,7 @@ mutable struct ZeKernel 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) @@ -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 diff --git a/src/compiler/compilation.jl b/src/compiler/compilation.jl index b49d17b2..a8c117a8 100644 --- a/src/compiler/compilation.jl +++ b/src/compiler/compilation.jl @@ -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()) diff --git a/src/compiler/execution.jl b/src/compiler/execution.jl index 327cf8c7..7cb37cff 100644 --- a/src/compiler/execution.jl +++ b/src/compiler/execution.jl @@ -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 diff --git a/src/context.jl b/src/context.jl index f625a40d..15a1037a 100644 --- a/src/context.jl +++ b/src/context.jl @@ -146,6 +146,7 @@ function is_integrated(dev::ZeDevice=device()) end const global_contexts = Dict{ZeDriver,ZeContext}() +const global_contexts_lock = ReentrantLock() """ context() -> ZeContext @@ -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 diff --git a/src/gpuarrays.jl b/src/gpuarrays.jl index e18b3c72..a0a38167 100644 --- a/src/gpuarrays.jl +++ b/src/gpuarrays.jl @@ -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)