diff --git a/Project.toml b/Project.toml index 81054df..c7ff0c2 100644 --- a/Project.toml +++ b/Project.toml @@ -1,6 +1,6 @@ name = "GPUToolbox" uuid = "096a3bc2-3ced-46d0-87f4-dd12716f4bfc" -version = "3.2.0" +version = "3.3.0" [deps] LLVM = "929cbde3-209d-540e-8aea-75f648917ca0" diff --git a/src/synchronization.jl b/src/synchronization.jl index 98d0140..42811db 100644 --- a/src/synchronization.jl +++ b/src/synchronization.jl @@ -17,6 +17,19 @@ export cooperative_wait const SPIN_BUSY_ITERATIONS = 32 const SPIN_ITERATIONS = 256 +# for devices that execute on the host's CPU cores, polling for that long competes with the +# operation for those cores, so they only busy-wait for a limited time (in nanoseconds) +function spin_until(isdone::F, obj, budget::UInt64) where {F} + isdone(obj) && return true + t0 = time_ns() + while time_ns() - t0 < budget + ccall(:jl_cpu_pause, Cvoid, ()) + GC.safepoint() + isdone(obj) && return true + end + return false +end + function spin_until(isdone::F, obj) where {F} isdone(obj) && return true for i in 1:SPIN_ITERATIONS @@ -157,24 +170,128 @@ function submit!(state, request::WaitRequest, wait::Bool) end +## slow path: completion notifications + +# some drivers can notify us when an operation completes, by calling back from a thread of +# their own. that avoids a worker thread, which matters for devices that execute on the +# host's CPU cores: waking a worker while the operation starts delays it considerably. +# +# the callback directly wakes the waiting task, like a worker does, instead of going through +# the event loop, which would add a thread hop (and depends on the thread running the event +# loop to be available). the calling thread is adopted by Julia, which is safe as long as +# driver calls that may be waiting on the callback are GC-safe. + +# a spin lock-based condition, so that signalling from a driver thread never blocks it in +# Julia's scheduler (as a `ReentrantLock` could) +mutable struct Completion + const cond::Base.ThreadSynchronizer + @atomic state::Int # one of the constants below + + Completion() = new(Base.ThreadSynchronizer(), PENDING) +end +const PENDING = 0 +const SIGNALLED = 1 +const RELEASED = 2 # the callback does not access the completion anymore + +# completions whose waiter stopped waiting (e.g., because it was cancelled), kept alive +# until the driver has signalled them +const abandoned_completions = Completion[] +const abandoned_completions_lock = ReentrantLock() + +function sweep_abandoned_completions() + filter!(c -> (@atomic :acquire c.state) != RELEASED, abandoned_completions) +end + +function abandon_completion(c::Completion) + @lock abandoned_completions_lock begin + sweep_abandoned_completions() + push!(abandoned_completions, c) + end + return +end + +""" + GPUToolbox.signal_completion(payload::Ptr{Cvoid}) + +Signal that the operation a `subscribe` function (see [`cooperative_wait`](@ref)) registered +a notification for has completed, waking up the waiting task. This can be called from a +thread that is not managed by Julia (e.g., one owned by the driver), and does not throw. +""" +function signal_completion(payload::Ptr{Cvoid}) + c = unsafe_pointer_to_objref(payload)::Completion + # releasing the lock runs pending finalizers. that should not happen here, as the driver + # may be holding locks a finalizer needs, so defer them to a thread managed by Julia. + # (debug builds of Julia do run them when finalizers are re-enabled below.) + ccall(:jl_gc_disable_finalizers_internal, Cvoid, ()) + lock(c.cond) + @atomic c.state = SIGNALLED + notify(c.cond) + unlock(c.cond) + ccall(:jl_gc_enable_finalizers_internal, Cvoid, ()) + # this is the last access, after which the completion may be freed + @atomic :release c.state = RELEASED + return +end + +# register a completion notification, or return the one registered already when resuming +# an interrupted wait +function subscribe!(subscribe::F, obj, state) where {F} + registered = state.completion + registered === nothing || return registered + + # reclaim completions that were abandoned by earlier waits + if !isempty(abandoned_completions) + @lock abandoned_completions_lock sweep_abandoned_completions() + end + + # registering the notification and keeping track of it must not be interrupted, or a + # retry would register it again, or we would lose track of it. registration may invoke + # the callback immediately, which is fine. + c = Completion() + uninterruptible() do + GC.@preserve c subscribe(obj, pointer_from_objref(c)) + state.completion = c + end + return c +end + +# returns the completion once it has been signalled +function wait_completion(subscribe::F, obj, state) where {F} + try + c = subscribe!(subscribe, obj, state) + @lock c.cond while (@atomic c.state) == PENDING + Base.wait(c.cond) + end + return c + catch + # when unwinding (e.g., because of an interrupt), the callback may still use the + # completion, so keep it alive until then. an interrupted wait that is resumed + # waits for the same completion, which is fine. + registered = state.completion + if registered !== nothing && (@atomic :acquire registered.state) != RELEASED + uninterruptible(() -> abandon_completion(registered)) + end + rethrow() + end +end + + ## entry point mutable struct WaitState # the request submitted to a worker, if any request::Union{Nothing,WaitRequest} + # the completion notification registered, if any + completion::Union{Nothing,Completion} # the first interrupt, to be thrown once the wait has completed interrupt::Union{Nothing,InterruptException} end # returns the completed request, or `nothing` if polling found `obj` to have completed -function wait_cooperatively(wait, obj, isdone, spin, state) +function wait_cooperatively(wait::W, obj, isdone::D, state) where {W,D} # when resuming an interrupted wait, keep waiting for the submitted request request = state.request if request === nothing - if isdone !== nothing && spin && spin_until(isdone, obj) - return nothing - end - request = WaitRequest(wait, obj) while !submit!(state, request, isdone === nothing) # all workers are busy: poll the object, periodically checking for a worker @@ -201,62 +318,135 @@ else end # run `f` without it being interrupted or cancelled -uninterruptible(f) = shielded(() -> disable_sigint(f)) +uninterruptible(f::F) where {F} = shielded(() -> disable_sigint(f)) """ - cooperative_wait(wait, obj; isdone=nothing, spin=true, cancellable=false) + cooperative_wait(wait, obj; subscribe=nothing, isdone=nothing, spin=true, + cancellable=false) Wait for `obj` to complete without blocking the calling thread, so that other tasks can run on it in the meantime. -`wait(obj)` should perform a blocking wait for `obj`. It is executed on a separate thread, -so it should block in a GC-safe manner (e.g., using [`@gcsafe_ccall`](@ref)) and set up any -thread-local state it relies on (e.g., the active context). +`wait(obj)` should perform a blocking wait for `obj`. By default, it is executed on a +separate thread, so it should block in a GC-safe manner (e.g., using +[`@gcsafe_ccall`](@ref)) and set up any thread-local state it relies on (e.g., the active +context). + +Alternatively, if the driver can notify completion by calling back, pass a function +`subscribe(obj, payload::Ptr{Cvoid})` that registers such a callback, which in turn calls +[`GPUToolbox.signal_completion(payload)`](@ref GPUToolbox.signal_completion). This avoids +involving a worker thread, which matters for devices that execute on the host's CPU cores, +where waking a worker delays the operation. `wait` is then only used where the calling +thread cannot switch tasks (see below). The callback: + +- has to be called exactly once, also when the operation fails, and may be called before + `subscribe` returns; +- must not be registered if `subscribe` throws; +- may be called from any thread, but must not let exceptions escape into the driver. + +`signal_completion` does not run finalizers on the calling thread, except on debug builds +of Julia, so with those, drivers that hold locks while invoking callbacks may deadlock if a +finalizer calls into the driver. Don't use `subscribe` with such drivers there. + +Registering the callback, and any other driver call that may wait for callbacks to finish, +should be GC-safe (e.g., using [`@gcsafe_ccall`](@ref)): a thread calling back into Julia +may have to wait for the garbage collector, which waits for all threads executing Julia +code. For example, with OpenCL: + +```julia +function notify_completion(event::cl_event, status::Cint, payload::Ptr{Cvoid}) + GPUToolbox.signal_completion(payload) + return +end +subscribe(event, payload) = + clSetEventCallback(event, CL_COMPLETE, + @cfunction(notify_completion, Cvoid, (cl_event, Cint, Ptr{Cvoid})), + payload) +``` -`isdone(obj)`, if given, should return whether `obj` has completed without blocking. It is -used to detect short operations without involving another thread (unless `spin=false`), -and to poll `obj` while all of these threads are busy. Objects that cannot be polled wait -for a thread to become available instead. +If the wait can be interrupted (see `cancellable`), the callback may be called after +`cooperative_wait` has returned, so it should not use anything else the caller may release +by then. -Returns `Some(wait(obj))`, or `nothing` if `wait` was not called because polling found `obj` -to have completed. In that case, calling `wait(obj)` returns without blocking, which may -still be needed, e.g., to synchronize memory or to check for errors. Errors thrown by `wait` -or `isdone` are rethrown. +`isdone(obj)`, if given, should return whether `obj` has completed without blocking. It is +used to detect short operations without involving another thread, and to poll `obj` while +all worker threads are busy. Objects that cannot be polled wait for a thread to become +available instead. + +`spin` determines how `obj` is polled before waiting as described above. By default, it is +polled for a while, first busy-waiting and then yielding to other tasks. For devices that +execute on the host's CPU cores, that competes with the operation for those cores, so pass +a duration in seconds (e.g., `spin=10e-6`) to only busy-wait for at most that long. With +`spin=false`, `obj` is not polled at all. + +Returns `Some(wait(obj))`, or `nothing` if `wait` was not called because polling or a +notification found `obj` to have completed. In that case, calling `wait(obj)` should not +block for long (it may have to wait for the notification callback to return), and may +still be needed, e.g., to synchronize memory or to check for errors. Errors thrown by +`wait`, `subscribe` or `isdone` are rethrown. If the wait is interrupted (i.e., an `InterruptException` is thrown) or the task is cancelled, the default is to keep waiting, and only throw once `obj` has completed: the -operation may be using memory that the caller would release when unwinding. Note that `wait(obj)` may not -have been called by then. Waits with `cancellable=true` throw immediately, while the -operation may still be executing. +operation may be using memory that the caller would release when unwinding. Note that +`wait(obj)` may not have been called by then. Waits with `cancellable=true` throw +immediately, while the operation may still be executing. In finalizers, where it is not possible to switch tasks, and while generating output (e.g., during precompilation), `wait(obj)` is called on the calling thread instead. """ -function cooperative_wait(wait::W, obj; isdone::D=nothing, spin::Bool=true, - cancellable::Bool=false) where {W,D} +function cooperative_wait(wait::W, obj; subscribe::S=nothing, isdone::D=nothing, + spin::Union{Bool,Real}=true, + cancellable::Bool=false) where {W,S,D} + spin isa Bool || spin >= 0 || throw(ArgumentError("spin duration must be non-negative")) if GC.in_finalizer() || generating_output() return Some(wait(obj)) end - state = WaitState(nothing, nothing) - request = if cancellable - wait_cooperatively(wait, obj, isdone, spin, state) + # fast path: poll the object, without allocating the state needed for the slow path. + # like the slow path (see `wait_uninterrupted`), polling is shielded from cancellation, + # and an interrupt is only thrown once the operation has completed, unless cancellable. + interrupt = nothing + if isdone !== nothing && spin !== false + poll = if spin === true + () -> spin_until(isdone, obj) + else + budget = round(UInt64, spin * 1e9) + () -> spin_until(isdone, obj, budget) + end + done = try + cancellable ? poll() : shielded(poll) + catch err + (cancellable || !(err isa InterruptException)) && rethrow() + interrupt = err + false + end + if done + cancellable || check_cancelled() + return nothing + end + end + + # slow path: wait for a completion notification, or hand the wait to a worker thread + state = WaitState(nothing, nothing, interrupt) + slow_wait = if subscribe === nothing + () -> wait_cooperatively(wait, obj, isdone, state) else - wait_uninterrupted(wait, obj, isdone, spin, state) + () -> wait_completion(subscribe, obj, state) end - request === nothing && return nothing - request.failed && throw(request.result) - return Some(request.result) + ret = cancellable ? slow_wait() : wait_uninterrupted(slow_wait, state) + ret isa WaitRequest || return nothing + ret.failed && throw(ret.result) + return Some(ret.result) end # an interrupt can be delivered wherever we yield or hit a safepoint. keep waiting, resuming # from where we were, and throw the interrupt once done. cancellation is deferred by running # in a shielded scope. -function wait_uninterrupted(wait, obj, isdone, spin, state) +function wait_uninterrupted(slow_wait::F, state) where {F} ret = shielded() do while true try - return wait_cooperatively(wait, obj, isdone, spin, state) + return slow_wait() catch err err isa InterruptException || rethrow() state.interrupt === nothing && (state.interrupt = err) diff --git a/test/synchronization.jl b/test/synchronization.jl index d2925cd..1b08f34 100644 --- a/test/synchronization.jl +++ b/test/synchronization.jl @@ -26,10 +26,48 @@ end # as it runs on the current thread (i.e., it was created with `@async`). waiting(op) = timedwait(() -> @atomic(op.waited), 30) === :ok +# signal a completion after a delay, from a thread that is not managed by Julia (like a +# driver's callback thread) +struct DelayedSignal + payload::Ptr{Cvoid} + ms::Cuint +end +function delayed_signal(arg::Ptr{DelayedSignal}) + signal = unsafe_load(arg) + Libc.free(arg) + @gcsafe_ccall uv_sleep(signal.ms::Cuint)::Cvoid + GPUToolbox.signal_completion(signal.payload) + return +end +function signal_later(payload, ms) + arg = convert(Ptr{DelayedSignal}, Libc.malloc(sizeof(DelayedSignal))) + unsafe_store!(arg, DelayedSignal(payload, ms)) + tid = Ref{NTuple{32, UInt8}}(ntuple(_ -> 0x0, 32)) + cb = @cfunction(delayed_signal, Cvoid, (Ptr{DelayedSignal},)) + err = ccall(:uv_thread_create, Cint, (Ptr{Cvoid}, Ptr{Cvoid}, Ptr{Cvoid}), tid, cb, arg) + @test err == 0 + ccall(:uv_thread_detach, Cint, (Ptr{Cvoid},), tid) + return +end + @testset "cooperative_wait" begin # completed operations are detected by polling @test cooperative_wait(blocking_wait, complete!(Operation()); isdone) === nothing + # polling can be limited to busy-waiting for some time + @test cooperative_wait(blocking_wait, complete!(Operation()); isdone, spin=1e-3) === nothing + @test cooperative_wait(blocking_wait, complete_after!(Operation(), 0.1); isdone, + spin=1e-6) == Some(:waited) + @test_throws ArgumentError cooperative_wait(blocking_wait, Operation(); spin=-1) + + # without allocating (other than to shield from task cancellation, on Julia 1.14+) + if !isdefined(Base, :CANCEL_TOKEN) + op = complete!(Operation()) + fast_wait(op) = cooperative_wait(blocking_wait, op; isdone) + fast_wait(op) + @test @allocated(fast_wait(op)) == 0 + end + # other tasks on this thread keep running while waiting on a worker for polled in (false, true) op = complete_after!(Operation(), 30) # in case the thread is blocked @@ -100,6 +138,31 @@ waiting(op) = timedwait(() -> @atomic(op.waited), 30) === :ok @test t.exception isa InterruptException end + # the same goes for interrupts while polling + for cancellable in (false, true) + op = Operation() + interrupted = Ref(false) + function interrupting_isdone(op) + interrupted[] || (interrupted[] = true; throw(InterruptException())) + return isdone(op) + end + t = @async cooperative_wait(blocking_wait, op; cancellable, + isdone=interrupting_isdone) + try + if cancellable + @test timedwait(() -> istaskdone(t), 30) === :ok + @test !isdone(op) + else + @test waiting(op) + @test !istaskdone(t) + end + finally + complete!(op) + end + @test_throws TaskFailedException wait(t) + @test t.exception isa InterruptException + end + # the same goes for task cancellation if isdefined(Base, :CancellationTokenSource) src = Base.CancellationTokenSource() @@ -118,6 +181,103 @@ waiting(op) = timedwait(() -> @atomic(op.waited), 30) === :ok complete!(op) end @test_throws TaskFailedException wait(t) + + # also when polling finds the operation to have completed + src = Base.CancellationTokenSource() + Base.cancel!(src) + op = complete!(Operation()) + t = Base.ScopedValues.with(Base.CANCEL_TOKEN => Base.CancellationToken(src)) do + @async cooperative_wait(blocking_wait, op; isdone) + end + @test_throws TaskFailedException wait(t) + @test t.exception isa Base.CancellationRequest + end + + # completion notifications from the driver, instead of a worker + let + # the callback may be invoked before registration returns + subscribed = Ref(0) + immediate = function (op, payload) + subscribed[] += 1 + GPUToolbox.signal_completion(payload) + end + @test cooperative_wait(blocking_wait, Operation(); subscribe=immediate) === nothing + @test subscribed[] == 1 + + # or later, from another thread, while other tasks on this thread keep running + ticks = Ref(0) + ticker = @async while ticks[] >= 0 + ticks[] += 1 + sleep(0.001) + end + op = Operation() + @test cooperative_wait(blocking_wait, op; subscribe=(_, p) -> signal_later(p, 100), + isdone, spin=false) === nothing + @test !@atomic(op.waited) + @test ticks[] > 0 + ticks[] = -1 + wait(ticker) + + # errors while registering are rethrown + failing = (_, _) -> error("oops") + @test_throws ErrorException("oops") cooperative_wait(blocking_wait, Operation(); + subscribe=failing) + + # interrupted waits keep waiting until notified, unless cancellable, in which case + # a later notification is harmless + for cancellable in (false, true) + notified = Base.Event() + payload = Ref{Ptr{Cvoid}}(C_NULL) + subscribe = (_, p) -> (payload[] = p; notify(notified)) + t = @async cooperative_wait(blocking_wait, Operation(); subscribe, cancellable) + wait(notified) + yield() + schedule(t, InterruptException(); error=true) + if cancellable + @test timedwait(() -> istaskdone(t), 30) === :ok + GC.gc() + signal_later(payload[], 10) + sleep(0.1) + else + for _ in 1:100 + yield() + end + @test !istaskdone(t) + signal_later(payload[], 10) + end + @test_throws TaskFailedException wait(t) + @test t.exception isa InterruptException + end + + # the same goes for task cancellation + if isdefined(Base, :CancellationTokenSource) + for cancellable in (false, true) + src = Base.CancellationTokenSource() + notified = Base.Event() + payload = Ref{Ptr{Cvoid}}(C_NULL) + subscribe = (_, p) -> (payload[] = p; notify(notified)) + token = Base.CancellationToken(src) + t = Base.ScopedValues.with(Base.CANCEL_TOKEN => token) do + @async cooperative_wait(blocking_wait, Operation(); subscribe, + cancellable) + end + wait(notified) + Base.cancel!(src) + if cancellable + @test timedwait(() -> istaskdone(t), 30) === :ok + GC.gc() + signal_later(payload[], 10) + sleep(0.1) + else + for _ in 1:100 + yield() + end + @test !istaskdone(t) + signal_later(payload[], 10) + end + @test_throws TaskFailedException wait(t) + end + end end # finalizers cannot switch tasks, so they wait on the calling thread