Skip to content
Draft
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
Expand Up @@ -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"
Expand Down
27 changes: 25 additions & 2 deletions src/algorithms/optimization/fixed_point_differentiation.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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
2 changes: 1 addition & 1 deletion test/gradients/c4v_ctmrg_gradients.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion test/gradients/ctmrg_gradients.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading