From c2754994f689d8c97be7c763b6141bac2814a52e Mon Sep 17 00:00:00 2001 From: Alex Robson Date: Wed, 10 Nov 2021 19:58:06 +0000 Subject: [PATCH 1/6] Bump project - check if breaking --- Project.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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" From 48bee81a0fd5c7e77a018bfa43d235a90249a86e Mon Sep 17 00:00:00 2001 From: Alex Robson Date: Wed, 10 Nov 2021 21:48:21 +0000 Subject: [PATCH 2/6] Remove changes. Add correct flatten function for Array{T} --- src/flatten.jl | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) 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 From 9926cf0cb7ce873c439678dd633c8c3584141b74 Mon Sep 17 00:00:00 2001 From: Alex Robson Date: Wed, 10 Nov 2021 21:48:32 +0000 Subject: [PATCH 3/6] Tidy up test utils --- src/test_utils.jl | 4 ++++ 1 file changed, 4 insertions(+) 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 From cf662a89ea834db2e25d0cd1c8ee577179193005 Mon Sep 17 00:00:00 2001 From: Alex Robson Date: Wed, 10 Nov 2021 21:52:02 +0000 Subject: [PATCH 4/6] Update deferred? --- src/parameters.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/parameters.jl b/src/parameters.jl index a0e18b9..454e5aa 100644 --- a/src/parameters.jl +++ b/src/parameters.jl @@ -148,9 +148,9 @@ Base.:(==)(a::Deferred, b::Deferred) = (a.f == b.f) && (a.args == b.args) value(x::Deferred) = x.f(value(x.args)...) -function flatten(::Type{T}, x::Deferred) where {T<:Real} +function flatten(::Type{T}, x::D) where {T<:Real, D<:Deferred} v, unflatten = flatten(T, x.args) - unflatten_Deferred(v_new::Vector{T}) = Deferred(x.f, unflatten(v_new)) + unflatten_Deferred(v_new::Vector{T}) = D(x.f, unflatten(v_new)) return v, unflatten_Deferred end From c8577bed0db9581bc7fe3f52980dbb659152206e Mon Sep 17 00:00:00 2001 From: Alex Robson Date: Thu, 11 Nov 2021 09:53:02 +0000 Subject: [PATCH 5/6] Just set no infer for Deferred --- src/parameters.jl | 4 ++-- test/parameters.jl | 3 ++- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/src/parameters.jl b/src/parameters.jl index 454e5aa..a0e18b9 100644 --- a/src/parameters.jl +++ b/src/parameters.jl @@ -148,9 +148,9 @@ Base.:(==)(a::Deferred, b::Deferred) = (a.f == b.f) && (a.args == b.args) value(x::Deferred) = x.f(value(x.args)...) -function flatten(::Type{T}, x::D) where {T<:Real, D<:Deferred} +function flatten(::Type{T}, x::Deferred) where {T<:Real} v, unflatten = flatten(T, x.args) - unflatten_Deferred(v_new::Vector{T}) = D(x.f, unflatten(v_new)) + unflatten_Deferred(v_new::Vector{T}) = Deferred(x.f, unflatten(v_new)) return v, unflatten_Deferred end diff --git a/test/parameters.jl b/test/parameters.jl index fff9268..4340a5b 100644 --- a/test/parameters.jl +++ b/test/parameters.jl @@ -39,11 +39,12 @@ pdiagmat(args...) = PDiagMat(args...) @testset "deferred" begin test_parameter_interface(deferred(sin, 0.5); check_inferred=tuple_infers) test_parameter_interface(deferred(sin, positive(0.5)); check_inferred=tuple_infers) + test_parameter_interface( deferred( mvnormal, fixed(randn(5)), deferred(pdiagmat, positive.(rand(5) .+ 1e-1)) ); - check_inferred=tuple_infers, + check_inferred=false, # flatten infers, unflatten doesn't ) end From 273f468d149f5cde49ad8cc479c93c8213f6f474 Mon Sep 17 00:00:00 2001 From: Alex Robson Date: Tue, 16 Nov 2021 15:02:05 +0000 Subject: [PATCH 6/6] inline deferred --- src/parameters.jl | 2 +- test/parameters.jl | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) 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/test/parameters.jl b/test/parameters.jl index 4340a5b..fff9268 100644 --- a/test/parameters.jl +++ b/test/parameters.jl @@ -39,12 +39,11 @@ pdiagmat(args...) = PDiagMat(args...) @testset "deferred" begin test_parameter_interface(deferred(sin, 0.5); check_inferred=tuple_infers) test_parameter_interface(deferred(sin, positive(0.5)); check_inferred=tuple_infers) - test_parameter_interface( deferred( mvnormal, fixed(randn(5)), deferred(pdiagmat, positive.(rand(5) .+ 1e-1)) ); - check_inferred=false, # flatten infers, unflatten doesn't + check_inferred=tuple_infers, ) end