Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 21 additions & 17 deletions src/apiutils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand All @@ -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
Expand Down
4 changes: 2 additions & 2 deletions src/derivative.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
6 changes: 3 additions & 3 deletions src/gradient.jl
Original file line number Diff line number Diff line change
Expand Up @@ -127,22 +127,22 @@ 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
i = ((c - 1) * N + 1)
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
Expand Down
16 changes: 8 additions & 8 deletions src/jacobian.jl
Original file line number Diff line number Diff line change
Expand Up @@ -191,15 +191,15 @@ 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
i = ((c - 1) * N + 1)
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
Expand All @@ -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)),
Expand All @@ -226,17 +226,17 @@ 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

@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)),
:(),
Expand All @@ -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
13 changes: 7 additions & 6 deletions test/AllocationsTest.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading