From 5e185e7e084be0d1d596ca8c27553225674056a6 Mon Sep 17 00:00:00 2001 From: JamesWrigley Date: Fri, 10 Jul 2026 13:58:37 +0200 Subject: [PATCH] Run FFTW in GC-safe regions Should help prevent GC pressure. --- src/fft.jl | 63 +++++++++++++++++++++++++++++++++--------------- src/providers.jl | 8 +++++- 2 files changed, 50 insertions(+), 21 deletions(-) diff --git a/src/fft.jl b/src/fft.jl index eb62f61..48417a3 100644 --- a/src/fft.jl +++ b/src/fft.jl @@ -192,6 +192,11 @@ function get_num_threads() end end +# Whether a plan created now with the current planner settings will re-enter +# Julia through the `spawnloop` threading callback when executed. Only needed +# with FFTW since MKL handles its own threading. +_plan_is_threaded() = fftw_provider == "fftw" && get_num_threads() > 1 + @exclusive function set_num_threads(f::Function, num_threads::Integer) orig_num_threads = get_num_threads() _set_num_threads(num_threads) @@ -261,11 +266,12 @@ for P in (:cFFTWPlan, :rFFTWPlan) # complex, r2c/c2r oalign::Int32 # alignment mod 16 of input flags::UInt32 # planner flags region::G # region (iterable) of dims that are transformed + threaded::Bool # whether execution re-enters Julia via the spawnloop callback pinv::ScaledPlan function $P{T,K,inplace,N,G}(plan::PlanPtr, flags::Integer, R::G, X::StridedArray{T,N}, Y::StridedArray) where {T<:fftwNumber,K,inplace,N,G} p = new(plan, size(X), size(Y), strides(X), strides(Y), - alignment_of(X), alignment_of(Y), flags, R) + alignment_of(X), alignment_of(Y), flags, R, _plan_is_threaded()) finalizer(maybe_destroy_plan, p) p end @@ -290,12 +296,13 @@ mutable struct r2rFFTWPlan{T<:fftwNumber,K,inplace,N,G} <: FFTWPlan{T,K,inplace} flags::UInt32 # planner flags region::G # region (iterable) of dims that are transformed kinds::K + threaded::Bool # whether execution re-enters Julia via the spawnloop callback pinv::ScaledPlan function r2rFFTWPlan{T,K,inplace,N,G}(plan::PlanPtr, flags::Integer, R::G, X::StridedArray{T,N}, Y::StridedArray, kinds::K) where {T<:fftwNumber,K,inplace,N,G} p = new(plan, size(X), size(Y), strides(X), strides(Y), - alignment_of(X), alignment_of(Y), flags, R, kinds) + alignment_of(X), alignment_of(Y), flags, R, kinds, _plan_is_threaded()) finalizer(maybe_destroy_plan, p) p end @@ -511,51 +518,67 @@ _colmajorstrides(p) = () # Execute +macro execute(plan, ccallexpr) + quote + if $(esc(plan)).threaded + # If the plan is multi-threaded then we execute it directly. It will + # end up calling spawnloop(), which will enter a GC-safe region. + $(esc(ccallexpr)) + else + # If the plan is single-threaded then it won't call back into Julia + # and we can execute it in a GC-safe region. + gc_state = @ccall jl_gc_safe_enter()::Int8 + $(esc(ccallexpr)) + @ccall jl_gc_safe_leave(gc_state::Int8)::Cvoid + end + end +end + unsafe_execute!(plan::FFTWPlan{<:fftwDouble}) = - ccall((:fftw_execute,libfftw3), Cvoid, (PlanPtr,), plan) + @execute plan ccall((:fftw_execute,libfftw3), Cvoid, (PlanPtr,), plan) unsafe_execute!(plan::FFTWPlan{<:fftwSingle}) = - ccall((:fftwf_execute,libfftw3f), Cvoid, (PlanPtr,), plan) + @execute plan ccall((:fftwf_execute,libfftw3f), Cvoid, (PlanPtr,), plan) unsafe_execute!(plan::cFFTWPlan{T}, X::StridedArray{T}, Y::StridedArray{T}) where {T<:fftwDouble} = - ccall((:fftw_execute_dft,libfftw3), Cvoid, - (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) + @execute plan ccall((:fftw_execute_dft,libfftw3), Cvoid, + (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) unsafe_execute!(plan::cFFTWPlan{T}, X::StridedArray{T}, Y::StridedArray{T}) where {T<:fftwSingle} = - ccall((:fftwf_execute_dft,libfftw3f), Cvoid, - (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) + @execute plan ccall((:fftwf_execute_dft,libfftw3f), Cvoid, + (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) unsafe_execute!(plan::rFFTWPlan{Float64,FORWARD}, X::StridedArray{Float64}, Y::StridedArray{Complex{Float64}}) = - ccall((:fftw_execute_dft_r2c,libfftw3), Cvoid, - (PlanPtr,Ptr{Float64},Ptr{Complex{Float64}}), plan, X, Y) + @execute plan ccall((:fftw_execute_dft_r2c,libfftw3), Cvoid, + (PlanPtr,Ptr{Float64},Ptr{Complex{Float64}}), plan, X, Y) unsafe_execute!(plan::rFFTWPlan{Float32,FORWARD}, X::StridedArray{Float32}, Y::StridedArray{Complex{Float32}}) = - ccall((:fftwf_execute_dft_r2c,libfftw3f), Cvoid, - (PlanPtr,Ptr{Float32},Ptr{Complex{Float32}}), plan, X, Y) + @execute plan ccall((:fftwf_execute_dft_r2c,libfftw3f), Cvoid, + (PlanPtr,Ptr{Float32},Ptr{Complex{Float32}}), plan, X, Y) unsafe_execute!(plan::rFFTWPlan{Complex{Float64},BACKWARD}, X::StridedArray{Complex{Float64}}, Y::StridedArray{Float64}) = - ccall((:fftw_execute_dft_c2r,libfftw3), Cvoid, - (PlanPtr,Ptr{Complex{Float64}},Ptr{Float64}), plan, X, Y) + @execute plan ccall((:fftw_execute_dft_c2r,libfftw3), Cvoid, + (PlanPtr,Ptr{Complex{Float64}},Ptr{Float64}), plan, X, Y) unsafe_execute!(plan::rFFTWPlan{Complex{Float32},BACKWARD}, X::StridedArray{Complex{Float32}}, Y::StridedArray{Float32}) = - ccall((:fftwf_execute_dft_c2r,libfftw3f), Cvoid, - (PlanPtr,Ptr{Complex{Float32}},Ptr{Float32}), plan, X, Y) + @execute plan ccall((:fftwf_execute_dft_c2r,libfftw3f), Cvoid, + (PlanPtr,Ptr{Complex{Float32}},Ptr{Float32}), plan, X, Y) unsafe_execute!(plan::r2rFFTWPlan{T}, X::StridedArray{T}, Y::StridedArray{T}) where {T<:fftwDouble} = - ccall((:fftw_execute_r2r,libfftw3), Cvoid, - (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) + @execute plan ccall((:fftw_execute_r2r,libfftw3), Cvoid, + (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) unsafe_execute!(plan::r2rFFTWPlan{T}, X::StridedArray{T}, Y::StridedArray{T}) where {T<:fftwSingle} = - ccall((:fftwf_execute_r2r,libfftw3f), Cvoid, - (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) + @execute plan ccall((:fftwf_execute_r2r,libfftw3f), Cvoid, + (PlanPtr,Ptr{T},Ptr{T}), plan, X, Y) # NOTE ON GC (garbage collection): # The FFTWPlan has a finalizer so that gc will destroy the plan, diff --git a/src/providers.jl b/src/providers.jl index 6764fd2..a5c3d35 100644 --- a/src/providers.jl +++ b/src/providers.jl @@ -46,7 +46,13 @@ end # tasks (FFTW/fftw3#175): function spawnloop(f::Ptr{Cvoid}, fdata::Ptr{Cvoid}, elsize::Csize_t, num::Cint, callback_data::Ptr{Cvoid}) @sync for i = 0:num-1 - Threads.@spawn ccall(f, Ptr{Cvoid}, (Ptr{Cvoid},), fdata + elsize*i) + Threads.@spawn begin + # Run the FFTW work chunk in a GC-safe region so the garbage + # collector can make progress while this thread is blocked. + gc_state = @ccall jl_gc_safe_enter()::Int8 + ccall(f, Ptr{Cvoid}, (Ptr{Cvoid},), fdata + elsize*i) + @ccall jl_gc_safe_leave(gc_state::Int8)::Cvoid + end end end