From 4d555f1aee3e1dbb67c2e9091ac1cc440f78da17 Mon Sep 17 00:00:00 2001 From: Tim Besard Date: Wed, 30 Sep 2026 23:57:17 +0200 Subject: [PATCH 1/3] Don't allocate when polling finds an operation completed `cooperative_wait` allocated its wait state before polling, even though that state is only needed when handing the wait to a worker. Poll first, keeping the same guarantees: polling is shielded from cancellation, and for waits that cannot be interrupted, an interrupt is only thrown once the operation has completed, or once the worker's wait has, if polling gets interrupted. Also specialize on the functions that are passed through, which Julia otherwise doesn't, causing dynamic dispatch and boxing. --- src/synchronization.jl | 47 +++++++++++++++++++++++++++++------------ test/synchronization.jl | 43 +++++++++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 14 deletions(-) diff --git a/src/synchronization.jl b/src/synchronization.jl index 98d0140..13a0c7c 100644 --- a/src/synchronization.jl +++ b/src/synchronization.jl @@ -167,14 +167,10 @@ mutable struct WaitState 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,7 +197,7 @@ 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) @@ -225,9 +221,9 @@ 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. @@ -238,11 +234,34 @@ function cooperative_wait(wait::W, obj; isdone::D=nothing, spin::Bool=true, return Some(wait(obj)) end - state = WaitState(nothing, nothing) + # 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 + done = try + if cancellable + spin_until(isdone, obj) + else + shielded(() -> spin_until(isdone, obj)) + end + catch err + (cancellable || !(err isa InterruptException)) && rethrow() + interrupt = err + false + end + if done + cancellable || check_cancelled() + return nothing + end + end + + # slow path: hand the wait to a worker thread + state = WaitState(nothing, interrupt) request = if cancellable - wait_cooperatively(wait, obj, isdone, spin, state) + wait_cooperatively(wait, obj, isdone, state) else - wait_uninterrupted(wait, obj, isdone, spin, state) + wait_uninterrupted(wait, obj, isdone, state) end request === nothing && return nothing request.failed && throw(request.result) @@ -252,11 +271,11 @@ 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(wait::W, obj, isdone::D, state) where {W,D} ret = shielded() do while true try - return wait_cooperatively(wait, obj, isdone, spin, state) + return wait_cooperatively(wait, obj, isdone, state) catch err err isa InterruptException || rethrow() state.interrupt === nothing && (state.interrupt = err) diff --git a/test/synchronization.jl b/test/synchronization.jl index d2925cd..c9d463e 100644 --- a/test/synchronization.jl +++ b/test/synchronization.jl @@ -30,6 +30,14 @@ waiting(op) = timedwait(() -> @atomic(op.waited), 30) === :ok # completed operations are detected by polling @test cooperative_wait(blocking_wait, complete!(Operation()); isdone) === nothing + # 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 +108,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 +151,16 @@ 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 # finalizers cannot switch tasks, so they wait on the calling thread From 9a671e36ebd98dc46b8573214dca45797aa2f1c3 Mon Sep 17 00:00:00 2001 From: Tim Besard Date: Wed, 30 Sep 2026 23:58:16 +0200 Subject: [PATCH 2/3] Support waiting for completion notifications from the driver Waiting on a worker thread works well for GPUs, but not for devices that execute on the host's CPU cores (e.g., PoCL, or Intel's CPU runtime): waking the worker as the operation starts, and polling, compete with the operation for those cores. Measured on 4 cores with PoCL, kernels of 0.2-1 ms took 25-35% longer than with a plain blocking wait. For those devices, `cooperative_wait` can now have the driver notify it when the operation completes: `subscribe(obj, payload)` registers a driver callback, which calls `GPUToolbox.signal_completion(payload)`. That directly wakes the waiting task (going through the event loop instead would add a thread hop), so no worker thread is involved, and kernels take as long as with a blocking wait. Signalling only takes a spin lock, and defers finalizers to a thread managed by Julia. Registration cannot be interrupted halfway, and a completion that its waiter gave up on (e.g., because it was cancelled) is kept alive until the driver has signalled it. --- Project.toml | 2 +- src/synchronization.jl | 198 +++++++++++++++++++++++++++++++++++----- test/synchronization.jl | 111 ++++++++++++++++++++++ 3 files changed, 287 insertions(+), 24 deletions(-) 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 13a0c7c..6133422 100644 --- a/src/synchronization.jl +++ b/src/synchronization.jl @@ -157,11 +157,119 @@ 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 @@ -200,24 +308,67 @@ end 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) +``` + +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. `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. +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. + +With `spin=false`, `obj` is not polled before waiting as described above. For devices that +execute on the host's CPU cores, that is recommended: polling competes with the operation +for those cores. -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. +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 @@ -228,8 +379,8 @@ 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::Bool=true, cancellable::Bool=false) where {W,S,D} if GC.in_finalizer() || generating_output() return Some(wait(obj)) end @@ -256,26 +407,27 @@ function cooperative_wait(wait::W, obj; isdone::D=nothing, spin::Bool=true, end end - # slow path: hand the wait to a worker thread - state = WaitState(nothing, interrupt) - request = if cancellable - wait_cooperatively(wait, obj, isdone, state) + # 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, 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::W, obj, isdone::D, state) where {W,D} +function wait_uninterrupted(slow_wait::F, state) where {F} ret = shielded() do while true try - return wait_cooperatively(wait, obj, isdone, 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 c9d463e..e6a6fe7 100644 --- a/test/synchronization.jl +++ b/test/synchronization.jl @@ -26,6 +26,30 @@ 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 @@ -163,6 +187,93 @@ waiting(op) = timedwait(() -> @atomic(op.waited), 30) === :ok @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 ret = Ref{Any}() @noinline function finalized_object() From 3767a5d024aa6d39fb08c49cc1ef21a0f9d44a78 Mon Sep 17 00:00:00 2001 From: Tim Besard Date: Thu, 1 Oct 2026 00:12:30 +0200 Subject: [PATCH 3/3] Support polling for a limited time MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit By default, `cooperative_wait` polls for a while, busy-waiting and then yielding to other tasks, which takes around 100 µs. For devices that execute on the host's CPU cores, that competes with the operation for those cores. Not polling at all makes short operations slower, though: with KernelAbstractions' PoCL back-end, trivial launches took 6.3 µs instead of 4.0 µs with a blocking wait. So also accept a duration for `spin`, to only busy-wait for at most that long. Polling for 10 µs before waiting for a completion notification, the saxpy benchmark of KernelAbstractions on 4 cores takes 3.1 instead of 4.0 µs for 1024 elements, and 12.2 instead of 13.3 µs for 262144. Longer operations are delayed by about half the polling time, which is why it should be kept short. --- src/synchronization.jl | 39 +++++++++++++++++++++++++++++---------- test/synchronization.jl | 6 ++++++ 2 files changed, 35 insertions(+), 10 deletions(-) diff --git a/src/synchronization.jl b/src/synchronization.jl index 6133422..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 @@ -360,9 +373,11 @@ used to detect short operations without involving another thread, and to poll `o all worker threads are busy. Objects that cannot be polled wait for a thread to become available instead. -With `spin=false`, `obj` is not polled before waiting as described above. For devices that -execute on the host's CPU cores, that is recommended: polling competes with the operation -for those cores. +`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 @@ -380,7 +395,9 @@ In finalizers, where it is not possible to switch tasks, and while generating ou during precompilation), `wait(obj)` is called on the calling thread instead. """ function cooperative_wait(wait::W, obj; subscribe::S=nothing, isdone::D=nothing, - spin::Bool=true, cancellable::Bool=false) where {W,S,D} + 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 @@ -389,13 +406,15 @@ function cooperative_wait(wait::W, obj; subscribe::S=nothing, isdone::D=nothing, # 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 + 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 - if cancellable - spin_until(isdone, obj) - else - shielded(() -> spin_until(isdone, obj)) - end + cancellable ? poll() : shielded(poll) catch err (cancellable || !(err isa InterruptException)) && rethrow() interrupt = err diff --git a/test/synchronization.jl b/test/synchronization.jl index e6a6fe7..1b08f34 100644 --- a/test/synchronization.jl +++ b/test/synchronization.jl @@ -54,6 +54,12 @@ end # 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())