diff --git a/src/apiutils.jl b/src/apiutils.jl index f401a3fc..4617c88e 100644 --- a/src/apiutils.jl +++ b/src/apiutils.jl @@ -27,7 +27,7 @@ end function vector_mode_dual_eval!(f!::F, cfg::JacobianConfig, y, x) where {F} ydual, xdual = cfg.duals seed!(xdual, x, cfg.seeds) - seed!(ydual, y) + unseed!(ydual, y) f!(ydual, xdual) return ydual end @@ -70,8 +70,10 @@ function structural_eachindex(x::Diagonal, y::AbstractArray) return diagind(x) end -function seed!(duals::AbstractArray{Dual{T,V,N}}, x, - seed::Partials{N,V} = zero(Partials{N,V})) where {T,V,N} +# Copies the values of `x` into `duals` with zero partials, i.e. removes any seeds +# `duals` is currently carrying. +function unseed!(duals::AbstractArray{Dual{T,V,N}}, x) where {T,V,N} + seed = zero(Partials{N,V}) if isbitstype(V) for idx in structural_eachindex(duals, x) duals[idx] = Dual{T,V,N}(x[idx], seed) @@ -88,16 +90,20 @@ function seed!(duals::AbstractArray{Dual{T,V,N}}, x, return duals end -function seed!(duals::AbstractArray{Dual{T,V,N}}, x, - seeds::NTuple{N,Partials{N,V}}) where {T,V,N} +# Unseeds at most `N` elements starting at `index`: chunk mode only ever needs to clear +# the N-wide chunk it just seeded, so writing through to the end of the array would be +# O(n) redundant work per chunk (O(n^2) per sweep). +function unseed!(duals::AbstractArray{Dual{T,V,N}}, x, index) where {T,V,N} + seed = zero(Partials{N,V}) + idxs = Iterators.take(Iterators.drop(structural_eachindex(duals, x), index - 1), N) if isbitstype(V) - for (i, idx) in zip(1:N, structural_eachindex(duals, x)) - duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) + for idx in idxs + duals[idx] = Dual{T,V,N}(x[idx], seed) end else - for (i, idx) in zip(1:N, structural_eachindex(duals, x)) + for idx in idxs if isassigned(x, idx) - duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) + duals[idx] = Dual{T,V,N}(x[idx], seed) else Base._unsetindex!(duals, idx) end @@ -106,18 +112,16 @@ function seed!(duals::AbstractArray{Dual{T,V,N}}, x, return duals end -function seed!(duals::AbstractArray{Dual{T,V,N}}, x, index, - seed::Partials{N,V} = zero(Partials{N,V})) where {T,V,N} - offset = index - 1 - idxs = Iterators.drop(structural_eachindex(duals, x), offset) +function seed!(duals::AbstractArray{Dual{T,V,N}}, x, + seeds::NTuple{N,Partials{N,V}}) where {T,V,N} if isbitstype(V) - for idx in idxs - duals[idx] = Dual{T,V,N}(x[idx], seed) + for (i, idx) in zip(1:N, structural_eachindex(duals, x)) + duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) end else - for idx in idxs + for (i, idx) in zip(1:N, structural_eachindex(duals, x)) if isassigned(x, idx) - duals[idx] = Dual{T,V,N}(x[idx], seed) + duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) else Base._unsetindex!(duals, idx) end diff --git a/src/derivative.jl b/src/derivative.jl index b39e2a48..d9fb355a 100644 --- a/src/derivative.jl +++ b/src/derivative.jl @@ -27,7 +27,7 @@ Set `check` to `Val{false}()` to disable tag checking. This can lead to perturba require_one_based_indexing(y) CHK && checktag(T, f!, x) ydual = cfg.duals - seed!(ydual, y) + unseed!(ydual, y) f!(ydual, Dual{T}(x, one(x))) map!(value, y, ydual) return extract_derivative(T, ydual) @@ -65,7 +65,7 @@ Set `check` to `Val{false}()` to disable tag checking. This can lead to perturba result isa DiffResult ? require_one_based_indexing(y) : require_one_based_indexing(result, y) CHK && checktag(T, f!, x) ydual = cfg.duals - seed!(ydual, y) + unseed!(ydual, y) f!(ydual, Dual{T}(x, one(x))) result = extract_value!(T, result, y, ydual) result = extract_derivative!(T, result, ydual) diff --git a/src/gradient.jl b/src/gradient.jl index 0832d354..d76a476e 100644 --- a/src/gradient.jl +++ b/src/gradient.jl @@ -127,14 +127,14 @@ function chunk_mode_gradient_expr(result_definition::Expr) # seed work vectors xdual = cfg.duals seeds = cfg.seeds - seed!(xdual, x) + unseed!(xdual, x) # do first chunk manually to calculate output type seed!(xdual, x, 1, seeds) ydual = f(xdual) $(result_definition) extract_gradient_chunk!(T, result, ydual, 1, N) - seed!(xdual, x, 1) + unseed!(xdual, x, 1) # do middle chunks for c in middlechunks @@ -142,7 +142,7 @@ function chunk_mode_gradient_expr(result_definition::Expr) seed!(xdual, x, i, seeds) ydual = f(xdual) extract_gradient_chunk!(T, result, ydual, i, N) - seed!(xdual, x, i) + unseed!(xdual, x, i) end # do final chunk diff --git a/src/jacobian.jl b/src/jacobian.jl index b8ce58fb..5f3a79fb 100644 --- a/src/jacobian.jl +++ b/src/jacobian.jl @@ -191,7 +191,7 @@ function jacobian_chunk_mode_expr(work_array_definition::Expr, compute_ydual::Ex $(result_definition) out_reshaped = reshape_jacobian(result, ydual, xdual) extract_jacobian_chunk!(T, out_reshaped, ydual, 1, N) - seed!(xdual, x, 1) + unseed!(xdual, x, 1) # do middle chunks for c in middlechunks @@ -199,7 +199,7 @@ function jacobian_chunk_mode_expr(work_array_definition::Expr, compute_ydual::Ex seed!(xdual, x, i, seeds) $(compute_ydual) extract_jacobian_chunk!(T, out_reshaped, ydual, i, N) - seed!(xdual, x, i) + unseed!(xdual, x, i) end # do final chunk @@ -216,7 +216,7 @@ end @eval function chunk_mode_jacobian(f::F, x, cfg::JacobianConfig{T,V,N}) where {F,T,V,N} $(jacobian_chunk_mode_expr(quote xdual = cfg.duals - seed!(xdual, x) + unseed!(xdual, x) end, :(ydual = f(xdual)), :(result = similar(ydual, valtype(T, eltype(ydual)), length(ydual), xlen)), @@ -226,9 +226,9 @@ end @eval function chunk_mode_jacobian(f!::F, y, x, cfg::JacobianConfig{T,V,N}) where {F,T,V,N} $(jacobian_chunk_mode_expr(quote ydual, xdual = cfg.duals - seed!(xdual, x) + unseed!(xdual, x) end, - :(f!(seed!(ydual, y), xdual)), + :(f!(unseed!(ydual, y), xdual)), :(result = similar(y, length(y), xlen)), :(map!(d -> value(T,d), y, ydual)))) end @@ -236,7 +236,7 @@ end @eval function chunk_mode_jacobian!(result, f::F, x, cfg::JacobianConfig{T,V,N}) where {F,T,V,N} $(jacobian_chunk_mode_expr(quote xdual = cfg.duals - seed!(xdual, x) + unseed!(xdual, x) end, :(ydual = f(xdual)), :(), @@ -246,9 +246,9 @@ end @eval function chunk_mode_jacobian!(result, f!::F, y, x, cfg::JacobianConfig{T,V,N}) where {F,T,V,N} $(jacobian_chunk_mode_expr(quote ydual, xdual = cfg.duals - seed!(xdual, x) + unseed!(xdual, x) end, - :(f!(seed!(ydual, y), xdual)), + :(f!(unseed!(ydual, y), xdual)), :(), :(extract_value!(T, result, y, ydual)))) end diff --git a/test/AllocationsTest.jl b/test/AllocationsTest.jl index af8d6e77..d75672b2 100644 --- a/test/AllocationsTest.jl +++ b/test/AllocationsTest.jl @@ -7,22 +7,23 @@ include(joinpath(dirname(@__FILE__), "utils.jl")) convert_test_574() = convert(ForwardDiff.Dual{Nothing,ForwardDiff.Dual{Nothing,ForwardDiff.Dual{Nothing,Float64,8},4},2}, 1.3) -@testset "Test seed! allocations" begin +@testset "Test seed!/unseed! allocations" begin x = rand(1000) cfg = ForwardDiff.GradientConfig(nothing, x) duals = cfg.duals seeds = cfg.seeds - seed = cfg.seeds[1] allocs_seed!(args...) = @allocated ForwardDiff.seed!(args...) allocs_seed!(duals, x, seeds) @test iszero(allocs_seed!(duals, x, seeds)) - allocs_seed!(duals, x, seed) - @test iszero(allocs_seed!(duals, x, seed)) allocs_seed!(duals, x, 1, seeds) @test iszero(allocs_seed!(duals, x, 1, seeds)) - allocs_seed!(duals, x, 1, seed) - @test iszero(allocs_seed!(duals, x, 1, seed)) + + allocs_unseed!(args...) = @allocated ForwardDiff.unseed!(args...) + allocs_unseed!(duals, x) + @test iszero(allocs_unseed!(duals, x)) + allocs_unseed!(duals, x, 1) + @test iszero(allocs_unseed!(duals, x, 1)) allocs_convert_test_574() = @allocated convert_test_574() allocs_convert_test_574()