From 3890c8c9199c0d7750d5e244aef8631763c6f16a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Wed, 23 Sep 2026 21:39:23 +0200 Subject: [PATCH 1/8] Extract values and partials irrespective of the nesting order `value(T, d)`, `partials(T, d, i)` and `valtype(T, D)` now descend through layers with other tags instead of throwing `DualMismatchError` when `T` is not the outermost tag. If `T` does not occur at all, the value is the number itself and the partials are zero, regardless of the order of tags. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/dual.jl | 44 +++++++++++++++++++++++++++++--------------- src/gradient.jl | 10 +++++++--- test/DualTest.jl | 18 ++++++++++++++---- 3 files changed, 50 insertions(+), 22 deletions(-) diff --git a/src/dual.jl b/src/dual.jl index 6d13dec3..4bd1c646 100644 --- a/src/dual.jl +++ b/src/dual.jl @@ -88,13 +88,22 @@ Dual{T,V,N}(x::Base.TwicePrecision) where {T,V,N} = @inline value(x) = x @inline value(d::Dual) = d.value +# Whether a `Dual` with tag `T` occurs anywhere in the nesting of `D` +@inline hastag(::Type{T}, ::Type) where {T} = false +@inline hastag(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} = S === T || hastag(T, V) + +@inline map_partials(f::F, ::Type{W}, p::Partials{N}) where {F,W,N} = Partials{N,W}(map(f, p.values)) + +# Extraction w.r.t. `T` descends through layers with other tags, so it does not depend on +# the order in which the layers are nested. Without a layer with tag `T` the value is the +# number itself and the partials are zero. @inline value(::Type{T}, x) where T = x @inline value(::Type{T}, d::Dual{T}) where T = value(d) -@inline function value(::Type{T}, d::Dual{S}) where {T,S} - if S ≺ T - d +@inline function value(::Type{T}, d::Dual{S,V}) where {T,S,V} + if hastag(T, typeof(d)) + Dual{S}(value(T, value(S, d)), map_partials(p -> value(T, p), valtype(T, V), partials(S, d))) else - throw(DualMismatchError(T,S)) + d end end @@ -107,18 +116,29 @@ end @inline Base.@propagate_inbounds partials(::Type{T}, x, i...) where T = partials(x, i...) @inline Base.@propagate_inbounds partials(::Type{T}, d::Dual{T}, i...) where T = partials(d, i...) -@inline function partials(::Type{T}, d::Dual{S}, i...) where {T,S} - if S ≺ T +@inline Base.@propagate_inbounds function partials(::Type{T}, d::Dual{S,V}, i...) where {T,S,V} + if hastag(T, typeof(d)) + Dual{S}(partials(T, value(S, d), i...), map_partials(p -> partials(T, p, i...), valtype(T, V), partials(S, d))) + else zero(d) + end +end +@inline partials(::Type{T}, d::Dual{T}) where {T} = partials(d) +@inline function partials(::Type{T}, d::Dual) where {T} + if hastag(T, typeof(d)) + Partials{npartials(T, typeof(d)),valtype(T, typeof(d))}(ntuple(i -> partials(T, d, i), Val(npartials(T, typeof(d))))) else - throw(DualMismatchError(T,S)) + Partials{0,typeof(d)}(tuple()) end end - @inline npartials(::Dual{T,V,N}) where {T,V,N} = N @inline npartials(::Type{Dual{T,V,N}}) where {T,V,N} = N +@inline npartials(::Type{T}, ::Type) where {T} = 0 +@inline npartials(::Type{T}, ::Type{Dual{T,V,N}}) where {T,V,N} = N +@inline npartials(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} = npartials(T, V) + @inline order(::Type{V}) where {V} = 0 @inline order(::Type{Dual{T,V,N}}) where {T,V,N} = 1 + order(V) @@ -130,13 +150,7 @@ end @inline valtype(::Type{T}, ::V) where {T,V} = valtype(T, V) @inline valtype(::Type, ::Type{V}) where {V} = V @inline valtype(::Type{T}, ::Type{Dual{T,V,N}}) where {T,V,N} = V -@inline function valtype(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} - if S ≺ T - Dual{S,V,N} - else - throw(DualMismatchError(T,S)) - end -end +@inline valtype(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} = Dual{S,valtype(T, V),N} @inline tagtype(::V) where {V} = Nothing @inline tagtype(::Type{V}) where {V} = Nothing diff --git a/src/gradient.jl b/src/gradient.jl index a5ef3dac..235ac06e 100644 --- a/src/gradient.jl +++ b/src/gradient.jl @@ -64,9 +64,13 @@ end extract_gradient!(::Type{T}, result::AbstractArray, y::Real) where {T} = fill!(result, zero(y)) function extract_gradient!(::Type{T}, result::AbstractArray, dual::Dual) where {T} - idxs = structural_eachindex(result) - for (i, idx) in zip(1:npartials(dual), idxs) - result[idx] = partials(T, dual, i) + if hastag(T, typeof(dual)) + idxs = structural_eachindex(result) + for (i, idx) in zip(1:npartials(T, typeof(dual)), idxs) + result[idx] = partials(T, dual, i) + end + else + fill!(result, zero(dual)) end return result end diff --git a/test/DualTest.jl b/test/DualTest.jl index 67d9c9f7..7a55fed8 100644 --- a/test/DualTest.jl +++ b/test/DualTest.jl @@ -112,10 +112,20 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test ForwardDiff.valtype(OuterTestTag, NESTED_FDNUM) == Dual{TestTag,Dual{TestTag,V,M},N} @test ForwardDiff.valtype(OuterTestTag, typeof(NESTED_FDNUM)) == Dual{TestTag,Dual{TestTag,V,M},N} - @test_throws ForwardDiff.DualMismatchError(TestTag, OuterTestTag) ForwardDiff.valtype(TestTag, Dual{OuterTestTag}(PRIMAL, PARTIALS)) - @test_throws ForwardDiff.DualMismatchError(TestTag, OuterTestTag) ForwardDiff.valtype(TestTag, typeof(Dual{OuterTestTag}(PRIMAL, PARTIALS))) - @test_throws ForwardDiff.DualMismatchError(TestTag, OuterTestTag) ForwardDiff.valtype(TestTag, Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS)) - @test_throws ForwardDiff.DualMismatchError(TestTag, OuterTestTag) ForwardDiff.valtype(TestTag, typeof(Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS))) + OUTER_FDNUM = Dual{OuterTestTag}(PRIMAL, PARTIALS) + @test ForwardDiff.valtype(TestTag, OUTER_FDNUM) == Dual{OuterTestTag,V,N} + @test ForwardDiff.valtype(TestTag, typeof(OUTER_FDNUM)) == Dual{OuterTestTag,V,N} + @test value(TestTag, OUTER_FDNUM) === OUTER_FDNUM + @test partials(TestTag, OUTER_FDNUM, 1) === zero(OUTER_FDNUM) + @test partials(TestTag, OUTER_FDNUM) === Partials{0,typeof(OUTER_FDNUM)}(()) + INNER_FDNUM = Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS) + @test ForwardDiff.valtype(TestTag, INNER_FDNUM) == Dual{OuterTestTag,V,N} + @test ForwardDiff.valtype(TestTag, typeof(INNER_FDNUM)) == Dual{OuterTestTag,V,N} + @test value(TestTag, INNER_FDNUM) === Dual{OuterTestTag}(PRIMAL, PARTIALS) + for j in 1:M + @test partials(TestTag, INNER_FDNUM, j) === Dual{OuterTestTag}(M_PARTIALS[j], zero(PARTIALS)) + end + @test partials(TestTag, INNER_FDNUM) === Partials{M,Dual{OuterTestTag,V,N}}(ntuple(j -> Dual{OuterTestTag}(M_PARTIALS[j], zero(PARTIALS)), M)) ##################### # Generic Functions # From ab28a5d30f42bc6c4118435c36e7bb27221559ad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Wed, 23 Sep 2026 22:45:33 +0200 Subject: [PATCH 2/8] Deprecate accessors without a tag and require tags to be types `value(x)`, `partials(x)`, `partials(x, i...)`, `npartials(d)` and `partials(T, x, i, j...)` extract the outermost layer of a nested `Dual`, which depends on the order of the tags. They are deprecated in favor of `value(T, x)`, `partials(T, x)`, `partials(T, x, i)` and `npartials(T, D)`, which all internal code now uses. Dispatching on the tag requires it to be a type, so constructing a `Dual` with a non-type tag such as `Dual{1}` now throws an `ArgumentError`. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/src/dev/how_it_works.md | 2 +- docs/src/user/advanced.md | 8 +- src/ForwardDiff.jl | 1 + src/deprecated.jl | 45 ++++++ src/derivative.jl | 2 +- src/dual.jl | 275 +++++++++++++++++------------------ src/gradient.jl | 9 +- src/hessian.jl | 4 +- test/DualTest.jl | 182 ++++++++++++----------- test/GradientTest.jl | 6 +- test/JacobianTest.jl | 4 +- test/MiscTest.jl | 2 +- test/SeedTest.jl | 6 +- 13 files changed, 300 insertions(+), 246 deletions(-) create mode 100644 src/deprecated.jl diff --git a/docs/src/dev/how_it_works.md b/docs/src/dev/how_it_works.md index c8746cc3..c3263e06 100644 --- a/docs/src/dev/how_it_works.md +++ b/docs/src/dev/how_it_works.md @@ -43,7 +43,7 @@ of the function, propagating the derivative via multiplication. For example, `Ba can be overloaded on `Dual` like so: ```julia -Base.sin(d::Dual{T}) where {T} = Dual{T}(sin(value(d)), cos(value(d)) * partials(d)) +Base.sin(d::Dual{T}) where {T} = Dual{T}(sin(value(T, d)), cos(value(T, d)) * partials(T, d)) ``` If we assume that a general function `f` is composed of entirely of these elementary diff --git a/docs/src/user/advanced.md b/docs/src/user/advanced.md index 69468751..6955ff1c 100644 --- a/docs/src/user/advanced.md +++ b/docs/src/user/advanced.md @@ -128,8 +128,8 @@ aren't sensitive to the input and thus cause ForwardDiff to incorrectly return ` ```julia-repl # the dual number's perturbation component is zero, so this # variable should not propagate derivative information -julia> log(ForwardDiff.Dual{:tag}(0.0, 0.0)) -Dual{:tag}(-Inf,NaN) # oops, this NaN should be 0.0 +julia> log(ForwardDiff.Dual(0.0, 0.0)) +Dual{Nothing}(-Inf,NaN) # oops, this NaN should be 0.0 ``` Here, ForwardDiff computes the derivative of `log(0.0)` as `NaN` and then propagates @@ -166,8 +166,8 @@ julia> set_preferences!(UUID("f6369f11-7733-5829-9624-2563aa707210"), "nansafe_m julia> using ForwardDiff -julia> log(ForwardDiff.Dual{:tag}(0.0, 0.0)) -Dual{:tag}(-Inf,0.0) +julia> log(ForwardDiff.Dual(0.0, 0.0)) +Dual{Nothing}(-Inf,0.0) ``` In the future, we plan on allowing users and downstream library authors to dynamically diff --git a/src/ForwardDiff.jl b/src/ForwardDiff.jl index b16b986b..574dbe8a 100644 --- a/src/ForwardDiff.jl +++ b/src/ForwardDiff.jl @@ -21,6 +21,7 @@ include("derivative.jl") include("gradient.jl") include("jacobian.jl") include("hessian.jl") +include("deprecated.jl") export DiffResults diff --git a/src/deprecated.jl b/src/deprecated.jl new file mode 100644 index 00000000..edc3324b --- /dev/null +++ b/src/deprecated.jl @@ -0,0 +1,45 @@ +# Accessors without a tag extract the outermost layer of a nested `Dual`, which depends on +# the order of the tags rather than on the derivative the caller is interested in. + +function depwarn_untagged(f::Symbol, replacement::String) + Base.depwarn("`ForwardDiff.$f` without a tag is deprecated, use `$replacement` with the tag `T` of the derivative instead.", f) +end + +@inline outer_partials(x, i) = zero(x) +@inline outer_partials(d::Dual{T}, i) where {T} = partials(T, d, i) +@inline outer_partials(x, i, j, k...) = outer_partials(outer_partials(x, i), j, k...) + +function value(x) + depwarn_untagged(:value, "ForwardDiff.value(T, x)") + return x +end +function value(d::Dual{T}) where {T} + depwarn_untagged(:value, "ForwardDiff.value(T, d)") + return value(T, d) +end + +function partials(x) + depwarn_untagged(:partials, "ForwardDiff.partials(T, x)") + return Partials{0,typeof(x)}(tuple()) +end +function partials(d::Dual{T}) where {T} + depwarn_untagged(:partials, "ForwardDiff.partials(T, d)") + return partials(T, d) +end +function partials(x, i, j...) + depwarn_untagged(:partials, "ForwardDiff.partials(T, x, i)") + return outer_partials(x, i, j...) +end +function partials(::Type{T}, x, i, j, k...) where {T} + Base.depwarn("`ForwardDiff.partials(T, x, i, j...)` is deprecated, use `ForwardDiff.partials(S, ForwardDiff.partials(T, x, i), j)` with the tag `S` of the inner derivative instead.", :partials) + return outer_partials(partials(T, x, i), j, k...) +end + +function npartials(::Dual{T,V,N}) where {T,V,N} + depwarn_untagged(:npartials, "ForwardDiff.npartials(T, typeof(d))") + return N +end +function npartials(::Type{Dual{T,V,N}}) where {T,V,N} + depwarn_untagged(:npartials, "ForwardDiff.npartials(T, D)") + return N +end diff --git a/src/derivative.jl b/src/derivative.jl index 0c8a6c05..892f21a7 100644 --- a/src/derivative.jl +++ b/src/derivative.jl @@ -29,7 +29,7 @@ Set `check` to `Val{false}()` to disable tag checking. This can lead to perturba ydual = cfg.duals seed_zero_partials!(ydual, y) f!(ydual, Dual{T}(x, one(x))) - map!(value, y, ydual) + map!(d -> value(T, d), y, ydual) return extract_derivative(T, ydual) end diff --git a/src/dual.jl b/src/dual.jl index 4bd1c646..7bd5dc5a 100644 --- a/src/dual.jl +++ b/src/dual.jl @@ -15,6 +15,7 @@ struct Dual{T,V,N} <: Real value::V partials::Partials{N,V} function Dual{T, V, N}(value::V, partials::Partials{N, V}) where {T, V, N} + T isa Type || throw_invalid_tag(T) can_dual(V) || throw_cannot_dual(V) new{T, V, N}(value, partials) end @@ -37,6 +38,10 @@ end Base.showerror(io::IO, e::DualMismatchError{A,B}) where {A,B} = print(io, "Cannot determine ordering of Dual tags ", e.a, " and ", e.b) +@noinline function throw_invalid_tag(T) + throw(ArgumentError(lazy"The tag of a Dual must be a type, got $(repr(T)).")) +end + @noinline function throw_cannot_dual(V::Type) throw(ArgumentError(lazy"Cannot create a dual over scalar type $V. If the type behaves as a scalar, define ForwardDiff.can_dual(::Type{$V}) = true.")) end @@ -85,9 +90,6 @@ Dual{T,V,N}(x::Base.TwicePrecision) where {T,V,N} = # Utility/Accessor Functions # ############################## -@inline value(x) = x -@inline value(d::Dual) = d.value - # Whether a `Dual` with tag `T` occurs anywhere in the nesting of `D` @inline hastag(::Type{T}, ::Type) where {T} = false @inline hastag(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} = S === T || hastag(T, V) @@ -98,7 +100,7 @@ Dual{T,V,N}(x::Base.TwicePrecision) where {T,V,N} = # the order in which the layers are nested. Without a layer with tag `T` the value is the # number itself and the partials are zero. @inline value(::Type{T}, x) where T = x -@inline value(::Type{T}, d::Dual{T}) where T = value(d) +@inline value(::Type{T}, d::Dual{T}) where T = d.value @inline function value(::Type{T}, d::Dual{S,V}) where {T,S,V} if hastag(T, typeof(d)) Dual{S}(value(T, value(S, d)), map_partials(p -> value(T, p), valtype(T, V), partials(S, d))) @@ -107,23 +109,17 @@ Dual{T,V,N}(x::Base.TwicePrecision) where {T,V,N} = end end -@inline partials(x) = Partials{0,typeof(x)}(tuple()) -@inline partials(d::Dual) = d.partials -@inline partials(x, i...) = zero(x) -@inline Base.@propagate_inbounds partials(d::Dual, i) = d.partials[i] -@inline Base.@propagate_inbounds partials(d::Dual, i, j) = partials(d, i).partials[j] -@inline Base.@propagate_inbounds partials(d::Dual, i, j, k...) = partials(partials(d, i, j), k...) - -@inline Base.@propagate_inbounds partials(::Type{T}, x, i...) where T = partials(x, i...) -@inline Base.@propagate_inbounds partials(::Type{T}, d::Dual{T}, i...) where T = partials(d, i...) -@inline Base.@propagate_inbounds function partials(::Type{T}, d::Dual{S,V}, i...) where {T,S,V} +@inline partials(::Type{T}, x) where T = Partials{0,typeof(x)}(tuple()) +@inline partials(::Type{T}, x, i) where T = zero(x) +@inline partials(::Type{T}, d::Dual{T}) where T = d.partials +@inline Base.@propagate_inbounds partials(::Type{T}, d::Dual{T}, i) where T = d.partials[i] +@inline Base.@propagate_inbounds function partials(::Type{T}, d::Dual{S,V}, i) where {T,S,V} if hastag(T, typeof(d)) - Dual{S}(partials(T, value(S, d), i...), map_partials(p -> partials(T, p, i...), valtype(T, V), partials(S, d))) + Dual{S}(partials(T, value(S, d), i), map_partials(p -> partials(T, p, i), valtype(T, V), partials(S, d))) else zero(d) end end -@inline partials(::Type{T}, d::Dual{T}) where {T} = partials(d) @inline function partials(::Type{T}, d::Dual) where {T} if hastag(T, typeof(d)) Partials{npartials(T, typeof(d)),valtype(T, typeof(d))}(ntuple(i -> partials(T, d, i), Val(npartials(T, typeof(d))))) @@ -132,9 +128,6 @@ end end end -@inline npartials(::Dual{T,V,N}) where {T,V,N} = N -@inline npartials(::Type{Dual{T,V,N}}) where {T,V,N} = N - @inline npartials(::Type{T}, ::Type) where {T} = 0 @inline npartials(::Type{T}, ::Type{Dual{T,V,N}}) where {T,V,N} = N @inline npartials(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} = npartials(T, V) @@ -264,9 +257,9 @@ function unary_dual_definition(M, f) end) return quote @inline function $M.$f(d::$FD.Dual{T}) where T - x = $FD.value(d) + x = $FD.value(T, d) $work - return $FD.dual_definition_retval(Val{T}(), val, deriv, $FD.partials(d)) + return $FD.dual_definition_retval(Val{T}(), val, deriv, $FD.partials(T, d)) end end end @@ -294,19 +287,19 @@ function binary_dual_definition(M, f) $FD.@define_binary_dual_op( $M.$f, begin - vx, vy = $FD.value(x), $FD.value(y) + vx, vy = $FD.value(Txy, x), $FD.value(Txy, y) $xy_work - return $FD.dual_definition_retval(Val{Txy}(), val, dvx, $FD.partials(x), dvy, $FD.partials(y)) + return $FD.dual_definition_retval(Val{Txy}(), val, dvx, $FD.partials(Txy, x), dvy, $FD.partials(Txy, y)) end, begin - vx = $FD.value(x) + vx = $FD.value(Tx, x) $x_work - return $FD.dual_definition_retval(Val{Tx}(), val, dvx, $FD.partials(x)) + return $FD.dual_definition_retval(Val{Tx}(), val, dvx, $FD.partials(Tx, x)) end, begin - vy = $FD.value(y) + vy = $FD.value(Ty, y) $y_work - return $FD.dual_definition_retval(Val{Ty}(), val, dvy, $FD.partials(y)) + return $FD.dual_definition_retval(Val{Ty}(), val, dvy, $FD.partials(Ty, y)) end ) end @@ -319,12 +312,12 @@ end Base.copy(d::Dual) = d -Base.eps(d::Dual) = eps(value(d)) +Base.eps(d::Dual{T}) where {T} = eps(value(T, d)) Base.eps(::Type{D}) where {D<:Dual} = eps(valtype(D)) # The `base` keyword was added in Julia 1.8: # https://github.com/JuliaLang/julia/pull/42428 -Base.precision(d::Dual; base::Integer=2) = precision(value(d); base=base) +Base.precision(d::Dual{T}; base::Integer=2) where {T} = precision(value(T, d); base=base) function Base.precision(::Type{D}; base::Integer=2) where {D<:Dual} precision(valtype(D); base=base) end @@ -341,26 +334,26 @@ Base.rtoldefault(::Type{D}) where {D<:Dual} = Base.rtoldefault(valtype(D)) # Base derives floor/ceil/trunc/round from `round(x, ::RoundingMode)`: # https://docs.julialang.org/en/v1/manual/interfaces/#man-rounding-interface -Base.round(d::Dual, r::RoundingMode) = round(value(d), r) +Base.round(d::Dual{T}, r::RoundingMode) where {T} = round(value(T, d), r) # Julia 1.11 added the generic `f(::Type{T}, x)` fallbacks, so these can be # dropped once 1.11 is the minimum supported version. if VERSION < v"1.11" - Base.floor(::Type{R}, d::Dual) where {R<:Real} = floor(R, value(d)) - Base.ceil(::Type{R}, d::Dual) where {R<:Real} = ceil(R, value(d)) - Base.trunc(::Type{R}, d::Dual) where {R<:Real} = trunc(R, value(d)) - Base.round(::Type{R}, d::Dual) where {R<:Real} = round(R, value(d)) + Base.floor(::Type{R}, d::Dual{T}) where {R<:Real,T} = floor(R, value(T, d)) + Base.ceil(::Type{R}, d::Dual{T}) where {R<:Real,T} = ceil(R, value(T, d)) + Base.trunc(::Type{R}, d::Dual{T}) where {R<:Real,T} = trunc(R, value(T, d)) + Base.round(::Type{R}, d::Dual{T}) where {R<:Real,T} = round(R, value(T, d)) end -Base.fld(x::Dual, y::Dual) = fld(value(x), value(y)) +Base.fld(x::Dual{Tx}, y::Dual{Ty}) where {Tx,Ty} = fld(value(Tx, x), value(Ty, y)) -Base.cld(x::Dual, y::Dual) = cld(value(x), value(y)) +Base.cld(x::Dual{Tx}, y::Dual{Ty}) where {Tx,Ty} = cld(value(Tx, x), value(Ty, y)) -Base.exponent(x::Dual) = exponent(value(x)) +Base.exponent(x::Dual{T}) where {T} = exponent(value(T, x)) -Base.div(x::Dual, y::Dual, r::RoundingMode) = div(value(x), value(y), r) +Base.div(x::Dual{Tx}, y::Dual{Ty}, r::RoundingMode) where {Tx,Ty} = div(value(Tx, x), value(Ty, y), r) -Base.hash(d::Dual, hsh::UInt) = hash(value(d), hsh) +Base.hash(d::Dual{T}, hsh::UInt) where {T} = hash(value(T, d), hsh) function Base.read(io::IO, ::Type{Dual{T,V,N}}) where {T,V,N} value = read(io, V) @@ -368,9 +361,9 @@ function Base.read(io::IO, ::Type{Dual{T,V,N}}) where {T,V,N} return Dual{T,V,N}(value, partials) end -function Base.write(io::IO, d::Dual) - write(io, value(d)) - write(io, partials(d)) +function Base.write(io::IO, d::Dual{T}) where {T} + write(io, value(T, d)) + write(io, partials(T, d)) end @inline Base.zero(d::Dual) = zero(typeof(d)) @@ -379,16 +372,16 @@ end @inline Base.one(d::Dual) = one(typeof(d)) @inline Base.one(::Type{Dual{T,V,N}}) where {T,V,N} = Dual{T}(one(V), zero(Partials{N,V})) -@inline function Base.Int(d::Dual) - all(iszero, partials(d)) || throw(InexactError(:Int, Int, d)) - Int(value(d)) +@inline function Base.Int(d::Dual{T}) where {T} + all(iszero, partials(T, d)) || throw(InexactError(:Int, Int, d)) + Int(value(T, d)) end -@inline function Base.Integer(d::Dual) - all(iszero, partials(d)) || throw(InexactError(:Integer, Integer, d)) - Integer(value(d)) +@inline function Base.Integer(d::Dual{T}) where {T} + all(iszero, partials(T, d)) || throw(InexactError(:Integer, Integer, d)) + Integer(value(T, d)) end -@inline Random.rand(rng::AbstractRNG, d::Dual) = rand(rng, value(d)) +@inline Random.rand(rng::AbstractRNG, d::Dual{T}) where {T} = rand(rng, value(T, d)) @inline Random.rand(::Type{Dual{T,V,N}}) where {T,V,N} = Dual{T}(rand(V), zero(Partials{N,V})) @inline Random.rand(rng::AbstractRNG, ::Type{Dual{T,V,N}}) where {T,V,N} = Dual{T}(rand(rng, V), zero(Partials{N,V})) @inline Random.randn(::Type{Dual{T,V,N}}) where {T,V,N} = Dual{T}(randn(V), zero(Partials{N,V})) @@ -399,10 +392,10 @@ end # Predicates # #------------# -isconstant(d::Dual) = iszero(partials(d)) +isconstant(d::Dual{T}) where {T} = iszero(partials(T, d)) for pred in UNARY_PREDICATES - @eval Base.$(pred)(d::Dual) = $(pred)(value(d)) + @eval Base.$(pred)(d::Dual{T}) where {T} = $(pred)(value(T, d)) end # Before PR#481 this loop ran over this list: @@ -410,33 +403,33 @@ end # Not a minimal set, as Base defines some in terms of others. @define_binary_dual_op( Base.:(<), - (value(x) < value(y)) || (value(x) == value(y) && (partials(x) < partials(y))), - (value(x) < y) || (value(x) == y && (partials(x) < zero(partials(x)))), - (x < value(y)) || (x == value(y) && (zero(partials(y)) < partials(y))), + (value(Txy, x) < value(Txy, y)) || (value(Txy, x) == value(Txy, y) && (partials(Txy, x) < partials(Txy, y))), + (value(Tx, x) < y) || (value(Tx, x) == y && (partials(Tx, x) < zero(partials(Tx, x)))), + (x < value(Ty, y)) || (x == value(Ty, y) && (zero(partials(Ty, y)) < partials(Ty, y))), ) @define_binary_dual_op( Base.:(<=), - (value(x) < value(y)) || (value(x) == value(y) && (partials(x) <= partials(y))), - (value(x) < y) || (value(x) == y && (partials(x) <= zero(partials(x)))), - (x < value(y)) || (x == value(y) && (zero(partials(y)) <= partials(y))), + (value(Txy, x) < value(Txy, y)) || (value(Txy, x) == value(Txy, y) && (partials(Txy, x) <= partials(Txy, y))), + (value(Tx, x) < y) || (value(Tx, x) == y && (partials(Tx, x) <= zero(partials(Tx, x)))), + (x < value(Ty, y)) || (x == value(Ty, y) && (zero(partials(Ty, y)) <= partials(Ty, y))), ) @define_binary_dual_op( Base.isless, - isless(value(x), value(y)) || (isequal(value(x), value(y)) && isless(partials(x), partials(y))), - isless(value(x), y) || (isequal(value(x), y) && isless(partials(x), zero(partials(x)))), - isless(x, value(y)) || (isequal(x, value(y)) && isless(zero(partials(y)), partials(y))), + isless(value(Txy, x), value(Txy, y)) || (isequal(value(Txy, x), value(Txy, y)) && isless(partials(Txy, x), partials(Txy, y))), + isless(value(Tx, x), y) || (isequal(value(Tx, x), y) && isless(partials(Tx, x), zero(partials(Tx, x)))), + isless(x, value(Ty, y)) || (isequal(x, value(Ty, y)) && isless(zero(partials(Ty, y)), partials(Ty, y))), ) -Base.iszero(x::Dual) = iszero(value(x)) && iszero(partials(x)) # shortcut, equivalent to x == zero(x) +Base.iszero(x::Dual{T}) where {T} = iszero(value(T, x)) && iszero(partials(T, x)) # shortcut, equivalent to x == zero(x) for pred in [:isequal, :(==)] @eval begin @define_binary_dual_op( Base.$(pred), - $(pred)(value(x), value(y)) && $(pred)(partials(x), partials(y)), - $(pred)(value(x), y) && iszero(partials(x)), - $(pred)(x, value(y)) && iszero(partials(y)), + $(pred)(value(Txy, x), value(Txy, y)) && $(pred)(partials(Txy, x), partials(Txy, y)), + $(pred)(value(Tx, x), y) && iszero(partials(Tx, x)), + $(pred)(x, value(Ty, y)) && iszero(partials(Ty, y)), ) end end @@ -474,7 +467,7 @@ for R in (AbstractIrrational, Real, BigFloat, Bool) end end -@inline Base.convert(::Type{Dual{T,V,N}}, d::Dual{T}) where {T,V,N} = Dual{T}(V(value(d)), convert(Partials{N,V}, partials(d))) +@inline Base.convert(::Type{Dual{T,V,N}}, d::Dual{T}) where {T,V,N} = Dual{T}(V(value(T, d)), convert(Partials{N,V}, partials(T, d))) @inline Base.convert(::Type{Dual{T,Dual{T,V,M},N}}, d::Dual{T,V,M}) where {T,V,N,M} = Dual{T}(d, Partials{N,Dual{T,V,M}}(zero_tuple(NTuple{N,Dual{T,V,M}}))) @inline Base.convert(::Type{Dual{T,V,N}}, x) where {T,V,N} = Dual{T}(V(x), zero(Partials{N,V})) @inline Base.convert(::Type{Dual{T,V,N}}, x::Number) where {T,V,N} = Dual{T}(V(x), zero(Partials{N,V})) @@ -513,29 +506,29 @@ end @define_binary_dual_op( Base.:+, begin - vx, vy = value(x), value(y) - Dual{Txy}(vx + vy, partials(x) + partials(y)) + vx, vy = value(Txy, x), value(Txy, y) + Dual{Txy}(vx + vy, partials(Txy, x) + partials(Txy, y)) end, - Dual{Tx}(value(x) + y, partials(x)), - Dual{Ty}(x + value(y), partials(y)) + Dual{Tx}(value(Tx, x) + y, partials(Tx, x)), + Dual{Ty}(x + value(Ty, y), partials(Ty, y)) ) @define_binary_dual_op( Base.:-, begin - vx, vy = value(x), value(y) - Dual{Txy}(vx - vy, partials(x) - partials(y)) + vx, vy = value(Txy, x), value(Txy, y) + Dual{Txy}(vx - vy, partials(Txy, x) - partials(Txy, y)) end, - Dual{Tx}(value(x) - y, partials(x)), - Dual{Ty}(x - value(y), -partials(y)) + Dual{Tx}(value(Tx, x) - y, partials(Tx, x)), + Dual{Ty}(x - value(Ty, y), -partials(Ty, y)) ) -@inline Base.:-(d::Dual{T}) where {T} = Dual{T}(-value(d), -partials(d)) +@inline Base.:-(d::Dual{T}) where {T} = Dual{T}(-value(T, d), -partials(T, d)) # * # #---# -@inline Base.:*(d::Dual, x::Bool) = x ? d : (signbit(value(d))==0 ? zero(d) : -zero(d)) +@inline Base.:*(d::Dual{T}, x::Bool) where {T} = x ? d : (signbit(value(T, d))==0 ? zero(d) : -zero(d)) @inline Base.:*(x::Bool, d::Dual) = d * x # / # @@ -546,14 +539,14 @@ end @define_binary_dual_op( Base.:/, begin - vx, vy = value(x), value(y) - Dual{Txy}(vx / vy, _div_partials(partials(x), partials(y), vx, vy)) + vx, vy = value(Txy, x), value(Txy, y) + Dual{Txy}(vx / vy, _div_partials(partials(Txy, x), partials(Txy, y), vx, vy)) end, - Dual{Tx}(value(x) / y, partials(x) / y), + Dual{Tx}(value(Tx, x) / y, partials(Tx, x) / y), begin - v = value(y) + v = value(Ty, y) divv = x / v - Dual{Ty}(divv, -(divv / v) * partials(y)) + Dual{Ty}(divv, -(divv / v) * partials(Ty, y)) end ) @@ -565,7 +558,7 @@ for (f, log) in ((:(Base.:^), :(Base.log)), (:(NaNMath.pow), :(NaNMath.log))) @define_binary_dual_op( $f, begin - vx, vy = value(x), value(y) + vx, vy = value(Txy, x), value(Txy, y) expv = ($f)(vx, vy) powval = vy * ($f)(vx, vy - 1) if isconstant(y) @@ -575,38 +568,38 @@ for (f, log) in ((:(Base.:^), :(Base.log)), (:(NaNMath.pow), :(NaNMath.log))) else logval = expv * ($log)(vx) end - new_partials = _mul_partials(partials(x), partials(y), powval, logval) + new_partials = _mul_partials(partials(Txy, x), partials(Txy, y), powval, logval) return Dual{Txy}(expv, new_partials) end, begin - v = value(x) + v = value(Tx, x) expv = ($f)(v, y) - if y == zero(y) || iszero(partials(x)) - new_partials = zero(partials(x)) + if y == zero(y) || iszero(partials(Tx, x)) + new_partials = zero(partials(Tx, x)) else - new_partials = partials(x) * y * ($f)(v, y - 1) + new_partials = partials(Tx, x) * y * ($f)(v, y - 1) end return Dual{Tx}(expv, new_partials) end, begin - v = value(y) + v = value(Ty, y) expv = ($f)(x, v) deriv = (iszero(x) && v > 0) ? zero(expv) : expv*($log)(oftype(expv, x)) - return Dual{Ty}(expv, deriv * partials(y)) + return Dual{Ty}(expv, deriv * partials(Ty, y)) end ) end end @inline Base.literal_pow(::typeof(^), x::Dual{T}, ::Val{0}) where {T} = - Dual{T}(one(value(x)), zero(partials(x))) + Dual{T}(one(value(T, x)), zero(partials(T, x))) for y in 1:3 @eval @inline function Base.literal_pow(::typeof(^), x::Dual{T}, ::Val{$y}) where {T} - v = value(x) + v = value(T, x) expv = v^$y deriv = $y * v^$(y - 1) - return Dual{T}(expv, deriv * partials(x)) + return Dual{T}(expv, deriv * partials(T, x)) end end @@ -614,11 +607,11 @@ end #-------# @inline function calc_hypot(x, y, z, ::Type{T}) where T - vx = value(x) - vy = value(y) - vz = value(z) + vx = value(T, x) + vy = value(T, y) + vz = value(T, z) h = hypot(vx, vy, vz) - p = (vx / h) * partials(x) + (vy / h) * partials(y) + (vz / h) * partials(z) + p = (vx / h) * partials(T, x) + (vy / h) * partials(T, y) + (vz / h) * partials(T, z) return Dual{T}(h, p) end @@ -639,27 +632,27 @@ end @generated function calc_fma_xyz(x::Dual{T,<:Any,N}, y::Dual{T,<:Any,N}, z::Dual{T,<:Any,N}) where {T,N} - ex = Expr(:tuple, [:(fma(value(x), partials(y)[$i], fma(value(y), partials(x)[$i], partials(z)[$i]))) for i in 1:N]...) + ex = Expr(:tuple, [:(fma(value(T, x), partials(T, y)[$i], fma(value(T, y), partials(T, x)[$i], partials(T, z)[$i]))) for i in 1:N]...) return quote $(Expr(:meta, :inline)) - v = fma(value(x), value(y), value(z)) + v = fma(value(T, x), value(T, y), value(T, z)) return Dual{T}(v, $ex) end end @inline function calc_fma_xy(x::Dual{T}, y::Dual{T}, z::Real) where T - vx, vy = value(x), value(y) + vx, vy = value(T, x), value(T, y) result = fma(vx, vy, z) - return Dual{T}(result, _mul_partials(partials(x), partials(y), vy, vx)) + return Dual{T}(result, _mul_partials(partials(T, x), partials(T, y), vy, vx)) end @generated function calc_fma_xz(x::Dual{T,<:Any,N}, y::Real, z::Dual{T,<:Any,N}) where {T,N} - ex = Expr(:tuple, [:(fma(partials(x)[$i], y, partials(z)[$i])) for i in 1:N]...) + ex = Expr(:tuple, [:(fma(partials(T, x)[$i], y, partials(T, z)[$i])) for i in 1:N]...) return quote $(Expr(:meta, :inline)) - v = fma(value(x), y, value(z)) + v = fma(value(T, x), y, value(T, z)) Dual{T}(v, $ex) end end @@ -670,9 +663,9 @@ end calc_fma_xy(x, y, z), # xy_body calc_fma_xz(x, y, z), # xz_body Base.fma(y, x, z), # yz_body - Dual{Tx}(fma(value(x), y, z), partials(x) * y), # x_body + Dual{Tx}(fma(value(Tx, x), y, z), partials(Tx, x) * y), # x_body Base.fma(y, x, z), # y_body - Dual{Tz}(fma(x, y, value(z)), partials(z)) # z_body + Dual{Tz}(fma(x, y, value(Tz, z)), partials(Tz, z)) # z_body ) # muladd # @@ -681,27 +674,27 @@ end @generated function calc_muladd_xyz(x::Dual{T,<:Any,N}, y::Dual{T,<:Any,N}, z::Dual{T,<:Any,N}) where {T,N} - ex = Expr(:tuple, [:(muladd(value(x), partials(y)[$i], muladd(value(y), partials(x)[$i], partials(z)[$i]))) for i in 1:N]...) + ex = Expr(:tuple, [:(muladd(value(T, x), partials(T, y)[$i], muladd(value(T, y), partials(T, x)[$i], partials(T, z)[$i]))) for i in 1:N]...) return quote $(Expr(:meta, :inline)) - v = muladd(value(x), value(y), value(z)) + v = muladd(value(T, x), value(T, y), value(T, z)) return Dual{T}(v, $ex) end end @inline function calc_muladd_xy(x::Dual{T}, y::Dual{T}, z::Real) where T - vx, vy = value(x), value(y) + vx, vy = value(T, x), value(T, y) result = muladd(vx, vy, z) - return Dual{T}(result, _mul_partials(partials(x), partials(y), vy, vx)) + return Dual{T}(result, _mul_partials(partials(T, x), partials(T, y), vy, vx)) end @generated function calc_muladd_xz(x::Dual{T,<:Any,N}, y::Real, z::Dual{T,<:Any,N}) where {T,N} - ex = Expr(:tuple, [:(muladd(partials(x)[$i], y, partials(z)[$i])) for i in 1:N]...) + ex = Expr(:tuple, [:(muladd(partials(T, x)[$i], y, partials(T, z)[$i])) for i in 1:N]...) return quote $(Expr(:meta, :inline)) - v = muladd(value(x), y, value(z)) + v = muladd(value(T, x), y, value(T, z)) Dual{T}(v, $ex) end end @@ -712,35 +705,35 @@ end calc_muladd_xy(x, y, z), # xy_body calc_muladd_xz(x, y, z), # xz_body Base.muladd(y, x, z), # yz_body - Dual{Tx}(muladd(value(x), y, z), partials(x) * y), # x_body + Dual{Tx}(muladd(value(Tx, x), y, z), partials(Tx, x) * y), # x_body Base.muladd(y, x, z), # y_body - Dual{Tz}(muladd(x, y, value(z)), partials(z)) # z_body + Dual{Tz}(muladd(x, y, value(Tz, z)), partials(Tz, z)) # z_body ) # sin/cos # #--------# function Base.sin(d::Dual{T}) where T - s, c = sincos(value(d)) - return Dual{T}(s, c * partials(d)) + s, c = sincos(value(T, d)) + return Dual{T}(s, c * partials(T, d)) end function Base.cos(d::Dual{T}) where T - s, c = sincos(value(d)) - return Dual{T}(c, -s * partials(d)) + s, c = sincos(value(T, d)) + return Dual{T}(c, -s * partials(T, d)) end @inline function Base.sincos(d::Dual{T}) where T - sd, cd = sincos(value(d)) - return (Dual{T}(sd, cd * partials(d)), Dual{T}(cd, -sd * partials(d))) + sd, cd = sincos(value(T, d)) + return (Dual{T}(sd, cd * partials(T, d)), Dual{T}(cd, -sd * partials(T, d))) end # sincospi # #----------# @inline function Base.sincospi(d::Dual{T}) where T - sd, cd = sincospi(value(d)) - return (Dual{T}(sd, cd * π * partials(d)), Dual{T}(cd, -sd * π * partials(d))) + sd, cd = sincospi(value(T, d)) + return (Dual{T}(sd, cd * π * partials(T, d)), Dual{T}(cd, -sd * π * partials(T, d))) end # LinearAlgebra.givensAlgorithm # @@ -758,35 +751,35 @@ end @define_binary_dual_op( LinearAlgebra.givensAlgorithm, begin - vx, vy = value(x), value(y) + vx, vy = value(Txy, x), value(Txy, y) c, s, u = LinearAlgebra.givensAlgorithm(vx, vy) ∂c∂x = s^2 / u ∂c∂y = ∂s∂x = -(c * s / u) ∂s∂y = c^2 / u - ∂x = partials(x) - ∂y = partials(y) + ∂x = partials(Txy, x) + ∂y = partials(Txy, y) ∂c = _mul_partials(∂x, ∂y, ∂c∂x, ∂c∂y) ∂s = _mul_partials(∂x, ∂y, ∂s∂x, ∂s∂y) ∂u = _mul_partials(∂x, ∂y, c, s) return Dual{Txy}(c, ∂c), Dual{Txy}(s, ∂s), Dual{Txy}(u, ∂u) end, begin - vx = value(x) + vx = value(Tx, x) c, s, u = LinearAlgebra.givensAlgorithm(vx, y) ∂c∂x = s^2 / u ∂s∂x = -(c * s / u) - ∂x = partials(x) + ∂x = partials(Tx, x) ∂c = ∂c∂x * ∂x ∂s = ∂s∂x * ∂x ∂u = c * ∂x return Dual{Tx}(c, ∂c), Dual{Tx}(s, ∂s), Dual{Tx}(u, ∂u) end, begin - vy = value(y) + vy = value(Ty, y) c, s, u = LinearAlgebra.givensAlgorithm(x, vy) ∂c∂y = -(c * s / u) ∂s∂y = c^2 / u - ∂y = partials(y) + ∂y = partials(Ty, y) ∂c = ∂c∂y * ∂y ∂s = ∂s∂y * ∂y ∂u = s * ∂y @@ -798,17 +791,17 @@ end #------------------------------------------------# # Extract structured matrices of primal values and partials -_structured_value(A::Symmetric{Dual{T,V,N}}) where {T,V,N} = Symmetric(map(value, parent(A)), A.uplo === 'U' ? :U : :L) -_structured_value(A::Hermitian{Dual{T,V,N}}) where {T,V,N} = Hermitian(map(value, parent(A)), A.uplo === 'U' ? :U : :L) -_structured_value(A::Hermitian{Complex{Dual{T,V,N}}}) where {T,V,N} = Hermitian(map(z -> splat(complex)(map(value, reim(z))), parent(A)), A.uplo === 'U' ? :U : :L) -_structured_value(A::SymTridiagonal{Dual{T,V,N}}) where {T,V,N} = SymTridiagonal(map(value, A.dv), map(value, A.ev)) +_structured_value(A::Symmetric{Dual{T,V,N}}) where {T,V,N} = Symmetric(map(Base.Fix1(value, T), parent(A)), A.uplo === 'U' ? :U : :L) +_structured_value(A::Hermitian{Dual{T,V,N}}) where {T,V,N} = Hermitian(map(Base.Fix1(value, T), parent(A)), A.uplo === 'U' ? :U : :L) +_structured_value(A::Hermitian{Complex{Dual{T,V,N}}}) where {T,V,N} = Hermitian(map(z -> splat(complex)(map(Base.Fix1(value, T), reim(z))), parent(A)), A.uplo === 'U' ? :U : :L) +_structured_value(A::SymTridiagonal{Dual{T,V,N}}) where {T,V,N} = SymTridiagonal(map(Base.Fix1(value, T), A.dv), map(Base.Fix1(value, T), A.ev)) -_structured_partials(A::Symmetric{Dual{T,V,N}}, j::Int) where {T,V,N} = Symmetric(partials.(parent(A), j), A.uplo === 'U' ? :U : :L) -_structured_partials(A::Hermitian{Dual{T,V,N}}, j::Int) where {T,V,N} = Hermitian(partials.(parent(A), j), A.uplo === 'U' ? :U : :L) +_structured_partials(A::Symmetric{Dual{T,V,N}}, j::Int) where {T,V,N} = Symmetric(partials.(T, parent(A), j), A.uplo === 'U' ? :U : :L) +_structured_partials(A::Hermitian{Dual{T,V,N}}, j::Int) where {T,V,N} = Hermitian(partials.(T, parent(A), j), A.uplo === 'U' ? :U : :L) function _structured_partials(A::Hermitian{Complex{Dual{T,V,N}}}, j::Int) where {T,V,N} - return Hermitian(complex.(partials.(real.(parent(A)), j), partials.(imag.(parent(A)), j)), A.uplo === 'U' ? :U : :L) + return Hermitian(complex.(partials.(T, real.(parent(A)), j), partials.(T, imag.(parent(A)), j)), A.uplo === 'U' ? :U : :L) end -_structured_partials(A::SymTridiagonal{Dual{T,V,N}}, j::Int) where {T,V,N} = SymTridiagonal(partials.(A.dv, j), partials.(A.ev, j)) +_structured_partials(A::SymTridiagonal{Dual{T,V,N}}, j::Int) where {T,V,N} = SymTridiagonal(partials.(T, A.dv, j), partials.(T, A.ev, j)) # Convert arrays of primal values and partials to arrays of Duals function _to_duals(::Val{T}, values::AbstractArray{<:Real}, partials::Tuple{Vararg{AbstractArray{<:Real}}}) where {T} @@ -883,16 +876,16 @@ end #---------------------------------------------------# function SpecialFunctions.logabsgamma(d::Dual{T,<:Real}) where {T} - x = value(d) + x = value(T, d) y, s = SpecialFunctions.logabsgamma(x) - return (Dual{T}(y, SpecialFunctions.digamma(x) * partials(d)), s) + return (Dual{T}(y, SpecialFunctions.digamma(x) * partials(T, d)), s) end # Derivatives wrt to first parameter and precision setting are not supported function SpecialFunctions.gamma_inc(a::Real, d::Dual{T,<:Real}, ind::Integer) where {T} - x = value(d) + x = value(T, d) p, q = SpecialFunctions.gamma_inc(a, x, ind) - ∂p = exp(-x) * x^(a - 1) / SpecialFunctions.gamma(a) * partials(d) + ∂p = exp(-x) * x^(a - 1) / SpecialFunctions.gamma(a) * partials(T, d) return (Dual{T}(p, ∂p), Dual{T}(q, -∂p)) end @@ -901,9 +894,9 @@ end ################### function Base.show(io::IO, d::Dual{T,V,N}) where {T,V,N} - print(io, "Dual{$(repr(T))}(", value(d)) + print(io, "Dual{$(repr(T))}(", value(T, d)) for i in 1:N - print(io, ",", partials(d, i)) + print(io, ",", partials(T, d, i)) end print(io, ")") end @@ -914,4 +907,4 @@ for op in (:(Base.typemin), :(Base.typemax), :(Base.floatmin), :(Base.floatmax)) end end -Printf.tofloat(d::Dual) = Printf.tofloat(value(d)) +Printf.tofloat(d::Dual{T}) where {T} = Printf.tofloat(value(T, d)) diff --git a/src/gradient.jl b/src/gradient.jl index 235ac06e..a74d1b78 100644 --- a/src/gradient.jl +++ b/src/gradient.jl @@ -57,8 +57,13 @@ function extract_gradient!(::Type{T}, result::DiffResult, y::Real) where {T} end function extract_gradient!(::Type{T}, result::DiffResult, dual::Dual) where {T} - result = DiffResults.value!(result, value(T, dual)) - result = DiffResults.gradient!(result, partials(T, dual)) + if hastag(T, typeof(dual)) + result = DiffResults.value!(result, value(T, dual)) + result = DiffResults.gradient!(result, partials(T, dual)) + else + result = DiffResults.value!(result, dual) + fill!(DiffResults.gradient(result), zero(dual)) + end return result end diff --git a/src/hessian.jl b/src/hessian.jl index 9c755c9a..686db496 100644 --- a/src/hessian.jl +++ b/src/hessian.jl @@ -47,10 +47,10 @@ mutable struct InnerGradientForHess{R,C,F} f::F end -function (g::InnerGradientForHess)(y, z) +function (g::InnerGradientForHess{R,<:HessianConfig{T}})(y, z) where {R,T} inner_result = DiffResult(zero(eltype(y)), y) gradient!(inner_result, g.f, z, g.cfg.gradient_config, Val{false}()) - g.result = DiffResults.value!(g.result, value(DiffResults.value(inner_result))) + g.result = DiffResults.value!(g.result, value(T, DiffResults.value(inner_result))) return y end diff --git a/test/DualTest.jl b/test/DualTest.jl index 7a55fed8..8ee4884f 100644 --- a/test/DualTest.jl +++ b/test/DualTest.jl @@ -14,6 +14,7 @@ import LinearAlgebra struct TestTag end struct OuterTestTag end +struct OutermostTestTag end samerng() = MersenneTwister(1) @@ -23,11 +24,9 @@ samerng() = MersenneTwister(1) intrand(V) = V == Int ? rand(2:10) : rand(V) dual_isapprox(a, b) = isapprox(a, b) -dual_isapprox(a::Dual{T,T1,T2}, b::Dual{T,T3,T4}) where {T,T1,T2,T3,T4} = isapprox(value(a), value(b)) && isapprox(partials(a), partials(b)) +dual_isapprox(a::Dual{T,T1,T2}, b::Dual{T,T3,T4}) where {T,T1,T2,T3,T4} = isapprox(value(T, a), value(T, b)) && isapprox(partials(T, a), partials(T, b)) dual_isapprox(a::Dual{T,T1,T2}, b::Dual{T3,T4,T5}) where {T,T1,T2,T3,T4,T5} = error("Tags don't match") -ForwardDiff.:≺(::Type{TestTag}, ::Int) = true -ForwardDiff.:≺(::Int, ::Type{TestTag}) = false ForwardDiff.:≺(::Type{TestTag}, ::Type{OuterTestTag}) = true ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @@ -67,6 +66,8 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test Dual{TestTag}(PRIMAL, PARTIALS...) === FDNUM @test Dual(PRIMAL, PARTIALS...) === Dual{Nothing}(PRIMAL, PARTIALS...) @test Dual(PRIMAL) === Dual{Nothing}(PRIMAL) + @test_throws ArgumentError("The tag of a Dual must be a type, got 1.") Dual{1}(PRIMAL, PARTIALS) + @test_throws ArgumentError("The tag of a Dual must be a type, got :tag.") Dual{:tag}(PRIMAL, PARTIALS) @test typeof(Dual{TestTag}(widen(V)(PRIMAL), PARTIALS)) === Dual{TestTag,widen(V),N} @test typeof(Dual{TestTag}(widen(V)(PRIMAL), PARTIALS.values)) === Dual{TestTag,widen(V),N} @@ -77,25 +78,34 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false # Accessors # ############# - @test value(PRIMAL) == PRIMAL - @test value(FDNUM) == PRIMAL - @test value(NESTED_FDNUM) === Dual{TestTag}(PRIMAL, M_PARTIALS) + @test value(TestTag, PRIMAL) == PRIMAL + @test value(TestTag, FDNUM) == PRIMAL + @test value(TestTag, NESTED_FDNUM) === Dual{TestTag}(PRIMAL, M_PARTIALS) - @test partials(PRIMAL) == Partials{0,V}(tuple()) - @test partials(FDNUM) == PARTIALS - @test partials(NESTED_FDNUM) === NESTED_PARTIALS + @test partials(TestTag, PRIMAL) == Partials{0,V}(tuple()) + @test partials(TestTag, FDNUM) == PARTIALS + @test partials(TestTag, NESTED_FDNUM) === NESTED_PARTIALS for i in 1:N - @test partials(FDNUM, i) == PARTIALS[i] + @test partials(TestTag, FDNUM, i) == PARTIALS[i] + end + + @test ForwardDiff.npartials(TestTag, typeof(FDNUM)) == N + @test ForwardDiff.npartials(TestTag, typeof(NESTED_FDNUM)) == N + + @test (@test_deprecated r"`ForwardDiff.value` without a tag is deprecated" value(PRIMAL)) == PRIMAL + @test (@test_deprecated r"`ForwardDiff.value` without a tag is deprecated" value(FDNUM)) == PRIMAL + @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(PRIMAL)) == Partials{0,V}(tuple()) + @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(FDNUM)) == PARTIALS + for i in 1:N + @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(FDNUM, i)) == PARTIALS[i] for j in 1:M - @test partials(NESTED_FDNUM, i, j) == partials(NESTED_PARTIALS[i], j) + @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(NESTED_FDNUM, i, j)) == partials(TestTag, NESTED_PARTIALS[i], j) + @test (@test_deprecated r"`ForwardDiff.partials\(T, x, i, j...\)` is deprecated" partials(TestTag, NESTED_FDNUM, i, j)) == partials(TestTag, NESTED_PARTIALS[i], j) end end - - @test ForwardDiff.npartials(FDNUM) == N - @test ForwardDiff.npartials(typeof(FDNUM)) == N - @test ForwardDiff.npartials(NESTED_FDNUM) == N - @test ForwardDiff.npartials(typeof(NESTED_FDNUM)) == N + @test (@test_deprecated r"`ForwardDiff.npartials` without a tag is deprecated" ForwardDiff.npartials(FDNUM)) == N + @test (@test_deprecated r"`ForwardDiff.npartials` without a tag is deprecated" ForwardDiff.npartials(typeof(FDNUM))) == N @test ForwardDiff.valtype(FDNUM) == V @test ForwardDiff.valtype(typeof(FDNUM)) == V @@ -254,8 +264,8 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test one(typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(one(V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) if V <: Integer - @test rand(samerng(), FDNUM) == rand(samerng(), value(FDNUM)) - @test rand(samerng(), NESTED_FDNUM) == rand(samerng(), value(NESTED_FDNUM)) + @test rand(samerng(), FDNUM) == rand(samerng(), value(TestTag, FDNUM)) + @test rand(samerng(), NESTED_FDNUM) == rand(samerng(), value(TestTag, NESTED_FDNUM)) elseif V <: AbstractFloat @test rand(samerng(), typeof(FDNUM)) === Dual{TestTag}(rand(samerng(), V), zero(Partials{N,V})) @test rand(samerng(), typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(rand(samerng(), V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) @@ -437,8 +447,8 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test typeof(WIDE_FDNUM) === Dual{TestTag,WIDE_T,N} @test typeof(WIDE_NESTED_FDNUM) === Dual{TestTag,Dual{TestTag,WIDE_T,M},N} - @test value(WIDE_FDNUM) == PRIMAL - @test (value(WIDE_NESTED_FDNUM) == PRIMAL) == (M == 0) + @test value(TestTag, WIDE_FDNUM) == PRIMAL + @test (value(TestTag, WIDE_NESTED_FDNUM) == PRIMAL) == (M == 0) @test convert(Dual, FDNUM) === FDNUM @test convert(Dual, NESTED_FDNUM) === NESTED_FDNUM @@ -456,34 +466,34 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false # Addition/Subtraction # #----------------------# - @test FDNUM + FDNUM2 === Dual{TestTag}(value(FDNUM) + value(FDNUM2), partials(FDNUM) + partials(FDNUM2)) - @test FDNUM + PRIMAL === Dual{TestTag}(value(FDNUM) + PRIMAL, partials(FDNUM)) - @test PRIMAL + FDNUM === Dual{TestTag}(value(FDNUM) + PRIMAL, partials(FDNUM)) + @test FDNUM + FDNUM2 === Dual{TestTag}(value(TestTag, FDNUM) + value(TestTag, FDNUM2), partials(TestTag, FDNUM) + partials(TestTag, FDNUM2)) + @test FDNUM + PRIMAL === Dual{TestTag}(value(TestTag, FDNUM) + PRIMAL, partials(TestTag, FDNUM)) + @test PRIMAL + FDNUM === Dual{TestTag}(value(TestTag, FDNUM) + PRIMAL, partials(TestTag, FDNUM)) - @test NESTED_FDNUM + NESTED_FDNUM2 === Dual{TestTag}(value(NESTED_FDNUM) + value(NESTED_FDNUM2), partials(NESTED_FDNUM) + partials(NESTED_FDNUM2)) - @test NESTED_FDNUM + PRIMAL === Dual{TestTag}(value(NESTED_FDNUM) + PRIMAL, partials(NESTED_FDNUM)) - @test PRIMAL + NESTED_FDNUM === Dual{TestTag}(value(NESTED_FDNUM) + PRIMAL, partials(NESTED_FDNUM)) + @test NESTED_FDNUM + NESTED_FDNUM2 === Dual{TestTag}(value(TestTag, NESTED_FDNUM) + value(TestTag, NESTED_FDNUM2), partials(TestTag, NESTED_FDNUM) + partials(TestTag, NESTED_FDNUM2)) + @test NESTED_FDNUM + PRIMAL === Dual{TestTag}(value(TestTag, NESTED_FDNUM) + PRIMAL, partials(TestTag, NESTED_FDNUM)) + @test PRIMAL + NESTED_FDNUM === Dual{TestTag}(value(TestTag, NESTED_FDNUM) + PRIMAL, partials(TestTag, NESTED_FDNUM)) - @test FDNUM - FDNUM2 === Dual{TestTag}(value(FDNUM) - value(FDNUM2), partials(FDNUM) - partials(FDNUM2)) - @test FDNUM - PRIMAL === Dual{TestTag}(value(FDNUM) - PRIMAL, partials(FDNUM)) - @test PRIMAL - FDNUM === Dual{TestTag}(PRIMAL - value(FDNUM), -(partials(FDNUM))) - @test -(FDNUM) === Dual{TestTag}(-(value(FDNUM)), -(partials(FDNUM))) + @test FDNUM - FDNUM2 === Dual{TestTag}(value(TestTag, FDNUM) - value(TestTag, FDNUM2), partials(TestTag, FDNUM) - partials(TestTag, FDNUM2)) + @test FDNUM - PRIMAL === Dual{TestTag}(value(TestTag, FDNUM) - PRIMAL, partials(TestTag, FDNUM)) + @test PRIMAL - FDNUM === Dual{TestTag}(PRIMAL - value(TestTag, FDNUM), -(partials(TestTag, FDNUM))) + @test -(FDNUM) === Dual{TestTag}(-(value(TestTag, FDNUM)), -(partials(TestTag, FDNUM))) - @test NESTED_FDNUM - NESTED_FDNUM2 === Dual{TestTag}(value(NESTED_FDNUM) - value(NESTED_FDNUM2), partials(NESTED_FDNUM) - partials(NESTED_FDNUM2)) - @test NESTED_FDNUM - PRIMAL === Dual{TestTag}(value(NESTED_FDNUM) - PRIMAL, partials(NESTED_FDNUM)) - @test PRIMAL - NESTED_FDNUM === Dual{TestTag}(PRIMAL - value(NESTED_FDNUM), -(partials(NESTED_FDNUM))) - @test -(NESTED_FDNUM) === Dual{TestTag}(-(value(NESTED_FDNUM)), -(partials(NESTED_FDNUM))) + @test NESTED_FDNUM - NESTED_FDNUM2 === Dual{TestTag}(value(TestTag, NESTED_FDNUM) - value(TestTag, NESTED_FDNUM2), partials(TestTag, NESTED_FDNUM) - partials(TestTag, NESTED_FDNUM2)) + @test NESTED_FDNUM - PRIMAL === Dual{TestTag}(value(TestTag, NESTED_FDNUM) - PRIMAL, partials(TestTag, NESTED_FDNUM)) + @test PRIMAL - NESTED_FDNUM === Dual{TestTag}(PRIMAL - value(TestTag, NESTED_FDNUM), -(partials(TestTag, NESTED_FDNUM))) + @test -(NESTED_FDNUM) === Dual{TestTag}(-(value(TestTag, NESTED_FDNUM)), -(partials(TestTag, NESTED_FDNUM))) # Multiplication # #----------------# - @test FDNUM * FDNUM2 === Dual{TestTag}(value(FDNUM) * value(FDNUM2), ForwardDiff._mul_partials(partials(FDNUM), partials(FDNUM2), value(FDNUM2), value(FDNUM))) - @test FDNUM * PRIMAL === Dual{TestTag}(value(FDNUM) * PRIMAL, partials(FDNUM) * PRIMAL) - @test PRIMAL * FDNUM === Dual{TestTag}(value(FDNUM) * PRIMAL, partials(FDNUM) * PRIMAL) + @test FDNUM * FDNUM2 === Dual{TestTag}(value(TestTag, FDNUM) * value(TestTag, FDNUM2), ForwardDiff._mul_partials(partials(TestTag, FDNUM), partials(TestTag, FDNUM2), value(TestTag, FDNUM2), value(TestTag, FDNUM))) + @test FDNUM * PRIMAL === Dual{TestTag}(value(TestTag, FDNUM) * PRIMAL, partials(TestTag, FDNUM) * PRIMAL) + @test PRIMAL * FDNUM === Dual{TestTag}(value(TestTag, FDNUM) * PRIMAL, partials(TestTag, FDNUM) * PRIMAL) - @test NESTED_FDNUM * NESTED_FDNUM2 === Dual{TestTag}(value(NESTED_FDNUM) * value(NESTED_FDNUM2), ForwardDiff._mul_partials(partials(NESTED_FDNUM), partials(NESTED_FDNUM2), value(NESTED_FDNUM2), value(NESTED_FDNUM))) - @test NESTED_FDNUM * PRIMAL === Dual{TestTag}(value(NESTED_FDNUM) * PRIMAL, partials(NESTED_FDNUM) * PRIMAL) - @test PRIMAL * NESTED_FDNUM === Dual{TestTag}(value(NESTED_FDNUM) * PRIMAL, partials(NESTED_FDNUM) * PRIMAL) + @test NESTED_FDNUM * NESTED_FDNUM2 === Dual{TestTag}(value(TestTag, NESTED_FDNUM) * value(TestTag, NESTED_FDNUM2), ForwardDiff._mul_partials(partials(TestTag, NESTED_FDNUM), partials(TestTag, NESTED_FDNUM2), value(TestTag, NESTED_FDNUM2), value(TestTag, NESTED_FDNUM))) + @test NESTED_FDNUM * PRIMAL === Dual{TestTag}(value(TestTag, NESTED_FDNUM) * PRIMAL, partials(TestTag, NESTED_FDNUM) * PRIMAL) + @test PRIMAL * NESTED_FDNUM === Dual{TestTag}(value(TestTag, NESTED_FDNUM) * PRIMAL, partials(TestTag, NESTED_FDNUM) * PRIMAL) # Division # #----------# @@ -491,21 +501,21 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false if M > 0 && N > 0 # Recall that FDNUM = Dual{TestTag}(PRIMAL, PARTIALS) has N partials, # all random numbers nonzero, and FDNUM2 another draw. M only affects NESTED_FDNUM. - @test Dual{1}(FDNUM) / Dual{1}(PRIMAL) === Dual{1}(FDNUM / PRIMAL) - @test Dual{1}(PRIMAL) / Dual{1}(FDNUM) === Dual{1}(PRIMAL / FDNUM) - @test Dual{1}(FDNUM) / FDNUM2 === Dual{1}(FDNUM / FDNUM2) - @test FDNUM / Dual{1}(FDNUM2) === Dual{1}(FDNUM / FDNUM2) + @test Dual{OuterTestTag}(FDNUM) / Dual{OuterTestTag}(PRIMAL) === Dual{OuterTestTag}(FDNUM / PRIMAL) + @test Dual{OuterTestTag}(PRIMAL) / Dual{OuterTestTag}(FDNUM) === Dual{OuterTestTag}(PRIMAL / FDNUM) + @test Dual{OuterTestTag}(FDNUM) / FDNUM2 === Dual{OuterTestTag}(FDNUM / FDNUM2) + @test FDNUM / Dual{OuterTestTag}(FDNUM2) === Dual{OuterTestTag}(FDNUM / FDNUM2) # following may not be exact, see #264 - @test Dual{1}(FDNUM / PRIMAL, FDNUM2 / PRIMAL) ≈ Dual{1}(FDNUM, FDNUM2) / PRIMAL + @test Dual{OuterTestTag}(FDNUM / PRIMAL, FDNUM2 / PRIMAL) ≈ Dual{OuterTestTag}(FDNUM, FDNUM2) / PRIMAL end - @test dual_isapprox(FDNUM / FDNUM2, Dual{TestTag}(value(FDNUM) / value(FDNUM2), ForwardDiff._div_partials(partials(FDNUM), partials(FDNUM2), value(FDNUM), value(FDNUM2)))) - @test dual_isapprox(FDNUM / PRIMAL, Dual{TestTag}(value(FDNUM) / PRIMAL, partials(FDNUM) / PRIMAL)) - @test dual_isapprox(PRIMAL / FDNUM, Dual{TestTag}(PRIMAL / value(FDNUM), (-(PRIMAL) / value(FDNUM)^2) * partials(FDNUM))) + @test dual_isapprox(FDNUM / FDNUM2, Dual{TestTag}(value(TestTag, FDNUM) / value(TestTag, FDNUM2), ForwardDiff._div_partials(partials(TestTag, FDNUM), partials(TestTag, FDNUM2), value(TestTag, FDNUM), value(TestTag, FDNUM2)))) + @test dual_isapprox(FDNUM / PRIMAL, Dual{TestTag}(value(TestTag, FDNUM) / PRIMAL, partials(TestTag, FDNUM) / PRIMAL)) + @test dual_isapprox(PRIMAL / FDNUM, Dual{TestTag}(PRIMAL / value(TestTag, FDNUM), (-(PRIMAL) / value(TestTag, FDNUM)^2) * partials(TestTag, FDNUM))) - @test dual_isapprox(NESTED_FDNUM / NESTED_FDNUM2, Dual{TestTag}(value(NESTED_FDNUM) / value(NESTED_FDNUM2), ForwardDiff._div_partials(partials(NESTED_FDNUM), partials(NESTED_FDNUM2), value(NESTED_FDNUM), value(NESTED_FDNUM2)))) - @test dual_isapprox(NESTED_FDNUM / PRIMAL, Dual{TestTag}(value(NESTED_FDNUM) / PRIMAL, partials(NESTED_FDNUM) / PRIMAL)) - @test dual_isapprox(PRIMAL / NESTED_FDNUM, Dual{TestTag}(PRIMAL / value(NESTED_FDNUM), (-(PRIMAL) / value(NESTED_FDNUM)^2) * partials(NESTED_FDNUM))) + @test dual_isapprox(NESTED_FDNUM / NESTED_FDNUM2, Dual{TestTag}(value(TestTag, NESTED_FDNUM) / value(TestTag, NESTED_FDNUM2), ForwardDiff._div_partials(partials(TestTag, NESTED_FDNUM), partials(TestTag, NESTED_FDNUM2), value(TestTag, NESTED_FDNUM), value(TestTag, NESTED_FDNUM2)))) + @test dual_isapprox(NESTED_FDNUM / PRIMAL, Dual{TestTag}(value(TestTag, NESTED_FDNUM) / PRIMAL, partials(TestTag, NESTED_FDNUM) / PRIMAL)) + @test dual_isapprox(PRIMAL / NESTED_FDNUM, Dual{TestTag}(PRIMAL / value(TestTag, NESTED_FDNUM), (-(PRIMAL) / value(TestTag, NESTED_FDNUM)^2) * partials(TestTag, NESTED_FDNUM))) # Exponentiation # #----------------# @@ -521,14 +531,14 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test dual_isapprox(1.0 * NESTED_FDNUM^PRIMAL, exp(PRIMAL * log(NESTED_FDNUM))) @test dual_isapprox(1.0 * PRIMAL^NESTED_FDNUM, exp(NESTED_FDNUM * log(PRIMAL))) - @test partials(NaNMath.pow(Dual{TestTag}(-2.0, 1.0), Dual{TestTag}(2.0, 0.0)), 1) == -4.0 + @test partials(TestTag, NaNMath.pow(Dual{TestTag}(-2.0, 1.0), Dual{TestTag}(2.0, 0.0)), 1) == -4.0 # differentiating must not widen the primal: `2^x` is Float32 for a Float32 `x` @testset "$f: real base keeps $W exponent" for f in (^, NaNMath.pow), W in (Float16, Float32, Float64) w = W(4)/W(3) - @test typeof(value(f(2, Dual{TestTag}(w, one(W))))) === typeof(f(2, w)) - @test typeof(value(f(2.0f0, Dual{TestTag}(w, one(W))))) === typeof(f(2.0f0, w)) + @test typeof(value(TestTag, f(2, Dual{TestTag}(w, one(W))))) === typeof(f(2, w)) + @test typeof(value(TestTag, f(2.0f0, Dual{TestTag}(w, one(W))))) === typeof(f(2.0f0, w)) end ################################### @@ -568,14 +578,14 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false actualval = $M.$f(x)::Union{Real,Complex} if actualval isa Real @test dx isa Dual{TestTag} - @test value(dx) == actualval - @test partials(dx, 1) == $deriv + @test value(TestTag, dx) == actualval + @test partials(TestTag, dx, 1) == $deriv else @test dx isa Complex{<:Dual{TestTag}} - @test value(real(dx)) == real(actualval) - @test value(imag(dx)) == imag(actualval) - @test partials(real(dx), 1) == real($deriv) - @test partials(imag(dx), 1) == imag($deriv) + @test value(TestTag, real(dx)) == real(actualval) + @test value(TestTag, imag(dx)) == imag(actualval) + @test partials(TestTag, real(dx), 1) == real($deriv) + @test partials(TestTag, imag(dx), 1) == imag($deriv) end end elseif arity == 2 @@ -597,25 +607,25 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false if actualval isa Real @test dx isa Dual{TestTag} @test dy isa Dual{TestTag} - @test value(dx) == actualval - @test value(dy) == actualval - @test partials(dx, 1) ≈ actualdx nans=true - @test partials(dy, 1) ≈ actualdy nans=true + @test value(TestTag, dx) == actualval + @test value(TestTag, dy) == actualval + @test partials(TestTag, dx, 1) ≈ actualdx nans=true + @test partials(TestTag, dy, 1) ≈ actualdy nans=true else @test dx isa Complex{<:Dual{TestTag}} @test dy isa Complex{<:Dual{TestTag}} - # @test real(value(dx)) == real(actualval) - # @test real(value(dy)) == real(actualval) - # @test imag(value(dx)) == imag(actualval) - # @test imag(value(dy)) == imag(actualval) - @test value(real(dx)) == real(actualval) - @test value(real(dy)) == real(actualval) - @test value(imag(dx)) == imag(actualval) - @test value(imag(dy)) == imag(actualval) - @test partials(real(dx), 1) ≈ real(actualdx) nans=true - @test partials(real(dy), 1) ≈ real(actualdy) nans=true - @test partials(imag(dx), 1) ≈ imag(actualdx) nans=true - @test partials(imag(dy), 1) ≈ imag(actualdy) nans=true + # @test real(value(TestTag, dx)) == real(actualval) + # @test real(value(TestTag, dy)) == real(actualval) + # @test imag(value(TestTag, dx)) == imag(actualval) + # @test imag(value(TestTag, dy)) == imag(actualval) + @test value(TestTag, real(dx)) == real(actualval) + @test value(TestTag, real(dy)) == real(actualval) + @test value(TestTag, imag(dx)) == imag(actualval) + @test value(TestTag, imag(dy)) == imag(actualval) + @test partials(TestTag, real(dx), 1) ≈ real(actualdx) nans=true + @test partials(TestTag, real(dy), 1) ≈ real(actualdy) nans=true + @test partials(TestTag, imag(dx), 1) ≈ imag(actualdx) nans=true + @test partials(TestTag, imag(dy), 1) ≈ imag(actualdy) nans=true end end end @@ -666,9 +676,9 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false tol = V === Float32 ? 5f-4 : 1e-5 tolval = tol^(one(tol) / 2^(isempty(ind) ? 0 : first(ind))) for i in 1:2 - @test value(pq[i]) ≈ gamma_inc(a, 1 + PRIMAL, ind...)[i] rtol=tolval + @test value(TestTag, pq[i]) ≈ gamma_inc(a, 1 + PRIMAL, ind...)[i] rtol=tolval der = Calculus.derivative(x -> gamma_inc(Float64(a), x, 0)[i], Float64(1 + PRIMAL)) - @test partials(pq[i]) ≈ PARTIALS * der rtol=tol + @test partials(TestTag, pq[i]) ≈ PARTIALS * der rtol=tol end end end @@ -679,16 +689,16 @@ end @testset "Exponentiation of zero" begin x0 = 0.0 - x1 = Dual{:t1}(x0, 1.0) - x2 = Dual{:t2}(x1, 1.0) - x3 = Dual{:t3}(x2, 1.0) + x1 = Dual{TestTag}(x0, 1.0) + x2 = Dual{OuterTestTag}(x1, 1.0) + x3 = Dual{OutermostTestTag}(x2, 1.0) pow = ^ # to call non-literal power @test pow(x3, 2) === x3^2 === x3 * x3 @test pow(x2, 1) === x2^1 === x2 - @test pow(x1, 0) === x1^0 === Dual{:t1}(1.0, 0.0) + @test pow(x1, 0) === x1^0 === Dual{TestTag}(1.0, 0.0) y = Dual{TestTag}(1.0, 0.0, 1.0); x = Dual{OuterTestTag}(0*y, 0*y); - @test iszero(ForwardDiff.partials(ForwardDiff.partials(x^y)[1])) + @test iszero(partials(TestTag, partials(OuterTestTag, x^y, 1))) end @testset "Type min/max" begin @@ -736,7 +746,7 @@ end @testset "float" begin # issue #492 @test float(Dual{Nothing, Int, 2}) === Dual{Nothing, Float64, 2} @test float(Dual(1)) isa Dual{Nothing, Float64, 0} - @test value.(float.(Dual.(1:4, 2:5, 3:6))) isa Vector{Float64} + @test value.(Nothing, float.(Dual.(1:4, 2:5, 3:6))) isa Vector{Float64} @test ForwardDiff.derivative(float, 1)::Float64 === 1.0 end @@ -761,10 +771,10 @@ end for (i, yi, yduali) in zip(1:3, y, ydual) # Primal values must match `LinearAlgebra.givensAlgorithm` with `Float64` inputs - @test ForwardDiff.value(yduali) ≈ yi + @test ForwardDiff.value(TestTag, yduali) ≈ yi # Partial derivatives must be zero (zero in - zero out) - @test iszero(ForwardDiff.partials(yduali)) + @test iszero(ForwardDiff.partials(TestTag, yduali)) end end end diff --git a/test/GradientTest.jl b/test/GradientTest.jl index bf121239..bc4a2a97 100644 --- a/test/GradientTest.jl +++ b/test/GradientTest.jl @@ -259,7 +259,7 @@ end npartials = Ref(0) y = ForwardDiff.gradient(T(randn(n, n))) do x fevals[] += 1 - npartials[] += ForwardDiff.npartials(eltype(x)) + npartials[] += ForwardDiff.npartials(ForwardDiff.tagtype(eltype(x)), eltype(x)) return sum(x) end if npartials[] <= ForwardDiff.DEFAULT_CHUNK_THRESHOLD @@ -278,8 +278,8 @@ end # issue #769 @testset "functions with `Dual` output" begin x = [Dual{OuterTestTag}(Dual{TestTag}(1.3, 2.1), Dual{TestTag}(0.3, -2.4))] - f(x) = sum(ForwardDiff.value, x) - der = ForwardDiff.derivative(ForwardDiff.value, only(x)) + f(x) = sum(Base.Fix1(ForwardDiff.value, OuterTestTag), x) + der = ForwardDiff.derivative(Base.Fix1(ForwardDiff.value, OuterTestTag), only(x)) # Vector mode grad = ForwardDiff.gradient(f, x) diff --git a/test/JacobianTest.jl b/test/JacobianTest.jl index 5050e23b..034539de 100644 --- a/test/JacobianTest.jl +++ b/test/JacobianTest.jl @@ -401,8 +401,8 @@ end # issue #769 @testset "functions with `Dual` output" begin x = [Dual{OuterTestTag}(Dual{TestTag}(1.3, 2.1), Dual{TestTag}(0.3, -2.4))] - f(x) = map(ForwardDiff.value, x) - der = ForwardDiff.derivative(ForwardDiff.value, only(x)) + f(x) = map(Base.Fix1(ForwardDiff.value, OuterTestTag), x) + der = ForwardDiff.derivative(Base.Fix1(ForwardDiff.value, OuterTestTag), only(x)) # Vector mode jac = ForwardDiff.jacobian(f, x) diff --git a/test/MiscTest.jl b/test/MiscTest.jl index e8068e95..8328461a 100644 --- a/test/MiscTest.jl +++ b/test/MiscTest.jl @@ -128,7 +128,7 @@ end # NaNs # #------# -@test ForwardDiff.partials(NaNMath.pow(ForwardDiff.Dual(-2.0,1.0),ForwardDiff.Dual(2.0,0.0)),1) == -4.0 +@test ForwardDiff.partials(Nothing, NaNMath.pow(ForwardDiff.Dual(-2.0,1.0),ForwardDiff.Dual(2.0,0.0)),1) == -4.0 # Partials{0} # #-------------# diff --git a/test/SeedTest.jl b/test/SeedTest.jl index 02b821c3..06c54d9b 100644 --- a/test/SeedTest.jl +++ b/test/SeedTest.jl @@ -27,12 +27,12 @@ const SEED_CASES = ( # Positions within `sidx` whose partials are zero. zeroed_positions(duals, sidx) = - [i for (i, idx) in enumerate(sidx) if iszero(ForwardDiff.partials(duals[idx]))] + [i for (i, idx) in enumerate(sidx) if iszero(ForwardDiff.partials(Nothing, duals[idx]))] # Compares over *every* index of `x`, not just the structural ones, so a bug misplacing values # outside the structural set is visible. Off-structure reads are safe: the wrapper types return # `zero(Dual)` without touching the (uninitialized) parent storage. -values_match(duals, x) = all(idx -> ForwardDiff.value(duals[idx]) == x[idx], eachindex(x)) +values_match(duals, x) = all(idx -> ForwardDiff.value(Nothing, duals[idx]) == x[idx], eachindex(x)) function fill_marker!(duals, x, sidx, marker) D = eltype(duals) @@ -45,7 +45,7 @@ end @testset "seed_zero_partials!: $(nameof(typeof(x)))" for (x, sidx) in SEED_CASES cfg = ForwardDiff.GradientConfig(nothing, x, ForwardDiff.Chunk{3}()) duals, seeds = cfg.duals, cfg.seeds - N = ForwardDiff.npartials(eltype(duals)) + N = ForwardDiff.npartials(Nothing, eltype(duals)) marker = Partials(ntuple(i -> Float64(i), N)) nstruct = length(sidx) From b46279aba7fff2d80c5f2cee0091890e2d24fe32 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Thu, 24 Sep 2026 11:06:42 +0200 Subject: [PATCH 3/8] Order tags structurally instead of by definition count MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Replace the impure `tagcount` counter with a total order on tag types that depends only on their structure, so it is the same in every session and unaffected by precompilation. A generated function computes a constant `(rank, key)` ID per tag, and `≺` compares the IDs, which is evaluated at compile time. The rank strictly increases from a type to any type containing it, so every tag is greater than the tags in its parameters and seeding outermost keeps nested `Dual`s sorted. The `Dual` constructor rejects a tag stored outside a greater tag. Since every pair of tags is now ordered, `DualMismatchError` is removed. Co-Authored-By: Claude Opus 5.5 (1M context) --- src/config.jl | 46 ++++++++++++++++++++++++++++++------------- src/dual.jl | 29 ++++++++++++++++----------- test/ConfusionTest.jl | 39 ++++++++++++++++++++++++++++++++++++ test/GradientTest.jl | 2 +- test/JacobianTest.jl | 2 +- 5 files changed, 90 insertions(+), 28 deletions(-) diff --git a/src/config.jl b/src/config.jl index 3c6c97e3..4ef6de80 100644 --- a/src/config.jl +++ b/src/config.jl @@ -5,25 +5,43 @@ struct Tag{F,V} end -const TAGCOUNT = Threads.Atomic{UInt}(0) +Tag(f::F, ::Type{V}) where {F,V} = Tag{F,V}() -# each tag is assigned a unique number -# tags which depend on other tags will be larger -@generated function tagcount(::Type{Tag{F,V}}) where {F,V} - :($(Threads.atomic_add!(TAGCOUNT, UInt(1)))) -end +Tag(::Nothing, ::Type{V}) where {V} = nothing -function Tag(f::F, ::Type{V}) where {F,V} - tagcount(Tag{F,V}) # trigger generated function - Tag{F,V}() +# Encodes a type (or type parameter) as a sequence of strings that depends only on its +# structure. Distinct objects have distinct keys, and the key of a parameter is a strict +# subsequence of the key of the type. +function typekey!(key::Vector{String}, x::DataType) + push!(key, "T", string(fullname(parentmodule(x))), String(nameof(x)), string(length(x.parameters))) + foreach(p -> typekey!(key, p), x.parameters) + return key +end +typekey!(key::Vector{String}, x::Symbol) = push!(key, "S", String(x)) +function typekey!(key::Vector{String}, x) + typekey!(push!(key, "V"), typeof(x)) + if isprimitivetype(typeof(x)) + push!(key, bytes2hex(reinterpret(UInt8, [x]))) + else + for i in 1:nfields(x) + if isdefined(x, i) + typekey!(key, getfield(x, i)) + else + push!(key, "#undef") + end + end + end + return key end -Tag(::Nothing, ::Type{V}) where {V} = nothing - +# `A ≺ B` compares `(rank, key)` lexicographically. The rank strictly increases from a type to +# any type containing it, so every tag is greater than the tags occurring in its parameters. +@generated tagid(::Type{T}) where {T} = (key = typekey!(String[], T); (length(key), key...)) -@inline function ≺(::Type{Tag{F1,V1}}, ::Type{Tag{F2,V2}}) where {F1,V1,F2,V2} - tagcount(Tag{F1,V1}) < tagcount(Tag{F2,V2}) -end +# Nested `Dual`s store the greatest tag outermost. A tag `Tag{F,V}` is greater than all tags +# in the type `V` of its input and in the type `F` of its function, so seeding it outermost +# keeps nested `Dual`s sorted. The comparison of the constant IDs is evaluated at compile time. +≺(::Type{A}, ::Type{B}) where {A,B} = isless(tagid(A), tagid(B)) struct InvalidTagException{E,O} <: Exception end diff --git a/src/dual.jl b/src/dual.jl index 7bd5dc5a..e419d4f6 100644 --- a/src/dual.jl +++ b/src/dual.jl @@ -16,6 +16,7 @@ struct Dual{T,V,N} <: Real partials::Partials{N,V} function Dual{T, V, N}(value::V, partials::Partials{N, V}) where {T, V, N} T isa Type || throw_invalid_tag(T) + check_tag_order(T, V) can_dual(V) || throw_cannot_dual(V) new{T, V, N}(value, partials) end @@ -30,14 +31,6 @@ Base.ArithmeticStyle(::Type{<:Dual{T,V}}) where {T,V} = Base.ArithmeticStyle(V) # Exceptions # ############## -struct DualMismatchError{A,B} <: Exception - a::A - b::B -end - -Base.showerror(io::IO, e::DualMismatchError{A,B}) where {A,B} = - print(io, "Cannot determine ordering of Dual tags ", e.a, " and ", e.b) - @noinline function throw_invalid_tag(T) throw(ArgumentError(lazy"The tag of a Dual must be a type, got $(repr(T)).")) end @@ -49,13 +42,25 @@ end """ ForwardDiff.≺(a, b)::Bool -Determines the order in which tagged `Dual` objects are composed. If true, then `Dual{b}` +Determines the order in which tagged `Dual` objects are stored. If true, then `Dual{b}` objects will appear outside `Dual{a}` objects. -This is important when working with nested differentiation: currently, only the outermost -tag can be extracted, so it should be used in the _innermost_ function. +Values and partials with respect to a tag can be extracted irrespective of this order. """ -≺(a,b) = throw(DualMismatchError(a,b)) +function ≺ end + +@noinline function throw_tag_order(T, S) + throw(ArgumentError(lazy"Cannot store a Dual with tag $T outside a Dual with tag $S, since $T ≺ $S.")) +end + +# Every `Dual` stores its greatest tag outermost, so it suffices to compare with the next layer +@inline check_tag_order(::Type{T}, ::Type) where {T} = nothing +@inline function check_tag_order(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} + if S !== T && T ≺ S + throw_tag_order(T, S) + end + return nothing +end ################ # Constructors # diff --git a/test/ConfusionTest.jl b/test/ConfusionTest.jl index 13c62ae9..64f6c5f4 100644 --- a/test/ConfusionTest.jl +++ b/test/ConfusionTest.jl @@ -72,5 +72,44 @@ end end end == 0.0 +# Nested derivatives whose tags are not ordered by containment +const captured = Ref{Any}() +const inner = Ref{Any}() +struct AFunction end +struct BFunction end +(::AFunction)(x) = captured[] * x^2 +(::BFunction)(x) = captured[] * x^2 +struct ADerivative end +struct BDerivative end +(::ADerivative)(x) = (captured[] = x; D(inner[], 2.0)) +(::BDerivative)(x) = (captured[] = x; D(inner[], 2.0)) +struct AGradient end +struct BGradient end +(::AGradient)(x) = (captured[] = prod(x); D(inner[], 2.0)) +(::BGradient)(x) = (captured[] = prod(x); D(inner[], 2.0)) +for (outer, f) in ((ADerivative(), BFunction()), (BDerivative(), AFunction())) + inner[] = f + @test D(outer, 3.0) == 4.0 +end +for (outer, f) in ((AGradient(), BFunction()), (BGradient(), AFunction())) + inner[] = f + @test ForwardDiff.gradient(outer, [3.0, 2.0]) == [8.0, 12.0] + @test ForwardDiff.hessian(outer, [3.0, 2.0]) == [0.0 4.0; 4.0 0.0] +end + +# Every tag is greater than the tags in its input type and its function type +let T = ForwardDiff.Tag{AFunction,Float64}, d = ForwardDiff.Dual{T}(1.0, 1.0), f = x -> d * x + for S in (ForwardDiff.Tag{BFunction,typeof(d)}, ForwardDiff.Tag{typeof(f),Float64}) + @test ForwardDiff.:≺(T, S) + @test !ForwardDiff.:≺(S, T) + end +end + +# Nested `Dual`s store the greatest tag outermost +struct ATag end +struct BTag end +@test ForwardDiff.Dual{BTag}(ForwardDiff.Dual{ATag}(1.0, 2.0), ForwardDiff.Dual{ATag}(3.0, 4.0)) isa ForwardDiff.Dual{BTag} +@test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag}(ForwardDiff.Dual{BTag}(1.0, 2.0), ForwardDiff.Dual{BTag}(3.0, 4.0)) + end # module diff --git a/test/GradientTest.jl b/test/GradientTest.jl index bc4a2a97..c84f9cd2 100644 --- a/test/GradientTest.jl +++ b/test/GradientTest.jl @@ -16,7 +16,7 @@ include(joinpath(dirname(@__FILE__), "utils.jl")) struct TestTag end struct OuterTestTag end ForwardDiff.:≺(::Type{TestTag}, ::Type{OuterTestTag}) = true -ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{<:Tag}) = true +ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false ################## # hardcoded test # diff --git a/test/JacobianTest.jl b/test/JacobianTest.jl index 034539de..9db1f1b0 100644 --- a/test/JacobianTest.jl +++ b/test/JacobianTest.jl @@ -14,7 +14,7 @@ include(joinpath(dirname(@__FILE__), "utils.jl")) struct TestTag end struct OuterTestTag end ForwardDiff.:≺(::Type{TestTag}, ::Type{OuterTestTag}) = true -ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{<:Tag}) = true +ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false ################## # hardcoded test # From da257339e41374f976cfa16dece5f4952c04bce8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Thu, 24 Sep 2026 23:04:33 +0200 Subject: [PATCH 4/8] Require the tags of nested `Dual`s to be unique Nesting a `Dual` in a `Dual` with the same tag merges two perturbations, so the `Dual` constructor now requires every tag to be strictly greater than the tags of its value and partials. It checks their runtime types, so tags hidden behind an abstract value type are rejected as well. Hence `HessianConfig` uses a separate tag `Tag{F,Dual{T,V,N}}` for its gradient layer (#845), which adds a type parameter, and so does `hessian!` for StaticArrays with immutable results. The latter now reads the gradient from the gradient layer, as all other Hessian paths do, so that all paths agree with the Jacobian of the gradient. Untagged `Dual`s and configs created with `f = nothing` use the tag `Tag{Nothing,V}` instead of `Nothing`, where `V` is the promoted value type, so that nesting them yields distinct tags. Promoting `Dual`s with the same tag but different numbers of partials throws an error instead of nesting the tag in itself. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/src/user/advanced.md | 4 +- ext/ForwardDiffStaticArraysExt.jl | 15 +-- src/config.jl | 17 +-- src/dual.jl | 57 ++++++---- test/AllocationsTest.jl | 4 +- test/ConfusionTest.jl | 45 ++++++++ test/DualTest.jl | 167 +++++++++++++++--------------- test/MiscTest.jl | 2 +- test/SeedTest.jl | 6 +- 9 files changed, 195 insertions(+), 122 deletions(-) diff --git a/docs/src/user/advanced.md b/docs/src/user/advanced.md index 6955ff1c..506b8f9d 100644 --- a/docs/src/user/advanced.md +++ b/docs/src/user/advanced.md @@ -129,7 +129,7 @@ aren't sensitive to the input and thus cause ForwardDiff to incorrectly return ` # the dual number's perturbation component is zero, so this # variable should not propagate derivative information julia> log(ForwardDiff.Dual(0.0, 0.0)) -Dual{Nothing}(-Inf,NaN) # oops, this NaN should be 0.0 +Dual{ForwardDiff.Tag{Nothing, Float64}}(-Inf,NaN) # oops, this NaN should be 0.0 ``` Here, ForwardDiff computes the derivative of `log(0.0)` as `NaN` and then propagates @@ -167,7 +167,7 @@ julia> set_preferences!(UUID("f6369f11-7733-5829-9624-2563aa707210"), "nansafe_m julia> using ForwardDiff julia> log(ForwardDiff.Dual(0.0, 0.0)) -Dual{Nothing}(-Inf,0.0) +Dual{ForwardDiff.Tag{Nothing, Float64}}(-Inf,0.0) ``` In the future, we plan on allowing users and downstream library authors to dynamically diff --git a/ext/ForwardDiffStaticArraysExt.jl b/ext/ForwardDiffStaticArraysExt.jl index 796bdf5c..575fc8a3 100644 --- a/ext/ForwardDiffStaticArraysExt.jl +++ b/ext/ForwardDiffStaticArraysExt.jl @@ -132,13 +132,16 @@ ForwardDiff.hessian!(result::ImmutableDiffResult, f::F, x::StaticArray, cfg::Hes ForwardDiff.hessian!(result::ImmutableDiffResult, f::F, x::StaticArray, cfg::HessianConfig, ::Val) where {F} = hessian!(result, f, x) function ForwardDiff.hessian!(result::ImmutableDiffResult, f::F, x::StaticArray) where {F} - T = typeof(Tag(f, eltype(x))) - d1 = dualize(T, x) - d2 = dualize(T, d1) + T1 = typeof(Tag(f, eltype(x))) + d1 = dualize(T1, x) + T2 = typeof(Tag(f, eltype(d1))) + d2 = dualize(T2, d1) fd2 = f(d2) - val = value(T,value(T,fd2)) - grad = extract_gradient(T,value(T,fd2), x) - hess = extract_jacobian(T,partials(T,fd2), x) + # Hessian = Jacobian (w.r.t. `T1`) of the gradient (w.r.t. `T2`), as in `hessian!` for other arrays + ∇fd2 = extract_gradient(T2, fd2, d1) + val = value(T1, value(T2, fd2)) + grad = map(Base.Fix1(value, T1), ∇fd2) + hess = extract_jacobian(T1, ∇fd2, x) result = DiffResults.hessian!(result, hess) result = DiffResults.gradient!(result, grad) result = DiffResults.value!(result, val) diff --git a/src/config.jl b/src/config.jl index 4ef6de80..f50762e4 100644 --- a/src/config.jl +++ b/src/config.jl @@ -7,7 +7,7 @@ end Tag(f::F, ::Type{V}) where {F,V} = Tag{F,V}() -Tag(::Nothing, ::Type{V}) where {V} = nothing +Tag(::Nothing, ::Type{V}) where {V} = Tag{Nothing,V}() # Encodes a type (or type parameter) as a sequence of strings that depends only on its # structure. Distinct objects have distinct keys, and the key of a parameter is a strict @@ -57,6 +57,9 @@ checktag(::Type{Tag{F,V}}, f::F, x::AbstractArray{V}) where {F,V} = true # no easy way to check Jacobian tag used with Hessians as multiple functions may be used checktag(::Type{Tag{FT,VT}}, f::F, x::AbstractArray{V}) where {FT<:Tuple,VT,F,V} = true +# tag of `nothing` configs, for any function +checktag(::Type{Tag{FT,VT}}, f::F, x::AbstractArray{V}) where {FT<:Nothing,VT,F,V} = true + # custom tag: you're on your own. checktag(z, f, x) = true @@ -213,9 +216,9 @@ Base.eltype(::Type{JacobianConfig{T,V,N,D}}) where {T,V,N,D} = Dual{T,V,N} # HessianConfig # ################# -struct HessianConfig{T,V,N,DG,DJ} <: AbstractConfig{N} +struct HessianConfig{T,V,N,DG,DJ,TG} <: AbstractConfig{N} jacobian_config::JacobianConfig{T,V,N,DJ} - gradient_config::GradientConfig{T,Dual{T,V,N},N,DG} + gradient_config::GradientConfig{TG,Dual{T,V,N},N,DG} end """ @@ -241,7 +244,7 @@ function HessianConfig(f::F, chunk::Chunk = Chunk(x), tag = Tag(f, V)) where {F,V} jacobian_config = JacobianConfig(f, x, chunk, tag) - gradient_config = GradientConfig(f, jacobian_config.duals, chunk, tag) + gradient_config = GradientConfig(f, jacobian_config.duals, chunk, Tag{F,eltype(jacobian_config)}()) return HessianConfig(jacobian_config, gradient_config) end @@ -266,10 +269,10 @@ function HessianConfig(f::F, chunk::Chunk = Chunk(x), tag = Tag(f, V)) where {F,V} jacobian_config = JacobianConfig((f,gradient), DiffResults.gradient(result), x, chunk, tag) - gradient_config = GradientConfig(f, jacobian_config.duals[2], chunk, tag) + gradient_config = GradientConfig(f, jacobian_config.duals[2], chunk, Tag{F,eltype(jacobian_config)}()) return HessianConfig(jacobian_config, gradient_config) end checktag(::HessianConfig{T},f,x) where {T} = checktag(T,f,x) -Base.eltype(::Type{HessianConfig{T,V,N,DG,DJ}}) where {T,V,N,DG,DJ} = - Dual{T,Dual{T,V,N},N} +Base.eltype(::Type{HessianConfig{T,V,N,DG,DJ,TG}}) where {T,V,N,DG,DJ,TG} = + Dual{TG,Dual{T,V,N},N} diff --git a/src/dual.jl b/src/dual.jl index e419d4f6..019b03a3 100644 --- a/src/dual.jl +++ b/src/dual.jl @@ -16,7 +16,8 @@ struct Dual{T,V,N} <: Real partials::Partials{N,V} function Dual{T, V, N}(value::V, partials::Partials{N, V}) where {T, V, N} T isa Type || throw_invalid_tag(T) - check_tag_order(T, V) + check_tag_order(T, value) + foreach(p -> check_tag_order(T, p), partials.values) can_dual(V) || throw_cannot_dual(V) new{T, V, N}(value, partials) end @@ -53,10 +54,16 @@ function ≺ end throw(ArgumentError(lazy"Cannot store a Dual with tag $T outside a Dual with tag $S, since $T ≺ $S.")) end -# Every `Dual` stores its greatest tag outermost, so it suffices to compare with the next layer -@inline check_tag_order(::Type{T}, ::Type) where {T} = nothing -@inline function check_tag_order(::Type{T}, ::Type{Dual{S,V,N}}) where {T,S,V,N} - if S !== T && T ≺ S +@noinline function throw_same_tag(T) + throw(ArgumentError(lazy"Cannot store a Dual with tag $T inside a Dual with the same tag.")) +end + +# Tags are strictly decreasing inwards, hence unique, so it suffices to check the next layer +@inline check_tag_order(::Type{T}, x) where {T} = nothing +@inline function check_tag_order(::Type{T}, ::Dual{S}) where {T,S} + if S === T + throw_same_tag(T) + elseif T ≺ S throw_tag_order(T, S) end return nothing @@ -66,20 +73,27 @@ end # Constructors # ################ -@inline Dual{T}(value::V, partials::Partials{N,V}) where {T,N,V} = Dual{T,V,N}(value, partials) - -@inline function Dual{T}(value::A, partials::Partials{N,B}) where {T,N,A,B} +# Converts the arguments of the constructors to a value and partials of the same type +@inline dual_args(value::V, partials::Partials{N,V}) where {N,V} = (value, partials) +@inline function dual_args(value::A, partials::Partials{N,B}) where {N,A,B} C = promote_type(A, B) - return Dual{T}(convert(C, value), convert(Partials{N,C}, partials)) + return (convert(C, value), convert(Partials{N,C}, partials)) +end +@inline dual_args(value, partials::Tuple) = dual_args(value, Partials(partials)) +@inline dual_args(value, partials::Tuple{}) = dual_args(value, Partials{0,typeof(value)}(partials)) +@inline dual_args(value) = dual_args(value, ()) +@inline dual_args(value, partial1, partials...) = dual_args(value, tuple(partial1, partials...)) +@inline dual_args(value::V, ::Chunk{N}, p::Val{i}) where {V,N,i} = dual_args(value, single_seed(Partials{N,V}, p)) + +@inline function Dual{T}(args...) where {T} + value, partials = dual_args(args...) + return Dual{T,typeof(value),length(partials)}(value, partials) end -@inline Dual{T}(value, partials::Tuple) where {T} = Dual{T}(value, Partials(partials)) -@inline Dual{T}(value, partials::Tuple{}) where {T} = Dual{T}(value, Partials{0,typeof(value)}(partials)) -@inline Dual{T}(value) where {T} = Dual{T}(value, ()) -@inline Dual{T}(x::Dual{T}) where {T} = Dual{T}(x, ()) -@inline Dual{T}(value, partial1, partials...) where {T} = Dual{T}(value, tuple(partial1, partials...)) -@inline Dual{T}(value::V, ::Chunk{N}, p::Val{i}) where {T,V,N,i} = Dual{T}(value, single_seed(Partials{N,V}, p)) -@inline Dual(args...) = Dual{Nothing}(args...) +@inline function Dual(args...) + value, partials = dual_args(args...) + return Dual{Tag{Nothing,typeof(value)}}(value, partials) +end # we define these special cases so that the "constructor <--> convert" pun holds for `Dual` @inline Dual{T,V,N}(x::Dual{T,V,N}) where {T,V,N} = x @@ -453,9 +467,13 @@ function Base.promote_rule(::Type{Dual{T1,V1,N1}}, end end -function Base.promote_rule(::Type{Dual{T,A,N}}, - ::Type{Dual{T,B,N}}) where {T,A,B,N} - return Dual{T,promote_type(A, B),N} +function Base.promote_rule(::Type{Dual{T,A,M}}, + ::Type{Dual{T,B,N}}) where {T,A,B,M,N} + if M === N + Dual{T,promote_type(A, B),N} + else + throw(ArgumentError(lazy"Cannot promote Duals with the same tag $T but $M and $N partials.")) + end end for R in (AbstractIrrational, Real, BigFloat, Bool) @@ -473,7 +491,6 @@ for R in (AbstractIrrational, Real, BigFloat, Bool) end @inline Base.convert(::Type{Dual{T,V,N}}, d::Dual{T}) where {T,V,N} = Dual{T}(V(value(T, d)), convert(Partials{N,V}, partials(T, d))) -@inline Base.convert(::Type{Dual{T,Dual{T,V,M},N}}, d::Dual{T,V,M}) where {T,V,N,M} = Dual{T}(d, Partials{N,Dual{T,V,M}}(zero_tuple(NTuple{N,Dual{T,V,M}}))) @inline Base.convert(::Type{Dual{T,V,N}}, x) where {T,V,N} = Dual{T}(V(x), zero(Partials{N,V})) @inline Base.convert(::Type{Dual{T,V,N}}, x::Number) where {T,V,N} = Dual{T}(V(x), zero(Partials{N,V})) Base.convert(::Type{D}, d::D) where {D<:Dual} = d diff --git a/test/AllocationsTest.jl b/test/AllocationsTest.jl index 94e7cddd..5c3a459f 100644 --- a/test/AllocationsTest.jl +++ b/test/AllocationsTest.jl @@ -5,7 +5,9 @@ using StaticArrays 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) +const D1_574 = ForwardDiff.Dual{ForwardDiff.Tag{Nothing,Float64},Float64,8} +const D2_574 = ForwardDiff.Dual{ForwardDiff.Tag{Nothing,D1_574},D1_574,4} +convert_test_574() = convert(ForwardDiff.Dual{ForwardDiff.Tag{Nothing,D2_574},D2_574,2}, 1.3) @testset "Test seed!/seed_zero_partials! allocations" begin x = rand(1000) diff --git a/test/ConfusionTest.jl b/test/ConfusionTest.jl index 64f6c5f4..9f503baf 100644 --- a/test/ConfusionTest.jl +++ b/test/ConfusionTest.jl @@ -4,6 +4,8 @@ using Test using ForwardDiff using LinearAlgebra +using StaticArrays +using DiffResults # Perturbation Confusion (Issue #83) # #------------------------------------# @@ -110,6 +112,49 @@ struct ATag end struct BTag end @test ForwardDiff.Dual{BTag}(ForwardDiff.Dual{ATag}(1.0, 2.0), ForwardDiff.Dual{ATag}(3.0, 4.0)) isa ForwardDiff.Dual{BTag} @test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag}(ForwardDiff.Dual{BTag}(1.0, 2.0), ForwardDiff.Dual{BTag}(3.0, 4.0)) +@test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag,Real,1}(ForwardDiff.Dual{BTag}(1.0, 2.0), ForwardDiff.Partials{1,Real}((3.0,))) +@test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag,Real,1}(1.0, ForwardDiff.Partials{1,Real}((ForwardDiff.Dual{BTag}(3.0, 4.0),))) + +# Tags of nested `Dual`s are unique +@test_throws ArgumentError("Cannot store a Dual with tag $ATag inside a Dual with the same tag.") ForwardDiff.Dual{ATag}(ForwardDiff.Dual{ATag}(1.0, 2.0), ForwardDiff.Dual{ATag}(3.0, 4.0)) +let d = ForwardDiff.Dual(ForwardDiff.Dual(1.0, 2.0), ForwardDiff.Dual(3.0, 4.0)) + T = ForwardDiff.Tag{Nothing,Float64} + S = ForwardDiff.Tag{Nothing,ForwardDiff.Dual{T,Float64,1}} + @test d isa ForwardDiff.Dual{S} + @test ForwardDiff.value(T, ForwardDiff.value(S, d)) == 1.0 + @test ForwardDiff.partials(T, ForwardDiff.value(S, d), 1) == 2.0 + @test ForwardDiff.value(T, ForwardDiff.partials(S, d, 1)) == 3.0 + @test ForwardDiff.partials(T, ForwardDiff.partials(S, d, 1), 1) == 4.0 +end +let f = x -> x[1]^2 * x[2], x = [3.0, 2.0] + @test ForwardDiff.hessian(f, x, ForwardDiff.HessianConfig(nothing, x)) == [4.0 6.0; 6.0 0.0] +end + +# Issue #845: all Hessian paths agree with the Jacobian of the gradient +strip_outer(x) = x +strip_outer(d::ForwardDiff.Dual{T}) where {T} = ForwardDiff.value(T, d) +f845a(z) = sum(abs2, z) + strip_outer(z[1]) * z[2] +f845b(z) = strip_outer(sum(abs2, z)) +for f in (f845a, f845b), x in ([1.0, 2.0, 3.0], SVector(1.0, 2.0, 3.0)) + H = ForwardDiff.jacobian(y -> ForwardDiff.gradient(f, y), x) + g = ForwardDiff.gradient(f, x) + @test ForwardDiff.hessian(f, x) == H + for c in 1:3 + @test ForwardDiff.hessian(f, x, ForwardDiff.HessianConfig(f, x, ForwardDiff.Chunk{c}())) == H + end + out = fill(NaN, 3, 3) + @test ForwardDiff.hessian!(out, f, x) === out + @test out == H + for result in (DiffResults.HessianResult(x), DiffResults.DiffResult(NaN, (fill(NaN, 3), fill(NaN, 3, 3)))) + r = ForwardDiff.hessian!(result, f, x) + if result isa DiffResults.MutableDiffResult + @test r === result + end + @test DiffResults.value(r) == f(x) + @test DiffResults.gradient(r) == g + @test DiffResults.hessian(r) == H + end +end end # module diff --git a/test/DualTest.jl b/test/DualTest.jl index 8ee4884f..752d7c40 100644 --- a/test/DualTest.jl +++ b/test/DualTest.jl @@ -30,7 +30,7 @@ dual_isapprox(a::Dual{T,T1,T2}, b::Dual{T3,T4,T5}) where {T,T1,T2,T3,T4,T5} = er ForwardDiff.:≺(::Type{TestTag}, ::Type{OuterTestTag}) = true ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false -@testset "Dual{Z,$V,$N} and Dual{Z,Dual{Z,$V,$M},$N}" for N in (0,3), M in (0,4), V in (Int, Float32) +@testset "Dual{Z,$V,$N} and Dual{Y,Dual{Z,$V,$M},$N}" for N in (0,3), M in (0,4), V in (Int, Float32) PARTIALS = Partials{N,V}(ntuple(n -> intrand(V), N)) PRIMAL = intrand(V) @@ -53,26 +53,27 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false M_PARTIALS = Partials{M,V}(ntuple(m -> intrand(V), M)) NESTED_PARTIALS = convert(Partials{N,Dual{TestTag,V,M}}, PARTIALS) - NESTED_FDNUM = Dual{TestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS) + NESTED_FDNUM = Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS) M_PARTIALS2 = Partials{M,V}(ntuple(m -> intrand(V), M)) NESTED_PARTIALS2 = convert(Partials{N,Dual{TestTag,V,M}}, PARTIALS2) - NESTED_FDNUM2 = Dual{TestTag}(Dual{TestTag}(PRIMAL2, M_PARTIALS2), NESTED_PARTIALS2) + NESTED_FDNUM2 = Dual{OuterTestTag}(Dual{TestTag}(PRIMAL2, M_PARTIALS2), NESTED_PARTIALS2) ################ # Constructors # ################ @test Dual{TestTag}(PRIMAL, PARTIALS...) === FDNUM - @test Dual(PRIMAL, PARTIALS...) === Dual{Nothing}(PRIMAL, PARTIALS...) - @test Dual(PRIMAL) === Dual{Nothing}(PRIMAL) + @test Dual(PRIMAL, PARTIALS...) === Dual{ForwardDiff.Tag{Nothing,V}}(PRIMAL, PARTIALS...) + @test Dual(PRIMAL) === Dual{ForwardDiff.Tag{Nothing,V}}(PRIMAL) + @test Dual(PRIMAL, convert(Partials{N,widen(V)}, PARTIALS)) isa Dual{ForwardDiff.Tag{Nothing,widen(V)},widen(V),N} @test_throws ArgumentError("The tag of a Dual must be a type, got 1.") Dual{1}(PRIMAL, PARTIALS) @test_throws ArgumentError("The tag of a Dual must be a type, got :tag.") Dual{:tag}(PRIMAL, PARTIALS) @test typeof(Dual{TestTag}(widen(V)(PRIMAL), PARTIALS)) === Dual{TestTag,widen(V),N} @test typeof(Dual{TestTag}(widen(V)(PRIMAL), PARTIALS.values)) === Dual{TestTag,widen(V),N} @test typeof(Dual{TestTag}(widen(V)(PRIMAL), PARTIALS...)) === Dual{TestTag,widen(V),N} - @test typeof(NESTED_FDNUM) == Dual{TestTag,Dual{TestTag,V,M},N} + @test typeof(NESTED_FDNUM) == Dual{OuterTestTag,Dual{TestTag,V,M},N} ############# # Accessors # @@ -80,18 +81,19 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test value(TestTag, PRIMAL) == PRIMAL @test value(TestTag, FDNUM) == PRIMAL - @test value(TestTag, NESTED_FDNUM) === Dual{TestTag}(PRIMAL, M_PARTIALS) + @test value(OuterTestTag, NESTED_FDNUM) === Dual{TestTag}(PRIMAL, M_PARTIALS) @test partials(TestTag, PRIMAL) == Partials{0,V}(tuple()) @test partials(TestTag, FDNUM) == PARTIALS - @test partials(TestTag, NESTED_FDNUM) === NESTED_PARTIALS + @test partials(OuterTestTag, NESTED_FDNUM) === NESTED_PARTIALS for i in 1:N @test partials(TestTag, FDNUM, i) == PARTIALS[i] end @test ForwardDiff.npartials(TestTag, typeof(FDNUM)) == N - @test ForwardDiff.npartials(TestTag, typeof(NESTED_FDNUM)) == N + @test ForwardDiff.npartials(TestTag, typeof(NESTED_FDNUM)) == M + @test ForwardDiff.npartials(OuterTestTag, typeof(NESTED_FDNUM)) == N @test (@test_deprecated r"`ForwardDiff.value` without a tag is deprecated" value(PRIMAL)) == PRIMAL @test (@test_deprecated r"`ForwardDiff.value` without a tag is deprecated" value(FDNUM)) == PRIMAL @@ -101,7 +103,7 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(FDNUM, i)) == PARTIALS[i] for j in 1:M @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(NESTED_FDNUM, i, j)) == partials(TestTag, NESTED_PARTIALS[i], j) - @test (@test_deprecated r"`ForwardDiff.partials\(T, x, i, j...\)` is deprecated" partials(TestTag, NESTED_FDNUM, i, j)) == partials(TestTag, NESTED_PARTIALS[i], j) + @test (@test_deprecated r"`ForwardDiff.partials\(T, x, i, j...\)` is deprecated" partials(OuterTestTag, NESTED_FDNUM, i, j)) == partials(TestTag, NESTED_PARTIALS[i], j) end end @test (@test_deprecated r"`ForwardDiff.npartials` without a tag is deprecated" ForwardDiff.npartials(FDNUM)) == N @@ -114,13 +116,13 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test ForwardDiff.valtype(TestTag, FDNUM) == V @test ForwardDiff.valtype(TestTag, typeof(FDNUM)) == V - @test ForwardDiff.valtype(TestTag, NESTED_FDNUM) == Dual{TestTag,V,M} - @test ForwardDiff.valtype(TestTag, typeof(NESTED_FDNUM)) == Dual{TestTag,V,M} + @test ForwardDiff.valtype(TestTag, NESTED_FDNUM) == Dual{OuterTestTag,V,N} + @test ForwardDiff.valtype(TestTag, typeof(NESTED_FDNUM)) == Dual{OuterTestTag,V,N} @test ForwardDiff.valtype(OuterTestTag, FDNUM) == Dual{TestTag,V,N} @test ForwardDiff.valtype(OuterTestTag, typeof(FDNUM)) == Dual{TestTag,V,N} - @test ForwardDiff.valtype(OuterTestTag, NESTED_FDNUM) == Dual{TestTag,Dual{TestTag,V,M},N} - @test ForwardDiff.valtype(OuterTestTag, typeof(NESTED_FDNUM)) == Dual{TestTag,Dual{TestTag,V,M},N} + @test ForwardDiff.valtype(OuterTestTag, NESTED_FDNUM) == Dual{TestTag,V,M} + @test ForwardDiff.valtype(OuterTestTag, typeof(NESTED_FDNUM)) == Dual{TestTag,V,M} OUTER_FDNUM = Dual{OuterTestTag}(PRIMAL, PARTIALS) @test ForwardDiff.valtype(TestTag, OUTER_FDNUM) == Dual{OuterTestTag,V,N} @@ -128,14 +130,11 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test value(TestTag, OUTER_FDNUM) === OUTER_FDNUM @test partials(TestTag, OUTER_FDNUM, 1) === zero(OUTER_FDNUM) @test partials(TestTag, OUTER_FDNUM) === Partials{0,typeof(OUTER_FDNUM)}(()) - INNER_FDNUM = Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS) - @test ForwardDiff.valtype(TestTag, INNER_FDNUM) == Dual{OuterTestTag,V,N} - @test ForwardDiff.valtype(TestTag, typeof(INNER_FDNUM)) == Dual{OuterTestTag,V,N} - @test value(TestTag, INNER_FDNUM) === Dual{OuterTestTag}(PRIMAL, PARTIALS) + @test value(TestTag, NESTED_FDNUM) === Dual{OuterTestTag}(PRIMAL, PARTIALS) for j in 1:M - @test partials(TestTag, INNER_FDNUM, j) === Dual{OuterTestTag}(M_PARTIALS[j], zero(PARTIALS)) + @test partials(TestTag, NESTED_FDNUM, j) === Dual{OuterTestTag}(M_PARTIALS[j], zero(PARTIALS)) end - @test partials(TestTag, INNER_FDNUM) === Partials{M,Dual{OuterTestTag,V,N}}(ntuple(j -> Dual{OuterTestTag}(M_PARTIALS[j], zero(PARTIALS)), M)) + @test partials(TestTag, NESTED_FDNUM) === Partials{M,Dual{OuterTestTag,V,N}}(ntuple(j -> Dual{OuterTestTag}(M_PARTIALS[j], zero(PARTIALS)), M)) ##################### # Generic Functions # @@ -255,25 +254,25 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test zero(FDNUM) === Dual{TestTag}(zero(PRIMAL), zero(PARTIALS)) @test zero(typeof(FDNUM)) === Dual{TestTag}(zero(V), zero(Partials{N,V})) - @test zero(NESTED_FDNUM) === Dual{TestTag}(Dual{TestTag}(zero(PRIMAL), zero(M_PARTIALS)), zero(NESTED_PARTIALS)) - @test zero(typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(zero(V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) + @test zero(NESTED_FDNUM) === Dual{OuterTestTag}(Dual{TestTag}(zero(PRIMAL), zero(M_PARTIALS)), zero(NESTED_PARTIALS)) + @test zero(typeof(NESTED_FDNUM)) === Dual{OuterTestTag}(Dual{TestTag}(zero(V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) @test one(FDNUM) === Dual{TestTag}(one(PRIMAL), zero(PARTIALS)) @test one(typeof(FDNUM)) === Dual{TestTag}(one(V), zero(Partials{N,V})) - @test one(NESTED_FDNUM) === Dual{TestTag}(Dual{TestTag}(one(PRIMAL), zero(M_PARTIALS)), zero(NESTED_PARTIALS)) - @test one(typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(one(V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) + @test one(NESTED_FDNUM) === Dual{OuterTestTag}(Dual{TestTag}(one(PRIMAL), zero(M_PARTIALS)), zero(NESTED_PARTIALS)) + @test one(typeof(NESTED_FDNUM)) === Dual{OuterTestTag}(Dual{TestTag}(one(V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) if V <: Integer @test rand(samerng(), FDNUM) == rand(samerng(), value(TestTag, FDNUM)) - @test rand(samerng(), NESTED_FDNUM) == rand(samerng(), value(TestTag, NESTED_FDNUM)) + @test rand(samerng(), NESTED_FDNUM) == rand(samerng(), value(OuterTestTag, NESTED_FDNUM)) elseif V <: AbstractFloat @test rand(samerng(), typeof(FDNUM)) === Dual{TestTag}(rand(samerng(), V), zero(Partials{N,V})) - @test rand(samerng(), typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(rand(samerng(), V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) + @test rand(samerng(), typeof(NESTED_FDNUM)) === Dual{OuterTestTag}(Dual{TestTag}(rand(samerng(), V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) @test randn(samerng(), typeof(FDNUM)) === Dual{TestTag}(randn(samerng(), V), zero(Partials{N,V})) - @test randn(samerng(), typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(randn(samerng(), V), zero(Partials{M,V})), + @test randn(samerng(), typeof(NESTED_FDNUM)) === Dual{OuterTestTag}(Dual{TestTag}(randn(samerng(), V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) @test randexp(samerng(), typeof(FDNUM)) === Dual{TestTag}(randexp(samerng(), V), zero(Partials{N,V})) - @test randexp(samerng(), typeof(NESTED_FDNUM)) === Dual{TestTag}(Dual{TestTag}(randexp(samerng(), V), zero(Partials{M,V})), + @test randexp(samerng(), typeof(NESTED_FDNUM)) === Dual{OuterTestTag}(Dual{TestTag}(randexp(samerng(), V), zero(Partials{M,V})), zero(Partials{N,Dual{TestTag,V,M}})) end @@ -291,13 +290,13 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false # Recall that FDNUM = Dual{TestTag}(PRIMAL, PARTIALS) has N partials, # and FDNUM2 has everything with a 2, and all random numbers nonzero. # M is the length of M_PARTIALS, which affects: - # NESTED_FDNUM = Dual{TestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS) + # NESTED_FDNUM = Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), NESTED_PARTIALS) NAN_PARTIALS = Partials{N,float(V)}(map(x -> oftype(float(x), NaN), PARTIALS.values)) @test (FDNUM == Dual{TestTag}(PRIMAL, PARTIALS2)) == (PARTIALS == PARTIALS2) @test isequal(FDNUM, Dual{TestTag}(PRIMAL, PARTIALS2)) == (PARTIALS == PARTIALS2) - @test isequal(NESTED_FDNUM, Dual{TestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS2), NESTED_PARTIALS2)) == ((M_PARTIALS == M_PARTIALS2) && (NESTED_PARTIALS == NESTED_PARTIALS2)) + @test isequal(NESTED_FDNUM, Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS2), NESTED_PARTIALS2)) == ((M_PARTIALS == M_PARTIALS2) && (NESTED_PARTIALS == NESTED_PARTIALS2)) if PRIMAL == PRIMAL2 @test isequal(FDNUM, Dual{TestTag}(PRIMAL, PARTIALS2)) == (PARTIALS == PARTIALS2) @@ -324,10 +323,10 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test !isless(Dual{TestTag}(1, NAN_PARTIALS), Dual{TestTag}(1, PARTIALS)) @test !(isless(Dual{TestTag}(2, PARTIALS), Dual{TestTag}(1, NAN_PARTIALS))) - @test isless(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS), Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2)) - @test isless(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS), Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === isless(NESTED_PARTIALS, NESTED_PARTIALS2) - @test !(isless(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS), Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS))) - @test !(isless(Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS), Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2))) + @test isless(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS), Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2)) + @test isless(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS), Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === isless(NESTED_PARTIALS, NESTED_PARTIALS2) + @test !(isless(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS), Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS))) + @test !(isless(Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS), Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2))) @test Dual{TestTag}(1, PARTIALS) < Dual{TestTag}(2, PARTIALS2) @test (Dual{TestTag}(1, PARTIALS) < Dual{TestTag}(1, PARTIALS2)) === (PARTIALS < PARTIALS2) @@ -338,10 +337,10 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test !(Dual{TestTag}(1, NAN_PARTIALS) < Dual{TestTag}(1, PARTIALS)) @test !(Dual{TestTag}(2, PARTIALS) < Dual{TestTag}(1, NAN_PARTIALS)) - @test Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) < Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2) - @test (Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) < Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS < NESTED_PARTIALS2) - @test !(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) < Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS)) - @test !(Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) < Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2)) + @test Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) < Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2) + @test (Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) < Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS < NESTED_PARTIALS2) + @test !(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) < Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS)) + @test !(Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) < Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2)) @test Dual{TestTag}(1, PARTIALS) <= Dual{TestTag}(2, PARTIALS2) @test (Dual{TestTag}(1, PARTIALS) <= Dual{TestTag}(1, PARTIALS2)) === (PARTIALS <= PARTIALS2) @@ -352,10 +351,10 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test (Dual{TestTag}(1, NAN_PARTIALS) <= Dual{TestTag}(1, PARTIALS)) === (N == 0) @test !(Dual{TestTag}(2, PARTIALS) <= Dual{TestTag}(1, NAN_PARTIALS)) - @test Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) <= Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2) - @test (Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) <= Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS <= NESTED_PARTIALS2) - @test Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) <= Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) - @test !(Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) <= Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2)) + @test Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) <= Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2) + @test (Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) <= Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS <= NESTED_PARTIALS2) + @test Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) <= Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) + @test !(Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) <= Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2)) @test Dual{TestTag}(2, PARTIALS) > Dual{TestTag}(1, PARTIALS2) @test (Dual{TestTag}(1, PARTIALS) > Dual{TestTag}(1, PARTIALS2)) === (PARTIALS > PARTIALS2) @@ -366,10 +365,10 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test !(Dual{TestTag}(1, NAN_PARTIALS) > Dual{TestTag}(1, PARTIALS)) @test Dual{TestTag}(2, PARTIALS) > Dual{TestTag}(1, NAN_PARTIALS) - @test Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) > Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2) - @test (Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) > Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS > NESTED_PARTIALS2) - @test !(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) > Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS)) - @test !(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) > Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2)) + @test Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) > Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2) + @test (Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) > Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS > NESTED_PARTIALS2) + @test !(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) > Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS)) + @test !(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) > Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2)) @test Dual{TestTag}(2, PARTIALS) >= Dual{TestTag}(1, PARTIALS2) @test (Dual{TestTag}(1, PARTIALS) >= Dual{TestTag}(1, PARTIALS2)) === (PARTIALS >= PARTIALS2) @@ -380,27 +379,27 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test ((Dual{TestTag}(1, NAN_PARTIALS) >= Dual{TestTag}(1, PARTIALS))) === (N == 0) @test Dual{TestTag}(2, PARTIALS) >= Dual{TestTag}(1, NAN_PARTIALS) - @test Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) >= Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2) - @test (Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) >= Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS >= NESTED_PARTIALS2) - @test Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) >= Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) - @test !(Dual{TestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) >= Dual{TestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2)) + @test Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS), NESTED_PARTIALS) >= Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS2), NESTED_PARTIALS2) + @test (Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) >= Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS2)) === (NESTED_PARTIALS >= NESTED_PARTIALS2) + @test Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) >= Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) + @test !(Dual{OuterTestTag}(Dual{TestTag}(1, M_PARTIALS), NESTED_PARTIALS) >= Dual{OuterTestTag}(Dual{TestTag}(2, M_PARTIALS2), NESTED_PARTIALS2)) @test isnan(Dual{TestTag}(NaN, PARTIALS)) @test !(isnan(FDNUM)) - @test isnan(Dual{TestTag}(Dual{TestTag}(NaN, M_PARTIALS), NESTED_PARTIALS)) + @test isnan(Dual{OuterTestTag}(Dual{TestTag}(NaN, M_PARTIALS), NESTED_PARTIALS)) @test !(isnan(NESTED_FDNUM)) @test isfinite(FDNUM) @test !(isfinite(Dual{TestTag}(Inf, PARTIALS))) @test isfinite(NESTED_FDNUM) - @test !(isfinite(Dual{TestTag}(Dual{TestTag}(NaN, M_PARTIALS), NESTED_PARTIALS))) + @test !(isfinite(Dual{OuterTestTag}(Dual{TestTag}(NaN, M_PARTIALS), NESTED_PARTIALS))) @test isinf(Dual{TestTag}(Inf, PARTIALS)) @test !(isinf(FDNUM)) - @test isinf(Dual{TestTag}(Dual{TestTag}(Inf, M_PARTIALS), NESTED_PARTIALS)) + @test isinf(Dual{OuterTestTag}(Dual{TestTag}(Inf, M_PARTIALS), NESTED_PARTIALS)) @test !(isinf(NESTED_FDNUM)) @test isreal(FDNUM) @@ -409,20 +408,20 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test isinteger(Dual{TestTag}(1.0, PARTIALS)) @test isinteger(FDNUM) == (V == Int) - @test isinteger(Dual{TestTag}(Dual{TestTag}(1.0, M_PARTIALS), NESTED_PARTIALS)) + @test isinteger(Dual{OuterTestTag}(Dual{TestTag}(1.0, M_PARTIALS), NESTED_PARTIALS)) @test isinteger(NESTED_FDNUM) == (V == Int) @test iseven(Dual{TestTag}(2)) @test !(iseven(Dual{TestTag}(1))) - @test iseven(Dual{TestTag}(Dual{TestTag}(2))) - @test !(iseven(Dual{TestTag}(Dual{TestTag}(1)))) + @test iseven(Dual{OuterTestTag}(Dual{TestTag}(2))) + @test !(iseven(Dual{OuterTestTag}(Dual{TestTag}(1)))) @test isodd(Dual{TestTag}(1)) @test !(isodd(Dual{TestTag}(2))) - @test isodd(Dual{TestTag}(Dual{TestTag}(1))) - @test !(isodd(Dual{TestTag}(Dual{TestTag}(2)))) + @test isodd(Dual{OuterTestTag}(Dual{TestTag}(1))) + @test !(isodd(Dual{OuterTestTag}(Dual{TestTag}(2)))) ######################## # Promotion/Conversion # @@ -435,29 +434,33 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test promote_type(Dual{TestTag,WIDE_T,N}, V) == Dual{TestTag,WIDE_T,N} @test promote_type(Dual{TestTag,V,N}, Dual{TestTag,V,N}) == Dual{TestTag,V,N} @test promote_type(Dual{TestTag,V,N}, Dual{TestTag,WIDE_T,N}) == Dual{TestTag,WIDE_T,N} - @test promote_type(Dual{TestTag,WIDE_T,N}, Dual{TestTag,Dual{TestTag,V,M},N}) == Dual{TestTag,Dual{TestTag,WIDE_T,M},N} + @test promote_type(Dual{OuterTestTag,WIDE_T,N}, Dual{OuterTestTag,Dual{TestTag,V,M},N}) == Dual{OuterTestTag,Dual{TestTag,WIDE_T,M},N} + @test promote_type(Dual{TestTag,V,M}, Dual{OuterTestTag,Dual{TestTag,V,M},N}) == Dual{OuterTestTag,Dual{TestTag,V,M},N} + if M != N + @test_throws ArgumentError("Cannot promote Duals with the same tag $TestTag but $M and $N partials.") promote_type(Dual{TestTag,V,M}, Dual{TestTag,V,N}) + end # issue #322 @test promote_type(Bool, Dual{TestTag,V,N}) == Dual{TestTag,promote_type(Bool, V),N} @test promote_type(BigFloat, Dual{TestTag,V,N}) == Dual{TestTag,promote_type(BigFloat, V),N} WIDE_FDNUM = convert(Dual{TestTag,WIDE_T,N}, FDNUM) - WIDE_NESTED_FDNUM = convert(Dual{TestTag,Dual{TestTag,WIDE_T,M},N}, NESTED_FDNUM) + WIDE_NESTED_FDNUM = convert(Dual{OuterTestTag,Dual{TestTag,WIDE_T,M},N}, NESTED_FDNUM) @test typeof(WIDE_FDNUM) === Dual{TestTag,WIDE_T,N} - @test typeof(WIDE_NESTED_FDNUM) === Dual{TestTag,Dual{TestTag,WIDE_T,M},N} + @test typeof(WIDE_NESTED_FDNUM) === Dual{OuterTestTag,Dual{TestTag,WIDE_T,M},N} @test value(TestTag, WIDE_FDNUM) == PRIMAL - @test (value(TestTag, WIDE_NESTED_FDNUM) == PRIMAL) == (M == 0) + @test (value(OuterTestTag, WIDE_NESTED_FDNUM) == PRIMAL) == (M == 0) @test convert(Dual, FDNUM) === FDNUM @test convert(Dual, NESTED_FDNUM) === NESTED_FDNUM @test convert(Dual{TestTag,V,N}, FDNUM) === FDNUM - @test convert(Dual{TestTag,Dual{TestTag,V,M},N}, NESTED_FDNUM) === NESTED_FDNUM + @test convert(Dual{OuterTestTag,Dual{TestTag,V,M},N}, NESTED_FDNUM) === NESTED_FDNUM @test convert(Dual{TestTag,WIDE_T,N}, PRIMAL) === Dual{TestTag}(WIDE_T(PRIMAL), zero(Partials{N,WIDE_T})) - @test convert(Dual{TestTag,Dual{TestTag,WIDE_T,M},N}, PRIMAL) === Dual{TestTag}(Dual{TestTag}(WIDE_T(PRIMAL), zero(Partials{M,WIDE_T})), zero(Partials{N,Dual{TestTag,V,M}})) - @test convert(Dual{TestTag,Dual{TestTag,V,M},N}, FDNUM) === Dual{TestTag}(convert(Dual{TestTag,V,M}, PRIMAL), convert(Partials{N,Dual{TestTag,V,M}}, PARTIALS)) - @test convert(Dual{TestTag,Dual{TestTag,WIDE_T,M},N}, FDNUM) === Dual{TestTag}(convert(Dual{TestTag,WIDE_T,M}, PRIMAL), convert(Partials{N,Dual{TestTag,WIDE_T,M}}, PARTIALS)) + @test convert(Dual{OuterTestTag,Dual{TestTag,WIDE_T,M},N}, PRIMAL) === Dual{OuterTestTag}(Dual{TestTag}(WIDE_T(PRIMAL), zero(Partials{M,WIDE_T})), zero(Partials{N,Dual{TestTag,V,M}})) + @test convert(Dual{OuterTestTag,Dual{TestTag,V,M},N}, Dual{TestTag}(PRIMAL, M_PARTIALS)) === Dual{OuterTestTag}(Dual{TestTag}(PRIMAL, M_PARTIALS), zero(Partials{N,Dual{TestTag,V,M}})) + @test convert(Dual{OuterTestTag,Dual{TestTag,WIDE_T,M},N}, Dual{TestTag}(PRIMAL, M_PARTIALS)) === Dual{OuterTestTag}(Dual{TestTag}(WIDE_T(PRIMAL), convert(Partials{M,WIDE_T}, M_PARTIALS)), zero(Partials{N,Dual{TestTag,WIDE_T,M}})) ############## # Arithmetic # @@ -470,19 +473,19 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test FDNUM + PRIMAL === Dual{TestTag}(value(TestTag, FDNUM) + PRIMAL, partials(TestTag, FDNUM)) @test PRIMAL + FDNUM === Dual{TestTag}(value(TestTag, FDNUM) + PRIMAL, partials(TestTag, FDNUM)) - @test NESTED_FDNUM + NESTED_FDNUM2 === Dual{TestTag}(value(TestTag, NESTED_FDNUM) + value(TestTag, NESTED_FDNUM2), partials(TestTag, NESTED_FDNUM) + partials(TestTag, NESTED_FDNUM2)) - @test NESTED_FDNUM + PRIMAL === Dual{TestTag}(value(TestTag, NESTED_FDNUM) + PRIMAL, partials(TestTag, NESTED_FDNUM)) - @test PRIMAL + NESTED_FDNUM === Dual{TestTag}(value(TestTag, NESTED_FDNUM) + PRIMAL, partials(TestTag, NESTED_FDNUM)) + @test NESTED_FDNUM + NESTED_FDNUM2 === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) + value(OuterTestTag, NESTED_FDNUM2), partials(OuterTestTag, NESTED_FDNUM) + partials(OuterTestTag, NESTED_FDNUM2)) + @test NESTED_FDNUM + PRIMAL === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) + PRIMAL, partials(OuterTestTag, NESTED_FDNUM)) + @test PRIMAL + NESTED_FDNUM === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) + PRIMAL, partials(OuterTestTag, NESTED_FDNUM)) @test FDNUM - FDNUM2 === Dual{TestTag}(value(TestTag, FDNUM) - value(TestTag, FDNUM2), partials(TestTag, FDNUM) - partials(TestTag, FDNUM2)) @test FDNUM - PRIMAL === Dual{TestTag}(value(TestTag, FDNUM) - PRIMAL, partials(TestTag, FDNUM)) @test PRIMAL - FDNUM === Dual{TestTag}(PRIMAL - value(TestTag, FDNUM), -(partials(TestTag, FDNUM))) @test -(FDNUM) === Dual{TestTag}(-(value(TestTag, FDNUM)), -(partials(TestTag, FDNUM))) - @test NESTED_FDNUM - NESTED_FDNUM2 === Dual{TestTag}(value(TestTag, NESTED_FDNUM) - value(TestTag, NESTED_FDNUM2), partials(TestTag, NESTED_FDNUM) - partials(TestTag, NESTED_FDNUM2)) - @test NESTED_FDNUM - PRIMAL === Dual{TestTag}(value(TestTag, NESTED_FDNUM) - PRIMAL, partials(TestTag, NESTED_FDNUM)) - @test PRIMAL - NESTED_FDNUM === Dual{TestTag}(PRIMAL - value(TestTag, NESTED_FDNUM), -(partials(TestTag, NESTED_FDNUM))) - @test -(NESTED_FDNUM) === Dual{TestTag}(-(value(TestTag, NESTED_FDNUM)), -(partials(TestTag, NESTED_FDNUM))) + @test NESTED_FDNUM - NESTED_FDNUM2 === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) - value(OuterTestTag, NESTED_FDNUM2), partials(OuterTestTag, NESTED_FDNUM) - partials(OuterTestTag, NESTED_FDNUM2)) + @test NESTED_FDNUM - PRIMAL === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) - PRIMAL, partials(OuterTestTag, NESTED_FDNUM)) + @test PRIMAL - NESTED_FDNUM === Dual{OuterTestTag}(PRIMAL - value(OuterTestTag, NESTED_FDNUM), -(partials(OuterTestTag, NESTED_FDNUM))) + @test -(NESTED_FDNUM) === Dual{OuterTestTag}(-(value(OuterTestTag, NESTED_FDNUM)), -(partials(OuterTestTag, NESTED_FDNUM))) # Multiplication # #----------------# @@ -491,9 +494,9 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test FDNUM * PRIMAL === Dual{TestTag}(value(TestTag, FDNUM) * PRIMAL, partials(TestTag, FDNUM) * PRIMAL) @test PRIMAL * FDNUM === Dual{TestTag}(value(TestTag, FDNUM) * PRIMAL, partials(TestTag, FDNUM) * PRIMAL) - @test NESTED_FDNUM * NESTED_FDNUM2 === Dual{TestTag}(value(TestTag, NESTED_FDNUM) * value(TestTag, NESTED_FDNUM2), ForwardDiff._mul_partials(partials(TestTag, NESTED_FDNUM), partials(TestTag, NESTED_FDNUM2), value(TestTag, NESTED_FDNUM2), value(TestTag, NESTED_FDNUM))) - @test NESTED_FDNUM * PRIMAL === Dual{TestTag}(value(TestTag, NESTED_FDNUM) * PRIMAL, partials(TestTag, NESTED_FDNUM) * PRIMAL) - @test PRIMAL * NESTED_FDNUM === Dual{TestTag}(value(TestTag, NESTED_FDNUM) * PRIMAL, partials(TestTag, NESTED_FDNUM) * PRIMAL) + @test NESTED_FDNUM * NESTED_FDNUM2 === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) * value(OuterTestTag, NESTED_FDNUM2), ForwardDiff._mul_partials(partials(OuterTestTag, NESTED_FDNUM), partials(OuterTestTag, NESTED_FDNUM2), value(OuterTestTag, NESTED_FDNUM2), value(OuterTestTag, NESTED_FDNUM))) + @test NESTED_FDNUM * PRIMAL === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) * PRIMAL, partials(OuterTestTag, NESTED_FDNUM) * PRIMAL) + @test PRIMAL * NESTED_FDNUM === Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) * PRIMAL, partials(OuterTestTag, NESTED_FDNUM) * PRIMAL) # Division # #----------# @@ -513,9 +516,9 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test dual_isapprox(FDNUM / PRIMAL, Dual{TestTag}(value(TestTag, FDNUM) / PRIMAL, partials(TestTag, FDNUM) / PRIMAL)) @test dual_isapprox(PRIMAL / FDNUM, Dual{TestTag}(PRIMAL / value(TestTag, FDNUM), (-(PRIMAL) / value(TestTag, FDNUM)^2) * partials(TestTag, FDNUM))) - @test dual_isapprox(NESTED_FDNUM / NESTED_FDNUM2, Dual{TestTag}(value(TestTag, NESTED_FDNUM) / value(TestTag, NESTED_FDNUM2), ForwardDiff._div_partials(partials(TestTag, NESTED_FDNUM), partials(TestTag, NESTED_FDNUM2), value(TestTag, NESTED_FDNUM), value(TestTag, NESTED_FDNUM2)))) - @test dual_isapprox(NESTED_FDNUM / PRIMAL, Dual{TestTag}(value(TestTag, NESTED_FDNUM) / PRIMAL, partials(TestTag, NESTED_FDNUM) / PRIMAL)) - @test dual_isapprox(PRIMAL / NESTED_FDNUM, Dual{TestTag}(PRIMAL / value(TestTag, NESTED_FDNUM), (-(PRIMAL) / value(TestTag, NESTED_FDNUM)^2) * partials(TestTag, NESTED_FDNUM))) + @test dual_isapprox(NESTED_FDNUM / NESTED_FDNUM2, Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) / value(OuterTestTag, NESTED_FDNUM2), ForwardDiff._div_partials(partials(OuterTestTag, NESTED_FDNUM), partials(OuterTestTag, NESTED_FDNUM2), value(OuterTestTag, NESTED_FDNUM), value(OuterTestTag, NESTED_FDNUM2)))) + @test dual_isapprox(NESTED_FDNUM / PRIMAL, Dual{OuterTestTag}(value(OuterTestTag, NESTED_FDNUM) / PRIMAL, partials(OuterTestTag, NESTED_FDNUM) / PRIMAL)) + @test dual_isapprox(PRIMAL / NESTED_FDNUM, Dual{OuterTestTag}(PRIMAL / value(OuterTestTag, NESTED_FDNUM), (-(PRIMAL) / value(OuterTestTag, NESTED_FDNUM)^2) * partials(OuterTestTag, NESTED_FDNUM))) # Exponentiation # #----------------# @@ -719,7 +722,7 @@ end @test isfinite(dfmin) @test isfinite(dfmax) - @test floatmin(Dual{Nothing, ForwardDiff.Dual{Nothing, Float64, 2}, 1}) === Dual{Nothing}(Dual{Nothing}(floatmin(Float64),0.0,0.0),Dual{Nothing}(0.0,0.0,0.0)) + @test floatmin(Dual{OuterTestTag, Dual{TestTag, Float64, 2}, 1}) === Dual{OuterTestTag}(Dual{TestTag}(floatmin(Float64),0.0,0.0),Dual{TestTag}(0.0,0.0,0.0)) end @testset "Integer" begin @@ -745,8 +748,8 @@ end @testset "float" begin # issue #492 @test float(Dual{Nothing, Int, 2}) === Dual{Nothing, Float64, 2} - @test float(Dual(1)) isa Dual{Nothing, Float64, 0} - @test value.(Nothing, float.(Dual.(1:4, 2:5, 3:6))) isa Vector{Float64} + @test float(Dual(1)) isa Dual{ForwardDiff.Tag{Nothing,Int}, Float64, 0} + @test value.(ForwardDiff.Tag{Nothing,Int}, float.(Dual.(1:4, 2:5, 3:6))) isa Vector{Float64} @test ForwardDiff.derivative(float, 1)::Float64 === 1.0 end diff --git a/test/MiscTest.jl b/test/MiscTest.jl index 8328461a..c3ccfa85 100644 --- a/test/MiscTest.jl +++ b/test/MiscTest.jl @@ -128,7 +128,7 @@ end # NaNs # #------# -@test ForwardDiff.partials(Nothing, NaNMath.pow(ForwardDiff.Dual(-2.0,1.0),ForwardDiff.Dual(2.0,0.0)),1) == -4.0 +@test ForwardDiff.partials(ForwardDiff.Tag{Nothing,Float64}, NaNMath.pow(ForwardDiff.Dual(-2.0,1.0),ForwardDiff.Dual(2.0,0.0)),1) == -4.0 # Partials{0} # #-------------# diff --git a/test/SeedTest.jl b/test/SeedTest.jl index 06c54d9b..ace810ae 100644 --- a/test/SeedTest.jl +++ b/test/SeedTest.jl @@ -27,12 +27,12 @@ const SEED_CASES = ( # Positions within `sidx` whose partials are zero. zeroed_positions(duals, sidx) = - [i for (i, idx) in enumerate(sidx) if iszero(ForwardDiff.partials(Nothing, duals[idx]))] + [i for (i, idx) in enumerate(sidx) if iszero(ForwardDiff.partials(ForwardDiff.Tag{Nothing,Float64}, duals[idx]))] # Compares over *every* index of `x`, not just the structural ones, so a bug misplacing values # outside the structural set is visible. Off-structure reads are safe: the wrapper types return # `zero(Dual)` without touching the (uninitialized) parent storage. -values_match(duals, x) = all(idx -> ForwardDiff.value(Nothing, duals[idx]) == x[idx], eachindex(x)) +values_match(duals, x) = all(idx -> ForwardDiff.value(ForwardDiff.Tag{Nothing,Float64}, duals[idx]) == x[idx], eachindex(x)) function fill_marker!(duals, x, sidx, marker) D = eltype(duals) @@ -45,7 +45,7 @@ end @testset "seed_zero_partials!: $(nameof(typeof(x)))" for (x, sidx) in SEED_CASES cfg = ForwardDiff.GradientConfig(nothing, x, ForwardDiff.Chunk{3}()) duals, seeds = cfg.duals, cfg.seeds - N = ForwardDiff.npartials(Nothing, eltype(duals)) + N = ForwardDiff.npartials(ForwardDiff.Tag{Nothing,Float64}, eltype(duals)) marker = Partials(ntuple(i -> Float64(i), N)) nstruct = length(sidx) From 621111489bba2ae5a6007906b1fd951d35234b21 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Mon, 28 Sep 2026 20:01:08 +0200 Subject: [PATCH 5/8] Test precompilation and deserialization of tags (#714, #801, #320) Co-Authored-By: Claude Opus 5.5 (1M context) --- Project.toml | 3 ++- test/ConfusionTest.jl | 24 ++++++++++++++++++++++++ test/packages/P714/src/P714.jl | 16 ++++++++++++++++ test/packages/PkgA/src/PkgA.jl | 12 ++++++++++++ test/packages/PkgB/src/PkgB.jl | 12 ++++++++++++ test/packages/TagDefs/src/TagDefs.jl | 10 ++++++++++ 6 files changed, 76 insertions(+), 1 deletion(-) create mode 100644 test/packages/P714/src/P714.jl create mode 100644 test/packages/PkgA/src/PkgA.jl create mode 100644 test/packages/PkgB/src/PkgB.jl create mode 100644 test/packages/TagDefs/src/TagDefs.jl diff --git a/Project.toml b/Project.toml index 801792c6..ca5d0736 100644 --- a/Project.toml +++ b/Project.toml @@ -41,9 +41,10 @@ DiffTests = "de460e47-3fe3-5279-bb4a-814414816d5d" InteractiveUtils = "b77e0a4c-d291-57a0-90e8-8db25a27a240" IrrationalConstants = "92d709cd-6900-40b7-9082-c6be49f344b6" JET = "c3a54625-cd67-489e-a8e7-0a5a0ff4e31b" +Serialization = "9e88b42a-f829-5b0c-bbe9-9e923198166b" SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf" StaticArrays = "90137ffa-7385-5640-81b9-e52037218182" Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40" [targets] -test = ["Calculus", "DiffTests", "IrrationalConstants", "JET", "SparseArrays", "StaticArrays", "Test", "InteractiveUtils"] +test = ["Calculus", "DiffTests", "IrrationalConstants", "JET", "Serialization", "SparseArrays", "StaticArrays", "Test", "InteractiveUtils"] diff --git a/test/ConfusionTest.jl b/test/ConfusionTest.jl index 9f503baf..44b31297 100644 --- a/test/ConfusionTest.jl +++ b/test/ConfusionTest.jl @@ -6,6 +6,7 @@ using ForwardDiff using LinearAlgebra using StaticArrays using DiffResults +using Serialization # Perturbation Confusion (Issue #83) # #------------------------------------# @@ -156,5 +157,28 @@ for f in (f845a, f845b), x in ([1.0, 2.0, 3.0], SVector(1.0, 2.0, 3.0)) end end +# Output of `code` in a new process, which can load the packages in `packages/` +julia_output(code) = readchomp(`$(Base.julia_cmd()) --project=$(Base.active_project()) -e "push!(LOAD_PATH, $(repr(joinpath(@__DIR__, "packages")))); $code"`) + +# Issue #714: nested derivatives in precompiled code +@test julia_output("using P714; print(P714.compute_derivative(1, 0))") == "2" + +# Issue #801: packages precompiling the same tags in different orders +@test julia_output("using PkgA, PkgB; print(PkgA.mix() === PkgB.mix())") == "true" + +# Issue #320: a function and its config deserialized separately have different tags +dir = mktempdir() +julia_output(""" + using ForwardDiff, Serialization + f = let c = 2.0 + x -> c * sum(abs2, x) + end + serialize($(repr(joinpath(dir, "f"))), f) + serialize($(repr(joinpath(dir, "cfg"))), ForwardDiff.GradientConfig(f, [1.0, 2.0])) + """) +f320 = deserialize(joinpath(dir, "f")) +cfg320 = deserialize(joinpath(dir, "cfg")) +@test_throws "Invalid Tag object" ForwardDiff.gradient(f320, [1.0, 2.0], cfg320) +@test ForwardDiff.gradient(f320, [1.0, 2.0], cfg320, Val(false)) == [4.0, 8.0] end # module diff --git a/test/packages/P714/src/P714.jl b/test/packages/P714/src/P714.jl new file mode 100644 index 00000000..890724cf --- /dev/null +++ b/test/packages/P714/src/P714.jl @@ -0,0 +1,16 @@ +module P714 + +using ForwardDiff + +dispatch(::Val{0}, x) = x +dispatch(::Val{1}, x) = ForwardDiff.derivative(z -> x + z^2, x) + +# prevents precompilation of `dispatch` +indirection = dispatch + +compute_derivative(α::Int, y) = ForwardDiff.derivative(x -> indirection(Val{α}(), x), y) + +ForwardDiff.derivative(x -> x^2, 1) +compute_derivative(0, 0) + +end diff --git a/test/packages/PkgA/src/PkgA.jl b/test/packages/PkgA/src/PkgA.jl new file mode 100644 index 00000000..0af76f13 --- /dev/null +++ b/test/packages/PkgA/src/PkgA.jl @@ -0,0 +1,12 @@ +module PkgA + +using ForwardDiff: Dual, Tag +using TagDefs: FA, FB, TA, TB + +mix() = Dual{TA}(1.0, 1.0) * Dual{TB}(2.0, 1.0) + +Tag(FA(), Float64) +Tag(FB(), Float64) +mix() + +end diff --git a/test/packages/PkgB/src/PkgB.jl b/test/packages/PkgB/src/PkgB.jl new file mode 100644 index 00000000..89cd6a26 --- /dev/null +++ b/test/packages/PkgB/src/PkgB.jl @@ -0,0 +1,12 @@ +module PkgB + +using ForwardDiff: Dual, Tag +using TagDefs: FA, FB, TA, TB + +mix() = Dual{TA}(1.0, 1.0) * Dual{TB}(2.0, 1.0) + +Tag(FB(), Float64) +Tag(FA(), Float64) +mix() + +end diff --git a/test/packages/TagDefs/src/TagDefs.jl b/test/packages/TagDefs/src/TagDefs.jl new file mode 100644 index 00000000..c2b434f2 --- /dev/null +++ b/test/packages/TagDefs/src/TagDefs.jl @@ -0,0 +1,10 @@ +module TagDefs + +using ForwardDiff: Tag + +struct FA end +struct FB end +const TA = Tag{FA,Float64} +const TB = Tag{FB,Float64} + +end From b3f4379063f19dff333342f99cd7145dfc8fa327 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Tue, 29 Sep 2026 14:57:38 +0200 Subject: [PATCH 6/8] Fix constructors for abstract value types, promote without throwing, and cover new lines MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - `Dual{T}(args...)` and `Dual(args...)` take the value type from the partials, so e.g. `zero(Dual{T,Real,N})` and `convert(Dual{T,Real,N}, x)` work again - `promote_rule` for the same tag with different numbers of partials returns `Union{}` instead of throwing, so `promote_type` falls back to `typejoin` - Include the package UUID in the tag key, so types of packages with the same name are distinguished - Remove the redundant `Tag(::Nothing, V)` method - Test `≺` with symbols, values and `Vararg`s as parameters, `npartials` without the tag, the deprecated `partials(x, i)` and `show` Co-Authored-By: Claude Opus 5.5 --- src/config.jl | 9 ++++----- src/dual.jl | 17 ++++++++--------- test/ConfusionTest.jl | 6 ++++++ test/DualTest.jl | 11 ++++++++++- 4 files changed, 28 insertions(+), 15 deletions(-) diff --git a/src/config.jl b/src/config.jl index f50762e4..1e0111c9 100644 --- a/src/config.jl +++ b/src/config.jl @@ -7,13 +7,12 @@ end Tag(f::F, ::Type{V}) where {F,V} = Tag{F,V}() -Tag(::Nothing, ::Type{V}) where {V} = Tag{Nothing,V}() - # Encodes a type (or type parameter) as a sequence of strings that depends only on its -# structure. Distinct objects have distinct keys, and the key of a parameter is a strict -# subsequence of the key of the type. +# structure. Distinct objects have distinct keys, except a type and its redefinition in the +# same session. The key of a parameter is a strict subsequence of the key of the type. function typekey!(key::Vector{String}, x::DataType) - push!(key, "T", string(fullname(parentmodule(x))), String(nameof(x)), string(length(x.parameters))) + mod = parentmodule(x) + push!(key, "T", string(Base.PkgId(Base.moduleroot(mod)).uuid), string(fullname(mod)), String(nameof(x)), string(length(x.parameters))) foreach(p -> typekey!(key, p), x.parameters) return key end diff --git a/src/dual.jl b/src/dual.jl index 019b03a3..63991879 100644 --- a/src/dual.jl +++ b/src/dual.jl @@ -87,12 +87,12 @@ end @inline function Dual{T}(args...) where {T} value, partials = dual_args(args...) - return Dual{T,typeof(value),length(partials)}(value, partials) + return Dual{T,eltype(partials),length(partials)}(value, partials) end @inline function Dual(args...) value, partials = dual_args(args...) - return Dual{Tag{Nothing,typeof(value)}}(value, partials) + return Dual{Tag{Nothing,eltype(partials)}}(value, partials) end # we define these special cases so that the "constructor <--> convert" pun holds for `Dual` @@ -467,15 +467,14 @@ function Base.promote_rule(::Type{Dual{T1,V1,N1}}, end end -function Base.promote_rule(::Type{Dual{T,A,M}}, - ::Type{Dual{T,B,N}}) where {T,A,B,M,N} - if M === N - Dual{T,promote_type(A, B),N} - else - throw(ArgumentError(lazy"Cannot promote Duals with the same tag $T but $M and $N partials.")) - end +function Base.promote_rule(::Type{Dual{T,A,N}}, + ::Type{Dual{T,B,N}}) where {T,A,B,N} + return Dual{T,promote_type(A, B),N} end +# no common type for different numbers of partials, `promote_type` falls back to `typejoin` +Base.promote_rule(::Type{Dual{T,A,M}}, ::Type{Dual{T,B,N}}) where {T,A,B,M,N} = Union{} + for R in (AbstractIrrational, Real, BigFloat, Bool) if isconcretetype(R) # issue #322 @eval begin diff --git a/test/ConfusionTest.jl b/test/ConfusionTest.jl index 44b31297..c5b66c90 100644 --- a/test/ConfusionTest.jl +++ b/test/ConfusionTest.jl @@ -108,6 +108,12 @@ let T = ForwardDiff.Tag{AFunction,Float64}, d = ForwardDiff.Dual{T}(1.0, 1.0), f end end +# Distinct tags are ordered, also with symbols, values or `Vararg`s as parameters +for (A, B) in ((Val{:a}, Val{:b}), (Val{(1, 2)}, Val{(2, 1)}), (Tuple{Vararg{Int}}, Tuple{Vararg{Float64}})) + T, S = ForwardDiff.Tag{A,Float64}, ForwardDiff.Tag{B,Float64} + @test ForwardDiff.:≺(T, S) != ForwardDiff.:≺(S, T) +end + # Nested `Dual`s store the greatest tag outermost struct ATag end struct BTag end diff --git a/test/DualTest.jl b/test/DualTest.jl index 752d7c40..87203db4 100644 --- a/test/DualTest.jl +++ b/test/DualTest.jl @@ -94,11 +94,14 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test ForwardDiff.npartials(TestTag, typeof(FDNUM)) == N @test ForwardDiff.npartials(TestTag, typeof(NESTED_FDNUM)) == M @test ForwardDiff.npartials(OuterTestTag, typeof(NESTED_FDNUM)) == N + @test ForwardDiff.npartials(TestTag, V) == 0 + @test ForwardDiff.npartials(OutermostTestTag, typeof(NESTED_FDNUM)) == 0 @test (@test_deprecated r"`ForwardDiff.value` without a tag is deprecated" value(PRIMAL)) == PRIMAL @test (@test_deprecated r"`ForwardDiff.value` without a tag is deprecated" value(FDNUM)) == PRIMAL @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(PRIMAL)) == Partials{0,V}(tuple()) @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(FDNUM)) == PARTIALS + @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(PRIMAL, 1)) == zero(PRIMAL) for i in 1:N @test (@test_deprecated r"`ForwardDiff.partials` without a tag is deprecated" partials(FDNUM, i)) == PARTIALS[i] for j in 1:M @@ -437,8 +440,10 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false @test promote_type(Dual{OuterTestTag,WIDE_T,N}, Dual{OuterTestTag,Dual{TestTag,V,M},N}) == Dual{OuterTestTag,Dual{TestTag,WIDE_T,M},N} @test promote_type(Dual{TestTag,V,M}, Dual{OuterTestTag,Dual{TestTag,V,M},N}) == Dual{OuterTestTag,Dual{TestTag,V,M},N} if M != N - @test_throws ArgumentError("Cannot promote Duals with the same tag $TestTag but $M and $N partials.") promote_type(Dual{TestTag,V,M}, Dual{TestTag,V,N}) + @test promote_type(Dual{TestTag,V,M}, Dual{TestTag,V,N}) == Dual{TestTag,V} end + @test zero(Dual{TestTag,Real,N}) isa Dual{TestTag,Real,N} + @test value(TestTag, convert(Dual{TestTag,Real,N}, PRIMAL)) == PRIMAL # issue #322 @test promote_type(Bool, Dual{TestTag,V,N}) == Dual{TestTag,promote_type(Bool, V),N} @@ -753,6 +758,10 @@ end @test ForwardDiff.derivative(float, 1)::Float64 === 1.0 end +@testset "show" begin + @test repr(Dual(1.0, 2.0)) == "Dual{ForwardDiff.Tag{Nothing, Float64}}(1.0,2.0)" +end + @testset "TwicePrecision" begin @test ForwardDiff.derivative(x -> sum(1 .+ x .* (0:0.1:1)), 1) == 5.5 end From c0bed6e3045594dd530b2d257ded054c80ef8650 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Tue, 29 Sep 2026 21:13:23 +0200 Subject: [PATCH 7/8] Keep greater tags outside when seeding With an abstract input type, the tag of an inner call can be smaller than tags of the input elements. Seeding now keeps layers with greater tags outside, and the work buffers use a bound on the seeded type (`Real` if the value type is abstract and can contain `Dual`s). Co-Authored-By: Claude Opus 5.5 --- src/apiutils.jl | 55 +++++++++++++++++++++++++++-------------- src/config.jl | 17 ++++++------- src/derivative.jl | 4 +-- src/gradient.jl | 12 ++++----- src/jacobian.jl | 16 ++++++------ test/AllocationsTest.jl | 29 +++++++++++----------- test/ConfusionTest.jl | 13 ++++++++++ test/SeedTest.jl | 14 +++++------ 8 files changed, 96 insertions(+), 64 deletions(-) diff --git a/src/apiutils.jl b/src/apiutils.jl index 0615fdb3..0d5401ec 100644 --- a/src/apiutils.jl +++ b/src/apiutils.jl @@ -20,14 +20,14 @@ end function vector_mode_dual_eval!(f::F, cfg::Union{JacobianConfig,GradientConfig}, x) where {F} xdual = cfg.duals - seed!(xdual, x, cfg.seeds) + seed!(eltype(cfg), xdual, x, cfg.seeds) return f(xdual) end -function vector_mode_dual_eval!(f!::F, cfg::JacobianConfig, y, x) where {F} +function vector_mode_dual_eval!(f!::F, cfg::JacobianConfig{T,V,N}, y, x) where {F,T,V,N} ydual, xdual = cfg.duals - seed!(xdual, x, cfg.seeds) - seed_zero_partials!(ydual, y) + seed!(eltype(cfg), xdual, x, cfg.seeds) + seed_zero_partials!(Dual{T,eltype(y),N}, ydual, y) f!(ydual, xdual) return ydual end @@ -40,6 +40,25 @@ end return Expr(:tuple, [:(single_seed(Partials{N,V}, Val{$i}())) for i in 1:N]...) end +# Seeds `x` with tag `T` and partials `p`. Layers of `x` with greater tags are kept outside, so +# nested `Dual`s stay sorted even if `T` is not greater than all tags in `x`. +@inline seed_dual(::Type{T}, x, p::Partials) where {T} = Dual{T}(x, p) +@inline function seed_dual(::Type{T}, x::Dual{S}, p::Partials) where {T,S} + T ≺ S || return Dual{T}(x, p) + # the seeds are constants, so their partials w.r.t. `S` are zero + q = map_partials(y -> value(S, y), valtype(S, eltype(p)), p) + return Dual{S}(seed_dual(T, value(S, x), q), map(y -> seed_dual(T, y, zero(q)), partials(S, x).values)) +end + +# Type of `seed_dual(T, x, p)` for `x::V` and `p::Partials{N,V}`. If `V` is abstract and can contain +# `Dual`s, their tags may be greater than `T`, so only `Real` is a bound. +function seed_type(::Type{Dual{T,V,N}}) where {T,V,N} + return isconcretetype(V) || typeintersect(V, Dual) === Union{} ? Dual{T,V,N} : Real +end +function seed_type(::Type{Dual{T,Dual{S,W,M},N}}) where {T,S,W,M,N} + return T ≺ S ? Dual{S,seed_type(Dual{T,W,N}),M} : Dual{T,Dual{S,W,M},N} +end + # Only seed indices that are structurally non-zero structural_eachindex(x::AbstractArray) = structural_eachindex(x, x) function structural_eachindex(x::AbstractArray, y::AbstractArray) @@ -73,29 +92,29 @@ end # Copies the values of `x` into `duals` with zero partials. Used both to remove seeds `duals` is # currently carrying and to initialize a freshly allocated work buffer, whose elements must all be # written before the target function reads them. -seed_zero_partials!(duals::AbstractArray{Dual{T,V,N}}, x) where {T,V,N} = - _seed_zero_partials!(duals, x, structural_eachindex(duals, x)) +seed_zero_partials!(::Type{D}, duals::AbstractArray, x) where {D<:Dual} = + _seed_zero_partials!(D, duals, x, structural_eachindex(duals, x)) # Zeroes the partials of `count` elements starting at structural position `index`. Chunk mode only # needs to clear the chunk it just seeded, so writing through to the end of the array would be O(n) # redundant work per chunk, i.e. O(n^2/N) per sweep. `count` mirrors the `chunksize` argument of -# `seed!(duals, x, index, seeds, chunksize)`. -function seed_zero_partials!(duals::AbstractArray{Dual{T,V,N}}, x, index, +# `seed!(D, duals, x, index, seeds, chunksize)`. +function seed_zero_partials!(::Type{Dual{T,V,N}}, duals::AbstractArray, x, index, count = N) where {T,V,N} idxs = Iterators.take(Iterators.drop(structural_eachindex(duals, x), index - 1), count) - return _seed_zero_partials!(duals, x, idxs) + return _seed_zero_partials!(Dual{T,V,N}, duals, x, idxs) end -function _seed_zero_partials!(duals::AbstractArray{Dual{T,V,N}}, x, idxs) where {T,V,N} +function _seed_zero_partials!(::Type{Dual{T,V,N}}, duals::AbstractArray, x, idxs) where {T,V,N} seed = zero(Partials{N,V}) if isbitstype(V) for idx in idxs - duals[idx] = Dual{T,V,N}(x[idx], seed) + duals[idx] = seed_dual(T, x[idx], seed) end else for idx in idxs if isassigned(x, idx) - duals[idx] = Dual{T,V,N}(x[idx], seed) + duals[idx] = seed_dual(T, x[idx], seed) else Base._unsetindex!(duals, idx) end @@ -104,16 +123,16 @@ function _seed_zero_partials!(duals::AbstractArray{Dual{T,V,N}}, x, idxs) where return duals end -function seed!(duals::AbstractArray{Dual{T,V,N}}, x, +function seed!(::Type{Dual{T,V,N}}, duals::AbstractArray, x, seeds::NTuple{N,Partials{N,V}}) where {T,V,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]) + duals[idx] = seed_dual(T, x[idx], seeds[i]) end else for (i, idx) in zip(1:N, structural_eachindex(duals, x)) if isassigned(x, idx) - duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) + duals[idx] = seed_dual(T, x[idx], seeds[i]) else Base._unsetindex!(duals, idx) end @@ -122,18 +141,18 @@ function seed!(duals::AbstractArray{Dual{T,V,N}}, x, return duals end -function seed!(duals::AbstractArray{Dual{T,V,N}}, x, index, +function seed!(::Type{Dual{T,V,N}}, duals::AbstractArray, x, index, seeds::NTuple{N,Partials{N,V}}, chunksize = N) where {T,V,N} offset = index - 1 idxs = Iterators.drop(structural_eachindex(duals, x), offset) if isbitstype(V) for (i, idx) in zip(1:chunksize, idxs) - duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) + duals[idx] = seed_dual(T, x[idx], seeds[i]) end else for (i, idx) in zip(1:chunksize, idxs) if isassigned(x, idx) - duals[idx] = Dual{T,V,N}(x[idx], seeds[i]) + duals[idx] = seed_dual(T, x[idx], seeds[i]) else Base._unsetindex!(duals, idx) end diff --git a/src/config.jl b/src/config.jl index 1e0111c9..af032d9c 100644 --- a/src/config.jl +++ b/src/config.jl @@ -103,7 +103,7 @@ function DerivativeConfig(f::F, y::AbstractArray{Y}, x::X, tag::T = Tag(f, X)) where {F,X<:Real,Y<:Real,T} - duals = similar(y, Dual{T,Y,1}) + duals = similar(y, seed_type(Dual{T,Y,1})) return DerivativeConfig{T,typeof(duals)}(duals) end @@ -139,7 +139,7 @@ function GradientConfig(f::F, ::Chunk{N} = Chunk(x), ::T = Tag(f, V)) where {F,V,N,T} seeds = construct_seeds(Partials{N,V}) - duals = similar(x, Dual{T,V,N}) + duals = similar(x, seed_type(Dual{T,V,N})) return GradientConfig{T,V,N,typeof(duals)}(seeds, duals) end @@ -176,7 +176,7 @@ function JacobianConfig(f::F, ::Chunk{N} = Chunk(x), ::T = Tag(f, V)) where {F,V,N,T} seeds = construct_seeds(Partials{N,V}) - duals = similar(x, Dual{T,V,N}) + duals = similar(x, seed_type(Dual{T,V,N})) return JacobianConfig{T,V,N,typeof(duals)}(seeds, duals) end @@ -202,8 +202,8 @@ function JacobianConfig(f::F, ::Chunk{N} = Chunk(x), ::T = Tag(f, X)) where {F,Y,X,N,T} seeds = construct_seeds(Partials{N,X}) - yduals = similar(y, Dual{T,Y,N}) - xduals = similar(x, Dual{T,X,N}) + yduals = similar(y, seed_type(Dual{T,Y,N})) + xduals = similar(x, seed_type(Dual{T,X,N})) duals = (yduals, xduals) return JacobianConfig{T,X,N,typeof(duals)}(seeds, duals) end @@ -215,9 +215,9 @@ Base.eltype(::Type{JacobianConfig{T,V,N,D}}) where {T,V,N,D} = Dual{T,V,N} # HessianConfig # ################# -struct HessianConfig{T,V,N,DG,DJ,TG} <: AbstractConfig{N} +struct HessianConfig{T,V,N,DJ,G<:GradientConfig} <: AbstractConfig{N} jacobian_config::JacobianConfig{T,V,N,DJ} - gradient_config::GradientConfig{TG,Dual{T,V,N},N,DG} + gradient_config::G end """ @@ -273,5 +273,4 @@ function HessianConfig(f::F, end checktag(::HessianConfig{T},f,x) where {T} = checktag(T,f,x) -Base.eltype(::Type{HessianConfig{T,V,N,DG,DJ,TG}}) where {T,V,N,DG,DJ,TG} = - Dual{TG,Dual{T,V,N},N} +Base.eltype(::Type{HessianConfig{T,V,N,DJ,G}}) where {T,V,N,DJ,G} = eltype(G) diff --git a/src/derivative.jl b/src/derivative.jl index 892f21a7..1e0d337b 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_zero_partials!(ydual, y) + seed_zero_partials!(Dual{T,eltype(y),1}, ydual, y) f!(ydual, Dual{T}(x, one(x))) map!(d -> value(T, d), 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_zero_partials!(ydual, y) + seed_zero_partials!(Dual{T,eltype(y),1}, 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 a74d1b78..96e1b930 100644 --- a/src/gradient.jl +++ b/src/gradient.jl @@ -139,24 +139,24 @@ function chunk_mode_gradient_expr(result_definition::Expr) # do first chunk manually to calculate output type. Seeding the first chunk and zeroing the # remaining elements partitions `xdual`, so every element is initialized exactly once. - seed!(xdual, x, 1, seeds) - seed_zero_partials!(xdual, x, N + 1, xlen - N) + seed!(Dual{T,V,N}, xdual, x, 1, seeds) + seed_zero_partials!(Dual{T,V,N}, xdual, x, N + 1, xlen - N) ydual = f(xdual) $(result_definition) extract_gradient_chunk!(T, result, ydual, 1, N) - seed_zero_partials!(xdual, x, 1) + seed_zero_partials!(Dual{T,V,N}, xdual, x, 1) # do middle chunks for c in middlechunks i = ((c - 1) * N + 1) - seed!(xdual, x, i, seeds) + seed!(Dual{T,V,N}, xdual, x, i, seeds) ydual = f(xdual) extract_gradient_chunk!(T, result, ydual, i, N) - seed_zero_partials!(xdual, x, i) + seed_zero_partials!(Dual{T,V,N}, xdual, x, i) end # do final chunk - seed!(xdual, x, lastchunkindex, seeds, lastchunksize) + seed!(Dual{T,V,N}, xdual, x, lastchunkindex, seeds, lastchunksize) ydual = f(xdual) extract_gradient_chunk!(T, result, ydual, lastchunkindex, lastchunksize) diff --git a/src/jacobian.jl b/src/jacobian.jl index f14a6a7b..67bf7f73 100644 --- a/src/jacobian.jl +++ b/src/jacobian.jl @@ -186,26 +186,26 @@ function jacobian_chunk_mode_expr(work_array_definition::Expr, compute_ydual::Ex # do first chunk manually to calculate output type. Seeding the first chunk and zeroing the # remaining elements partitions `xdual`, so every element is initialized exactly once. - seed!(xdual, x, 1, seeds) - seed_zero_partials!(xdual, x, N + 1, xlen - N) + seed!(Dual{T,V,N}, xdual, x, 1, seeds) + seed_zero_partials!(Dual{T,V,N}, xdual, x, N + 1, xlen - N) $(compute_ydual) ydual isa AbstractArray || throw(JACOBIAN_ERROR) $(result_definition) out_reshaped = reshape_jacobian(result, ydual, xdual) extract_jacobian_chunk!(T, out_reshaped, ydual, 1, N) - seed_zero_partials!(xdual, x, 1) + seed_zero_partials!(Dual{T,V,N}, xdual, x, 1) # do middle chunks for c in middlechunks i = ((c - 1) * N + 1) - seed!(xdual, x, i, seeds) + seed!(Dual{T,V,N}, xdual, x, i, seeds) $(compute_ydual) extract_jacobian_chunk!(T, out_reshaped, ydual, i, N) - seed_zero_partials!(xdual, x, i) + seed_zero_partials!(Dual{T,V,N}, xdual, x, i) end # do final chunk - seed!(xdual, x, lastchunkindex, seeds, lastchunksize) + seed!(Dual{T,V,N}, xdual, x, lastchunkindex, seeds, lastchunksize) $(compute_ydual) extract_jacobian_chunk!(T, out_reshaped, ydual, lastchunkindex, lastchunksize) @@ -224,7 +224,7 @@ end @eval function chunk_mode_jacobian(f!::F, y, x, cfg::JacobianConfig{T,V,N}) where {F,T,V,N} $(jacobian_chunk_mode_expr(:((ydual, xdual) = cfg.duals), - :(f!(seed_zero_partials!(ydual, y), xdual)), + :(f!(seed_zero_partials!(Dual{T,eltype(y),N}, ydual, y), xdual)), :(result = similar(y, length(y), xlen)), :(map!(d -> value(T,d), y, ydual)))) end @@ -238,7 +238,7 @@ 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(:((ydual, xdual) = cfg.duals), - :(f!(seed_zero_partials!(ydual, y), xdual)), + :(f!(seed_zero_partials!(Dual{T,eltype(y),N}, ydual, y), xdual)), :(), :(extract_value!(T, result, y, ydual)))) end diff --git a/test/AllocationsTest.jl b/test/AllocationsTest.jl index 5c3a459f..1f61ac86 100644 --- a/test/AllocationsTest.jl +++ b/test/AllocationsTest.jl @@ -12,24 +12,25 @@ convert_test_574() = convert(ForwardDiff.Dual{ForwardDiff.Tag{Nothing,D2_574},D2 @testset "Test seed!/seed_zero_partials! allocations" begin x = rand(1000) cfg = ForwardDiff.GradientConfig(nothing, x) - duals = cfg.duals + D, duals = eltype(cfg), cfg.duals seeds = cfg.seeds - allocs_seed!(args...) = @allocated ForwardDiff.seed!(args...) - allocs_seed!(duals, x, seeds) - @test iszero(allocs_seed!(duals, x, seeds)) - allocs_seed!(duals, x, 1, seeds) - @test iszero(allocs_seed!(duals, x, 1, seeds)) + allocs_seed!(::Type{D}, duals, x, seeds) where {D} = @allocated ForwardDiff.seed!(D, duals, x, seeds) + allocs_seed!(::Type{D}, duals, x, index, seeds) where {D} = @allocated ForwardDiff.seed!(D, duals, x, index, seeds) + allocs_seed!(D, duals, x, seeds) + @test iszero(allocs_seed!(D, duals, x, seeds)) + allocs_seed!(D, duals, x, 1, seeds) + @test iszero(allocs_seed!(D, duals, x, 1, seeds)) - # the 4-arg form passes `count` as a runtime value, so it catches an inference regression at the + # the form with `count` passes `count` as a runtime value, so it catches an inference regression at the # `_seed_zero_partials!` boundary that the forms defaulting `count` to `N` could hide - allocs_szp!(args...) = @allocated ForwardDiff.seed_zero_partials!(args...) - allocs_szp!(duals, x) - @test iszero(allocs_szp!(duals, x)) - allocs_szp!(duals, x, 1) - @test iszero(allocs_szp!(duals, x, 1)) - allocs_szp!(duals, x, 1, 4) - @test iszero(allocs_szp!(duals, x, 1, 4)) + allocs_szp!(::Type{D}, args...) where {D} = @allocated ForwardDiff.seed_zero_partials!(D, args...) + allocs_szp!(D, duals, x) + @test iszero(allocs_szp!(D, duals, x)) + allocs_szp!(D, duals, x, 1) + @test iszero(allocs_szp!(D, duals, x, 1)) + allocs_szp!(D, duals, x, 1, 4) + @test iszero(allocs_szp!(D, duals, x, 1, 4)) allocs_convert_test_574() = @allocated convert_test_574() allocs_convert_test_574() diff --git a/test/ConfusionTest.jl b/test/ConfusionTest.jl index c5b66c90..fbb7138a 100644 --- a/test/ConfusionTest.jl +++ b/test/ConfusionTest.jl @@ -137,6 +137,19 @@ let f = x -> x[1]^2 * x[2], x = [3.0, 2.0] @test ForwardDiff.hessian(f, x, ForwardDiff.HessianConfig(nothing, x)) == [4.0 6.0; 6.0 0.0] end +# Seeding keeps greater tags of the input outside, e.g. if the input type is abstract +let f = x -> x[1]^2 * x[2], g! = (y, x) -> (y[1] = f(x); y) + @test ForwardDiff.derivative(t -> ForwardDiff.gradient(f, Real[t, 2t])[1], 1.0) == 8.0 + @test ForwardDiff.derivative(t -> ForwardDiff.gradient(f, Real[t, 2t], ForwardDiff.GradientConfig(f, Real[t, 2t], ForwardDiff.Chunk{1}()))[1], 1.0) == 8.0 + @test ForwardDiff.derivative(t -> ForwardDiff.jacobian(x -> [f(x)], Real[t, 2t])[1, 1], 1.0) == 8.0 + @test ForwardDiff.derivative(t -> ForwardDiff.jacobian(g!, Real[0.0], Real[t, 2t])[1, 1], 1.0) == 8.0 + @test ForwardDiff.derivative(t -> ForwardDiff.jacobian(g!, Real[t], Real[t, 2t])[1, 1], 1.0) == 8.0 + @test ForwardDiff.derivative(t -> ForwardDiff.jacobian(g!, Real[0.0], Real[t, 2t], ForwardDiff.JacobianConfig(g!, Real[0.0], Real[t, 2t], ForwardDiff.Chunk{1}()))[1, 1], 1.0) == 8.0 + @test ForwardDiff.derivative(t -> ForwardDiff.hessian(f, Real[t, 2t])[1, 1], 1.0) == 4.0 + x = [ForwardDiff.Dual{BTag}(1.0, 1.0), ForwardDiff.Dual{BTag}(2.0, 2.0)] + @test ForwardDiff.gradient(f, x, ForwardDiff.GradientConfig(f, x, ForwardDiff.Chunk{2}(), ATag())) == [ForwardDiff.Dual{BTag}(4.0, 8.0), ForwardDiff.Dual{BTag}(1.0, 2.0)] +end + # Issue #845: all Hessian paths agree with the Jacobian of the gradient strip_outer(x) = x strip_outer(d::ForwardDiff.Dual{T}) where {T} = ForwardDiff.value(T, d) diff --git a/test/SeedTest.jl b/test/SeedTest.jl index ace810ae..e769574b 100644 --- a/test/SeedTest.jl +++ b/test/SeedTest.jl @@ -44,7 +44,7 @@ end @testset "seed_zero_partials!: $(nameof(typeof(x)))" for (x, sidx) in SEED_CASES cfg = ForwardDiff.GradientConfig(nothing, x, ForwardDiff.Chunk{3}()) - duals, seeds = cfg.duals, cfg.seeds + D, duals, seeds = eltype(cfg), cfg.duals, cfg.seeds N = ForwardDiff.npartials(ForwardDiff.Tag{Nothing,Float64}, eltype(duals)) marker = Partials(ntuple(i -> Float64(i), N)) nstruct = length(sidx) @@ -55,7 +55,7 @@ end # `count` defaults to N fill_marker!(duals, x, sidx, marker) - ForwardDiff.seed_zero_partials!(duals, x, 4) + ForwardDiff.seed_zero_partials!(D, duals, x, 4) @test zeroed_positions(duals, sidx) == collect(4:(4 + N - 1)) @test values_match(duals, x) @@ -67,24 +67,24 @@ end (nstruct - 1, N, (nstruct - 1):nstruct), (1, 0, 1:0)) fill_marker!(duals, x, sidx, marker) - ForwardDiff.seed_zero_partials!(duals, x, index, count) + ForwardDiff.seed_zero_partials!(D, duals, x, index, count) @test zeroed_positions(duals, sidx) == collect(expected) @test values_match(duals, x) end - # the 2-arg form clears every structural position + # the form without `index` clears every structural position fill_marker!(duals, x, sidx, marker) - ForwardDiff.seed_zero_partials!(duals, x) + ForwardDiff.seed_zero_partials!(D, duals, x) @test zeroed_positions(duals, sidx) == collect(1:nstruct) @test values_match(duals, x) # `seed!` and `seed_zero_partials!` must agree on what "the chunk at `index`" is, or chunk mode # would leave stale seeds behind. `duals` enters each iteration fully cleared. @testset "round-trips seed! at index=$index" for index in unique((1, 4, nstruct - N + 1)) - ForwardDiff.seed!(duals, x, index, seeds) + ForwardDiff.seed!(D, duals, x, index, seeds) @test zeroed_positions(duals, sidx) == [i for i in 1:nstruct if !(index <= i <= index + N - 1)] - ForwardDiff.seed_zero_partials!(duals, x, index) + ForwardDiff.seed_zero_partials!(D, duals, x, index) @test zeroed_positions(duals, sidx) == collect(1:nstruct) @test values_match(duals, x) end From 8fdf29f50fc316f4bf3fdfd7a775553e9741c587 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?David=20M=C3=BCller-Widmann?= Date: Tue, 29 Sep 2026 23:45:08 +0200 Subject: [PATCH 8/8] Require a concrete value type in `Dual`s With an abstract value type, the components of a `Dual` can carry different tag layers, so the type no longer determines where a tag occurs. Seeding now converts the seeds to the type of each element, so elements of abstract input arrays are seeded with concrete value types. Co-Authored-By: Claude Opus 5.5 --- src/apiutils.jl | 17 ++++++++--------- src/dual.jl | 12 ++++++++---- test/ConfusionTest.jl | 2 -- test/DualTest.jl | 4 ++-- 4 files changed, 18 insertions(+), 17 deletions(-) diff --git a/src/apiutils.jl b/src/apiutils.jl index 0d5401ec..d23148f5 100644 --- a/src/apiutils.jl +++ b/src/apiutils.jl @@ -40,21 +40,20 @@ end return Expr(:tuple, [:(single_seed(Partials{N,V}, Val{$i}())) for i in 1:N]...) end -# Seeds `x` with tag `T` and partials `p`. Layers of `x` with greater tags are kept outside, so -# nested `Dual`s stay sorted even if `T` is not greater than all tags in `x`. -@inline seed_dual(::Type{T}, x, p::Partials) where {T} = Dual{T}(x, p) -@inline function seed_dual(::Type{T}, x::Dual{S}, p::Partials) where {T,S} +# Seeds `x` with tag `T` and partials `p`, converted to the type of `x`. Layers of `x` with greater +# tags are kept outside, so nested `Dual`s stay sorted even if `T` is not greater than all tags in `x`. +@inline seed_dual(::Type{T}, x, p::Partials{N}) where {T,N} = Dual{T}(x, convert(Partials{N,typeof(x)}, p)) +@inline function seed_dual(::Type{T}, x::Dual{S}, p::Partials{N}) where {T,S,N} + p = convert(Partials{N,typeof(x)}, p) T ≺ S || return Dual{T}(x, p) # the seeds are constants, so their partials w.r.t. `S` are zero q = map_partials(y -> value(S, y), valtype(S, eltype(p)), p) return Dual{S}(seed_dual(T, value(S, x), q), map(y -> seed_dual(T, y, zero(q)), partials(S, x).values)) end -# Type of `seed_dual(T, x, p)` for `x::V` and `p::Partials{N,V}`. If `V` is abstract and can contain -# `Dual`s, their tags may be greater than `T`, so only `Real` is a bound. -function seed_type(::Type{Dual{T,V,N}}) where {T,V,N} - return isconcretetype(V) || typeintersect(V, Dual) === Union{} ? Dual{T,V,N} : Real -end +# Type of `seed_dual(T, x, p)` for `x::V` and `p::Partials{N,V}`. If `V` is abstract, each element is +# seeded with its own type, so only `Real` is a bound. +seed_type(::Type{Dual{T,V,N}}) where {T,V,N} = isconcretetype(V) ? Dual{T,V,N} : Real function seed_type(::Type{Dual{T,Dual{S,W,M},N}}) where {T,S,W,M,N} return T ≺ S ? Dual{S,seed_type(Dual{T,W,N}),M} : Dual{T,Dual{S,W,M},N} end diff --git a/src/dual.jl b/src/dual.jl index 63991879..b6e6170f 100644 --- a/src/dual.jl +++ b/src/dual.jl @@ -16,8 +16,8 @@ struct Dual{T,V,N} <: Real partials::Partials{N,V} function Dual{T, V, N}(value::V, partials::Partials{N, V}) where {T, V, N} T isa Type || throw_invalid_tag(T) - check_tag_order(T, value) - foreach(p -> check_tag_order(T, p), partials.values) + isconcretetype(V) || throw_abstract_value(V) + check_tag_order(T, V) can_dual(V) || throw_cannot_dual(V) new{T, V, N}(value, partials) end @@ -36,6 +36,10 @@ Base.ArithmeticStyle(::Type{<:Dual{T,V}}) where {T,V} = Base.ArithmeticStyle(V) throw(ArgumentError(lazy"The tag of a Dual must be a type, got $(repr(T)).")) end +@noinline function throw_abstract_value(V::Type) + throw(ArgumentError(lazy"The value type of a Dual must be concrete, got $V.")) +end + @noinline function throw_cannot_dual(V::Type) throw(ArgumentError(lazy"Cannot create a dual over scalar type $V. If the type behaves as a scalar, define ForwardDiff.can_dual(::Type{$V}) = true.")) end @@ -59,8 +63,8 @@ end end # Tags are strictly decreasing inwards, hence unique, so it suffices to check the next layer -@inline check_tag_order(::Type{T}, x) where {T} = nothing -@inline function check_tag_order(::Type{T}, ::Dual{S}) where {T,S} +@inline check_tag_order(::Type{T}, ::Type) where {T} = nothing +@inline function check_tag_order(::Type{T}, ::Type{<:Dual{S}}) where {T,S} if S === T throw_same_tag(T) elseif T ≺ S diff --git a/test/ConfusionTest.jl b/test/ConfusionTest.jl index fbb7138a..2c59ab3f 100644 --- a/test/ConfusionTest.jl +++ b/test/ConfusionTest.jl @@ -119,8 +119,6 @@ struct ATag end struct BTag end @test ForwardDiff.Dual{BTag}(ForwardDiff.Dual{ATag}(1.0, 2.0), ForwardDiff.Dual{ATag}(3.0, 4.0)) isa ForwardDiff.Dual{BTag} @test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag}(ForwardDiff.Dual{BTag}(1.0, 2.0), ForwardDiff.Dual{BTag}(3.0, 4.0)) -@test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag,Real,1}(ForwardDiff.Dual{BTag}(1.0, 2.0), ForwardDiff.Partials{1,Real}((3.0,))) -@test_throws ArgumentError("Cannot store a Dual with tag $ATag outside a Dual with tag $BTag, since $ATag ≺ $BTag.") ForwardDiff.Dual{ATag,Real,1}(1.0, ForwardDiff.Partials{1,Real}((ForwardDiff.Dual{BTag}(3.0, 4.0),))) # Tags of nested `Dual`s are unique @test_throws ArgumentError("Cannot store a Dual with tag $ATag inside a Dual with the same tag.") ForwardDiff.Dual{ATag}(ForwardDiff.Dual{ATag}(1.0, 2.0), ForwardDiff.Dual{ATag}(3.0, 4.0)) diff --git a/test/DualTest.jl b/test/DualTest.jl index 87203db4..87f2d20f 100644 --- a/test/DualTest.jl +++ b/test/DualTest.jl @@ -442,8 +442,8 @@ ForwardDiff.:≺(::Type{OuterTestTag}, ::Type{TestTag}) = false if M != N @test promote_type(Dual{TestTag,V,M}, Dual{TestTag,V,N}) == Dual{TestTag,V} end - @test zero(Dual{TestTag,Real,N}) isa Dual{TestTag,Real,N} - @test value(TestTag, convert(Dual{TestTag,Real,N}, PRIMAL)) == PRIMAL + @test_throws ArgumentError("The value type of a Dual must be concrete, got Real.") convert(Dual{TestTag,Real,N}, PRIMAL) + @test_throws ArgumentError("The value type of a Dual must be concrete, got Real.") Dual{TestTag}(PRIMAL, Partials{N,Real}(ntuple(i -> PRIMAL, N))) # issue #322 @test promote_type(Bool, Dual{TestTag,V,N}) == Dual{TestTag,promote_type(Bool, V),N}