diff --git a/Project.toml b/Project.toml index a1322d7..c1cec25 100644 --- a/Project.toml +++ b/Project.toml @@ -1,7 +1,7 @@ name = "ParameterHandling" uuid = "2412ca09-6db7-441c-8e3a-88d5709968c5" authors = ["Invenia Technical Computing Corporation"] -version = "0.4.0" +version = "0.4.1" [deps] ChainRulesCore = "d360d2e6-b24c-11e9-a2a3-2a2ae2dbcce4" diff --git a/src/flatten.jl b/src/flatten.jl index 717b27b..3859a7a 100644 --- a/src/flatten.jl +++ b/src/flatten.jl @@ -31,7 +31,10 @@ function flatten(::Type{T}, x::R) where {T<:Real,R<:Real} return v, unflatten_to_Real end -flatten(::Type{T}, x::Vector{R}) where {T<:Real,R<:Real} = (Vector{T}(x), Vector{R}) +function flatten(::Type{T}, x::Vector{R}) where {T<:Real,R<:Real} + unflatten_to_Vector(v) = Vector{R}(v) + return Vector{T}(x), unflatten_to_Vector +end function _flatten_vector_integer(::Type{T}, x::AbstractVector{<:Integer}) where {T<:Real} unflatten_to_Vector_Integer(x_vec) = x diff --git a/src/parameters.jl b/src/parameters.jl index a0e18b9..cf782a2 100644 --- a/src/parameters.jl +++ b/src/parameters.jl @@ -150,7 +150,7 @@ value(x::Deferred) = x.f(value(x.args)...) function flatten(::Type{T}, x::Deferred) where {T<:Real} v, unflatten = flatten(T, x.args) - unflatten_Deferred(v_new::Vector{T}) = Deferred(x.f, unflatten(v_new)) + @inline unflatten_Deferred(v_new::Vector{T}) = Deferred(x.f, unflatten(v_new)) return v, unflatten_Deferred end diff --git a/src/test_utils.jl b/src/test_utils.jl index 22ce3d3..2a1e5fa 100644 --- a/src/test_utils.jl +++ b/src/test_utils.jl @@ -43,6 +43,7 @@ function test_flatten_interface(x::T; check_inferred::Bool=true) where {T} # Check that everything infers properly. check_inferred && @inferred flatten(x) + check_inferred && @inferred unflatten(v) # Test with different precisions @testset "Float64" begin @@ -54,6 +55,7 @@ function test_flatten_interface(x::T; check_inferred::Bool=true) where {T} # Check that everything infers properly. check_inferred && @inferred flatten(Float64, x) + check_inferred && @inferred _unflatten(_v) end @testset "Float32" begin _v, _unflatten = flatten(Float32, x) @@ -63,6 +65,7 @@ function test_flatten_interface(x::T; check_inferred::Bool=true) where {T} # Check that everything infers properly. check_inferred && @inferred flatten(Float32, x) + check_inferred && @inferred _unflatten(_v) end @testset "Float16" begin _v, _unflatten = flatten(Float16, x) @@ -72,6 +75,7 @@ function test_flatten_interface(x::T; check_inferred::Bool=true) where {T} # Check that everything infers properly. check_inferred && @inferred flatten(Float16, x) + check_inferred && @inferred _unflatten(_v) end end