From e6254b6553ed8bef32c03fd7bdc579bdd6c4fe76 Mon Sep 17 00:00:00 2001 From: Miha Zgubic Date: Tue, 17 May 2022 15:48:31 +0100 Subject: [PATCH 1/5] Draft Composite transform --- src/FeatureTransforms.jl | 2 ++ src/composite.jl | 57 ++++++++++++++++++++++++++++++++++++++++ src/traits.jl | 10 +++++++ test/composite.jl | 45 +++++++++++++++++++++++++++++++ test/runtests.jl | 1 + test/traits.jl | 22 ++++++++++++++++ 6 files changed, 137 insertions(+) create mode 100644 src/composite.jl create mode 100644 test/composite.jl diff --git a/src/FeatureTransforms.jl b/src/FeatureTransforms.jl index 03e3623..5e81e15 100644 --- a/src/FeatureTransforms.jl +++ b/src/FeatureTransforms.jl @@ -9,6 +9,7 @@ export Transform, transform, transform! export HoD, LinearCombination, OneHotEncoding, Periodic, Power export AbstractScaling, IdentityScaling, MeanStdScaling, StandardScaling export LogTransform, InverseHyperbolicSine +export Composite include("utils.jl") include("traits.jl") @@ -24,6 +25,7 @@ include("periodic.jl") include("power.jl") include("scaling.jl") include("temporal.jl") +include("composite.jl") include("test_utils.jl") diff --git a/src/composite.jl b/src/composite.jl new file mode 100644 index 0000000..4ff17db --- /dev/null +++ b/src/composite.jl @@ -0,0 +1,57 @@ +""" + Composite <: Transform + +A `Composite` transform is a composition of `Transform`s, currently limited to `OneToOne()` +cardinality. It can be fit and applied in a single step. + +The transforms in `Composite([t1, t2, t3])` are applied in `t1`, `t2`, `t3` order, where +the output of `t1` is the input to `t2` etc. When using `∘` to create transforms, the order +is `t3 ∘ t2 ∘ t1`, as in function composition. + +```jldoctest composite +julia> id = IdentityScaling(); + +julia> power = Power(2.0); + +julia> id ∘ power == Composite([power, id]) +true +``` +""" +struct Composite <: Transform + transforms::Vector{<:Transform} + + function Composite(transforms::Vector{<:Transform}) + if all(==(OneToOne()), map(cardinality, transforms)) + return new(transforms) + else + throw(ArgumentError("Only OneToOne() transforms are supported.")) + end + end +end + +cardinality(c::Composite) = ∘(map(cardinality, c.transforms)...) + +function fit!(c::Composite, data; kwargs...) + for t in c.transforms + fit!(t, data,; kwargs...) + data = t(data) + end + return c +end + +function _apply(x, c::Composite; kwargs...) + data = deepcopy(x) + for t in c.transforms + data = _apply(data, t; kwargs...) + end + return data +end + +# creating composite transforms: reverse the order so that c.transforms[1] is the first +# transforms that gets applied +Base.:(∘)(f::Transform, g::Transform) = Composite([g, f]) +Base.:(∘)(c::Composite, t::Transform) = Composite([t, c.transforms...]) +Base.:(∘)(t::Transform, c::Composite) = Composite([c.transforms..., t]) +Base.:(∘)(c::Composite, c2::Composite) = Composite([c2.transforms..., c.transforms...]) + +Base.:(==)(c::Composite, d::Composite) = return all(map(==, c.transforms, d.transforms)) diff --git a/src/traits.jl b/src/traits.jl index 7872985..ec9595d 100644 --- a/src/traits.jl +++ b/src/traits.jl @@ -46,3 +46,13 @@ struct ManyToMany <: Cardinality end Returns the [`Cardinality`](@ref) of the `transform`. """ function cardinality end + +Base.:(∘)(::OneToOne, ::OneToOne) = OneToOne() +Base.:(∘)(::OneToMany, ::OneToOne) = OneToMany() +Base.:(∘)(::ManyToOne, ::OneToMany) = OneToOne() +Base.:(∘)(::ManyToMany, ::OneToMany) = OneToMany() +Base.:(∘)(::OneToOne, ::ManyToOne) = ManyToOne() +Base.:(∘)(::OneToMany, ::ManyToOne) = ManyToMany() +Base.:(∘)(::ManyToOne, ::ManyToMany) = ManyToOne() +Base.:(∘)(::ManyToMany, ::ManyToMany) = ManyToMany() +Base.:(∘)(c2::Cardinality, c1::Cardinality) = throw(ArgumentError("$c2 ∘ $c1 is undefined.")) diff --git a/test/composite.jl b/test/composite.jl new file mode 100644 index 0000000..27ba4b3 --- /dev/null +++ b/test/composite.jl @@ -0,0 +1,45 @@ +@testset "composite.jl" begin + @testset "constructor" begin + id = IdentityScaling() + power = Power(3.0) + logt = LogTransform() + @test id ∘ id == Composite([id, id]) + @test id ∘ id ∘ power == Composite([power, id, id]) + @test power ∘ id ∘ power == Composite([power, id, power]) + + @test power ∘ (id ∘ logt) == Composite([logt, id, power]) + @test (power ∘ id) ∘ logt == Composite([logt, id, power]) + @test (power ∘ id) ∘ (logt ∘ id) == Composite([id, logt, id, power]) + + @test_throws ArgumentError id ∘ LinearCombination([1, 2, 3]) + @test_throws ArgumentError OneHotEncoding([1, 2]) ∘ id + end + + @testset "apply" begin + p = Power(4.0) + c = Power(2.0) ∘ Power(2.0) ∘ IdentityScaling() + x = [1, 2, 3] + @test FeatureTransforms.apply(x, p) == FeatureTransforms.apply(x, c) + @test p(x) == c(x) + end + + @testset "fit!" begin + s = StandardScaling() + c = StandardScaling() ∘ IdentityScaling() ∘ StandardScaling() + x = rand(10) + x_copy = deepcopy(x) + + fit!(s, x) + fit!(c, x) + + @test c(x) ≈ s(x) + + # did not change the input data + @test x_copy == x + + # but make sure that it is fit and transformed on the already transformed data, in + # this case leaving the second scaling redundant, i.e. centered at 0.0 and std = 1.0 + @test isapprox(0.0, c.transforms[3].μ; atol=1e-15) + @test isapprox(1.0, c.transforms[3].σ; atol=1e-15) + end +end diff --git a/test/runtests.jl b/test/runtests.jl index 61a1d55..f4f2fb4 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -25,6 +25,7 @@ using TimeZones include("scaling.jl") include("temporal.jl") include("traits.jl") + include("composite.jl") include("test_utils.jl") include("types/tables.jl") diff --git a/test/traits.jl b/test/traits.jl index ce28e22..fb4df3c 100644 --- a/test/traits.jl +++ b/test/traits.jl @@ -2,4 +2,26 @@ for t in (OneToOne(), OneToMany(), ManyToOne(), ManyToMany()) @test t isa FeatureTransforms.Cardinality end + + @testset "composite" begin + @test OneToOne() == OneToOne() ∘ OneToOne() + @test OneToMany() == OneToMany() ∘ OneToOne() + @test_throws ArgumentError ManyToOne() ∘ OneToOne() + @test_throws ArgumentError ManyToMany() ∘ OneToOne() + + @test ManyToOne() == OneToOne() ∘ ManyToOne() + @test ManyToMany() == OneToMany() ∘ ManyToOne() + @test_throws ArgumentError ManyToOne() ∘ ManyToOne() + @test_throws ArgumentError ManyToMany() ∘ ManyToOne() + + @test_throws ArgumentError OneToOne() ∘ OneToMany() + @test_throws ArgumentError OneToMany() ∘ OneToMany() + @test OneToOne() == ManyToOne() ∘ OneToMany() + @test OneToMany() == ManyToMany() ∘ OneToMany() + + @test_throws ArgumentError OneToOne() ∘ ManyToMany() + @test_throws ArgumentError OneToMany() ∘ ManyToMany() + @test ManyToOne() == ManyToOne() ∘ ManyToMany() + @test ManyToMany() == ManyToMany() ∘ ManyToMany() + end end From 491b99b17daedbda997361a94b8060dd478766ca Mon Sep 17 00:00:00 2001 From: Miha Zgubic Date: Tue, 17 May 2022 15:50:54 +0100 Subject: [PATCH 2/5] typo and formatting --- src/composite.jl | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/src/composite.jl b/src/composite.jl index 4ff17db..f8d5af3 100644 --- a/src/composite.jl +++ b/src/composite.jl @@ -21,11 +21,8 @@ struct Composite <: Transform transforms::Vector{<:Transform} function Composite(transforms::Vector{<:Transform}) - if all(==(OneToOne()), map(cardinality, transforms)) - return new(transforms) - else - throw(ArgumentError("Only OneToOne() transforms are supported.")) - end + all(==(OneToOne()), map(cardinality, transforms)) && return new(transforms) + throw(ArgumentError("Only OneToOne() transforms are supported.")) end end @@ -33,7 +30,7 @@ cardinality(c::Composite) = ∘(map(cardinality, c.transforms)...) function fit!(c::Composite, data; kwargs...) for t in c.transforms - fit!(t, data,; kwargs...) + fit!(t, data; kwargs...) data = t(data) end return c From 557fe18cc9527fafa03eed69a88df9236d4ff098 Mon Sep 17 00:00:00 2001 From: Miha Zgubic Date: Tue, 17 May 2022 18:20:05 +0100 Subject: [PATCH 3/5] Apply suggestions from code review Co-authored-by: Glenn Moynihan --- src/composite.jl | 2 +- src/traits.jl | 4 +++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/src/composite.jl b/src/composite.jl index f8d5af3..18d3e9e 100644 --- a/src/composite.jl +++ b/src/composite.jl @@ -18,7 +18,7 @@ true ``` """ struct Composite <: Transform - transforms::Vector{<:Transform} + transforms::Tuple{Vararg{Transform}} function Composite(transforms::Vector{<:Transform}) all(==(OneToOne()), map(cardinality, transforms)) && return new(transforms) diff --git a/src/traits.jl b/src/traits.jl index ec9595d..1a03510 100644 --- a/src/traits.jl +++ b/src/traits.jl @@ -55,4 +55,6 @@ Base.:(∘)(::OneToOne, ::ManyToOne) = ManyToOne() Base.:(∘)(::OneToMany, ::ManyToOne) = ManyToMany() Base.:(∘)(::ManyToOne, ::ManyToMany) = ManyToOne() Base.:(∘)(::ManyToMany, ::ManyToMany) = ManyToMany() -Base.:(∘)(c2::Cardinality, c1::Cardinality) = throw(ArgumentError("$c2 ∘ $c1 is undefined.")) +function Base.:(∘)(c2::Cardinality, c1::Cardinality) + return throw(ArgumentError("Cannot compose cardinalities: $c2 ∘ $c1.")) +end From 705c30f45227c58bb2bcf052fb31f0b1beb49348 Mon Sep 17 00:00:00 2001 From: Miha Zgubic Date: Tue, 17 May 2022 18:29:25 +0100 Subject: [PATCH 4/5] change to tuple --- src/composite.jl | 8 ++++---- test/composite.jl | 12 ++++++------ 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/src/composite.jl b/src/composite.jl index 18d3e9e..3cc1888 100644 --- a/src/composite.jl +++ b/src/composite.jl @@ -46,9 +46,9 @@ end # creating composite transforms: reverse the order so that c.transforms[1] is the first # transforms that gets applied -Base.:(∘)(f::Transform, g::Transform) = Composite([g, f]) -Base.:(∘)(c::Composite, t::Transform) = Composite([t, c.transforms...]) -Base.:(∘)(t::Transform, c::Composite) = Composite([c.transforms..., t]) -Base.:(∘)(c::Composite, c2::Composite) = Composite([c2.transforms..., c.transforms...]) +Base.:(∘)(f::Transform, g::Transform) = Composite((g, f)) +Base.:(∘)(c::Composite, t::Transform) = Composite((t, c.transforms...)) +Base.:(∘)(t::Transform, c::Composite) = Composite((c.transforms..., t)) +Base.:(∘)(c::Composite, c2::Composite) = Composite((c2.transforms..., c.transforms...)) Base.:(==)(c::Composite, d::Composite) = return all(map(==, c.transforms, d.transforms)) diff --git a/test/composite.jl b/test/composite.jl index 27ba4b3..c998580 100644 --- a/test/composite.jl +++ b/test/composite.jl @@ -3,13 +3,13 @@ id = IdentityScaling() power = Power(3.0) logt = LogTransform() - @test id ∘ id == Composite([id, id]) - @test id ∘ id ∘ power == Composite([power, id, id]) - @test power ∘ id ∘ power == Composite([power, id, power]) + @test id ∘ id == Composite((id, id)) + @test id ∘ id ∘ power == Composite((power, id, id)) + @test power ∘ id ∘ power == Composite((power, id, power)) - @test power ∘ (id ∘ logt) == Composite([logt, id, power]) - @test (power ∘ id) ∘ logt == Composite([logt, id, power]) - @test (power ∘ id) ∘ (logt ∘ id) == Composite([id, logt, id, power]) + @test power ∘ (id ∘ logt) == Composite((logt, id, power)) + @test (power ∘ id) ∘ logt == Composite((logt, id, power)) + @test (power ∘ id) ∘ (logt ∘ id) == Composite((id, logt, id, power)) @test_throws ArgumentError id ∘ LinearCombination([1, 2, 3]) @test_throws ArgumentError OneHotEncoding([1, 2]) ∘ id From d18b668149c95e13ef43a6036b37bfb116a6f34b Mon Sep 17 00:00:00 2001 From: Miha Zgubic Date: Tue, 17 May 2022 18:32:31 +0100 Subject: [PATCH 5/5] change constructor to tuple --- src/composite.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/composite.jl b/src/composite.jl index 3cc1888..8fd2144 100644 --- a/src/composite.jl +++ b/src/composite.jl @@ -20,7 +20,7 @@ true struct Composite <: Transform transforms::Tuple{Vararg{Transform}} - function Composite(transforms::Vector{<:Transform}) + function Composite(transforms::Tuple{Vararg{Transform}}) all(==(OneToOne()), map(cardinality, transforms)) && return new(transforms) throw(ArgumentError("Only OneToOne() transforms are supported.")) end