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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -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"
Expand Down
5 changes: 4 additions & 1 deletion src/flatten.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/parameters.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
4 changes: 4 additions & 0 deletions src/test_utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -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

Expand Down