From 8697ddf4ab48b16da71062a59c3366d3821dbdd4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Tam=C3=A1s=20K=2E=20Papp?= Date: Fri, 7 Aug 2026 17:40:01 +0200 Subject: [PATCH] add support for bridge --- src/experimental.jl | 61 ++++++++++++++++++++++++++++++--------- test/Project.toml | 1 + test/test_experimental.jl | 15 ++++++++++ 3 files changed, 64 insertions(+), 13 deletions(-) diff --git a/src/experimental.jl b/src/experimental.jl index d39e168..b9c39ec 100644 --- a/src/experimental.jl +++ b/src/experimental.jl @@ -23,13 +23,15 @@ using Compat: @compat @compat public model_parameters_dimension, make_model_parameters, calculate_derived_quantities, make_approximation_basis, describe_policy_transformations, policy_coefficients_dimension, make_policy_functions, constant_initial_guess, - calculate_initial_guess, sum_of_squared_residuals + calculate_initial_guess, sum_of_squared_residuals, bridge -import ..SpectralKit +using ..SpectralKit: SpectralKit, Chebyshev, TransformedBasis, SmolyakBasis, PM1, + AbstractUnivariateTransformation, dimension, linear_combination, grid, domain, + transform_from, transform_to using ArgCheck: @argcheck using DocStringExtensions: FUNCTIONNAME, SIGNATURES -using InverseFunctions: inverse +import InverseFunctions #### #### utilities @@ -63,18 +65,18 @@ y` for all `x` in the domain. """ function constant_coefficients end -function constant_coefficients(basis::SpectralKit.Chebyshev, y) +function constant_coefficients(basis::Chebyshev, y) θ = zeros(basis.N) θ[1] = y θ end -function constant_coefficients(basis::SpectralKit.TransformedBasis, y) +function constant_coefficients(basis::TransformedBasis, y) constant_coefficients(parent(basis), y) end -function constant_coefficients(basis::SpectralKit.SmolyakBasis, y) - θ = zeros(SpectralKit.dimension(basis)) +function constant_coefficients(basis::SmolyakBasis, y) + θ = zeros(dimension(basis)) θ[1] = y θ end @@ -148,18 +150,18 @@ $(USERNOTE) function calculate_residuals end function policy_coefficients_dimension(policy_transformations::NamedTuple, approximation_basis) - SpectralKit.dimension(approximation_basis) * length(policy_transformations) + dimension(approximation_basis) * length(policy_transformations) end function make_policy_functions(model_family, policy_transformations::NamedTuple, approximation_basis, coefficients) - d = SpectralKit.dimension(approximation_basis) + d = dimension(approximation_basis) # QUESTION line below assumes all univariate, generalize? ranges = named_cumulative_ranges(map(_ -> d, policy_transformations)) @argcheck firstindex(coefficients) == 1 @argcheck lastindex(coefficients) == last(last(ranges)) map(ranges, policy_transformations) do r, t - t ∘ SpectralKit.linear_combination(approximation_basis, @view coefficients[r]) + t ∘ linear_combination(approximation_basis, @view coefficients[r]) end end @@ -184,14 +186,14 @@ function calculate_initial_guess(model_family, model_parameters, derived_quantit policy_transformations::NamedTuple{N}, approximation_basis) where N constant_guess = constant_initial_guess(model_family, model_parameters, derived_quantities) - d = SpectralKit.dimension(approximation_basis) + d = dimension(approximation_basis) # QUESTION line below assumes all univariate, generalize? ranges = named_cumulative_ranges(map(_ -> d, policy_transformations)) coefficients = zeros(last(last(ranges))) for (name, transformation) in pairs(policy_transformations) r = getproperty(ranges, name) transformed_y = getproperty(constant_guess, name) - y = inverse(transformation)(transformed_y) + y = InverseFunctions.inverse(transformation)(transformed_y) # FIXME a constant_coefficient! API would have fewer allocations coefficients[r] .= constant_coefficients(approximation_basis, y) end @@ -206,7 +208,7 @@ grid that corresponds to the approximation basis. """ function make_approximation_grid(model_family, model_parameters, approximation_parameters, derived_quantities, approximation_basis) - SpectralKit.grid(approximation_basis) + grid(approximation_basis) end """ @@ -230,4 +232,37 @@ function sum_of_squared_residuals(model_family, model_parameters, policy_functio end end +#### +#### bridge — expose the univariate algebraic transformations +#### + +""" +Implementation of [`bridge`](@ref). Not part of the (experimental) API. +""" +struct Bridge{O<:AbstractUnivariateTransformation,I<:AbstractUnivariateTransformation} + outer_transformation::O + inner_transformation::I +end + +""" +$(SIGNATURES) + +Transform using the `outer_transformation⁻¹ ∘ inner_transformation`. Return a callable +that supports [`InverseFunctions.inverse`](@ref). +""" +function bridge(outer_transformation::AbstractUnivariateTransformation, + inner_transformation::AbstractUnivariateTransformation) + Bridge(outer_transformation, inner_transformation) +end + +function (b::Bridge)(x) + (; outer_transformation, inner_transformation) = b + d = PM1() + transform_from(d, outer_transformation, transform_to(d, inner_transformation, x)) +end + +SpectralKit.domain(b::Bridge) = domain(b.inner_transformation) + +InverseFunctions.inverse(b::Bridge) = Bridge(b.inner_transformation, b.outer_transformation) + end diff --git a/test/Project.toml b/test/Project.toml index 3ed46b1..abd6b3b 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -3,6 +3,7 @@ Aqua = "4c88cf16-eb10-579e-8560-4a9242c79595" BenchmarkTools = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" DocStringExtensions = "ffbed154-4ef7-542d-bbb7-c09d3a79fcae" FiniteDifferences = "26cc04aa-876d-5657-8c51-4c34ba976000" +InverseFunctions = "3587e190-3f89-42d0-90ee-14403ec27112" JET = "c3a54625-cd67-489e-a8e7-0a5a0ff4e31b" LogExpFunctions = "2ab3a3ac-af41-5b50-aa03-7779005ae688" Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c" diff --git a/test/test_experimental.jl b/test/test_experimental.jl index a86da41..47b99b4 100644 --- a/test/test_experimental.jl +++ b/test/test_experimental.jl @@ -2,6 +2,7 @@ import SpectralKit.Experimental as SKX using LogExpFunctions: logistic using Test using SpectralKit +using InverseFunctions: inverse #### #### generic api @@ -148,3 +149,17 @@ approximation_grid = SKX.make_approximation_grid(model_family, model_parameters, @test @inferred SKX.sum_of_squared_residuals(model_family, model_parameters, policy_functions, approximation_grid) isa Float64 + +@testset "bridge" begin + i = InfRational(0, 1) + o = BoundedLinear(3, 5) + b = SKX.bridge(o, i) + ib = inverse(b) + @test domain(b) == domain(i) + @test domain(ib) == domain(o) + for _ in 1:100 + x = randn() * 10 + @test ib(b(x)) ≈ x + @test 3 ≤ b(x) ≤ 5 + end +end