From 39ff233ddd301d802d746650031c0cad1829ff32 Mon Sep 17 00:00:00 2001 From: leburgel Date: Mon, 27 Jul 2026 10:55:46 +0200 Subject: [PATCH] Add `fixedpoint_gradient` implementation for `OptimKit.FixedPointAlgorithm` solvers --- Project.toml | 2 +- .../fixed_point_differentiation.jl | 27 +++++++++++++++++-- test/gradients/c4v_ctmrg_gradients.jl | 2 +- test/gradients/ctmrg_gradients.jl | 2 +- 4 files changed, 28 insertions(+), 5 deletions(-) diff --git a/Project.toml b/Project.toml index 68ece85cd..b315d5659 100644 --- a/Project.toml +++ b/Project.toml @@ -42,7 +42,7 @@ MPSKit = "0.13.9" MPSKitModels = "0.4" MatrixAlgebraKit = "0.6.5" OhMyThreads = "0.7, 0.8" -OptimKit = "0.4" +OptimKit = "0.5" Printf = "1" Random = "1" Statistics = "1" diff --git a/src/algorithms/optimization/fixed_point_differentiation.jl b/src/algorithms/optimization/fixed_point_differentiation.jl index 2c2a54991..4b83124b5 100644 --- a/src/algorithms/optimization/fixed_point_differentiation.jl +++ b/src/algorithms/optimization/fixed_point_differentiation.jl @@ -93,8 +93,12 @@ end FixedPointGradient(; kwargs...) = GradientAlgorithm(; alg = :FixedPointGradient, kwargs...) GRADIENT_ALGORITHM_SYMBOLS[:FixedPointGradient] = FixedPointGradient -const FIXEDPOINT_SOLVER_SYMBOLS = IdDict{Symbol, Type{<:Any}}( - :GMRES => GMRES, :BiCGStab => BiCGStab, :Arnoldi => Arnoldi, +const FIXEDPOINT_SOLVER_SYMBOLS = IdDict{Symbol, Any}( + :GMRES => GMRES, + :BiCGStab => BiCGStab, + :Arnoldi => Arnoldi, + :AndersonMixing => AndersonMixing, + :SimpleIteration => SimpleIteration, ) _default_solver_alg(::Type{<:FixedPointGradient}) = Defaults.gradient_fixedpoint_solver_alg @@ -107,6 +111,15 @@ function _pad_solver_kwargs(::Type{<:Arnoldi}, solver_kwargs) ) return solver_kwargs end +function _pad_solver_kwargs(::Type{<:OptimKit.FixedPointAlgorithm}, solver_kwargs) + solver_kwargs = (; + gradtol = solver_kwargs.tol, # patch for OptimKit.FixedPointAlgorithm gradient tolerance kwarg name + solver_kwargs..., + ) + solver_kwargs = Base.structdiff(solver_kwargs, (; tol = nothing)) + + return solver_kwargs +end """ $(TYPEDEF) @@ -360,3 +373,13 @@ function fixedpoint_gradient(∂E∂x, ∂f∂x, ∂f∂A, x₀, alg::KrylovKit. return ∂f∂A(y) end + +function fixedpoint_gradient(x̆, ∂ₓf, ∂ₚf, y₀, alg::OptimKit.FixedPointAlgorithm) + fp(y) = x̆ + ∂ₓf(y) # fixed-point condition for geometric sum + y, g, = OptimKit.fixedpoint(fp, y₀, alg) + if alg.verbosity > 0 && norm(g) > alg.gradtol + @warn("gradient fixed-point iteration reached maximal number of iterations without converging: ‖g‖ = $(norm(g))") + end + + return ∂ₚf(y) +end diff --git a/test/gradients/c4v_ctmrg_gradients.jl b/test/gradients/c4v_ctmrg_gradients.jl index 2a0c4dab7..1d860adb5 100644 --- a/test/gradients/c4v_ctmrg_gradients.jl +++ b/test/gradients/c4v_ctmrg_gradients.jl @@ -25,7 +25,7 @@ ctmrg_algs = [[:C4vCTMRG]] projector_algs = [[:C4vEighProjector, :C4vQRProjector]] decomposition_rrule_algs = [[:FullPullback, :TruncPullback]] gradient_algs = [[nothing, :FixedPointGradient]] -gradient_solver_algs = [[:GeomSum, :ManualIter, :GMRES, :BiCGStab, :Arnoldi]] +gradient_solver_algs = [[:GeomSum, :ManualIter, :GMRES, :BiCGStab, :Arnoldi, :AndersonMixing]] steps = -0.01:0.005:0.01 # record which rrule alg is compatible with which projector alg diff --git a/test/gradients/ctmrg_gradients.jl b/test/gradients/ctmrg_gradients.jl index 8618e139a..b1bd1974b 100644 --- a/test/gradients/ctmrg_gradients.jl +++ b/test/gradients/ctmrg_gradients.jl @@ -23,7 +23,7 @@ projector_algs = [[:HalfInfiniteProjector, :FullInfiniteProjector], [:HalfInfini svd_rrule_algs = [[:FullPullback, :TruncPullback, :Arnoldi], [:FullPullback, :Arnoldi]] gradient_algs = [[nothing, :FixedPointGradient], [:FixedPointGradient]] gradient_solver_algs = [ - [:GeomSum, :ManualIter, :GMRES, :BiCGStab, :Arnoldi], + [:GeomSum, :ManualIter, :GMRES, :BiCGStab, :Arnoldi, :AndersonMixing], [:GeomSum, :ManualIter, :GMRES, :BiCGStab, :Arnoldi], ] steps = -0.01:0.005:0.01