From ee58c0bb5ba373ca1275004daa4f62767d2a0b45 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 18 Jun 2026 06:12:48 +0200 Subject: [PATCH 1/3] update Riemannian Conjugate Gradient tunable --- Project.toml | 2 +- src/solvers/rcg.jl | 236 ++++++++++++++++++++++------------ src/solvers/solve_dispatch.jl | 12 +- 3 files changed, 165 insertions(+), 85 deletions(-) diff --git a/Project.toml b/Project.toml index cd61afc..a3bb5e9 100644 --- a/Project.toml +++ b/Project.toml @@ -15,7 +15,7 @@ TensorOperations = "6aa20fa7-93e2-5fca-9bc0-fbd0db3c71a2" [compat] JuliaFormatter = "2.8.5" -Manifolds = "0.11.20" +Manifolds = "0.11.28" ManifoldsBase = "2.3.5" Manopt = "0.5.37" ProgressMeter = "1.11.0" diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index 7204a10..8cea02c 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -1,93 +1,97 @@ # solvers/rcg.jl — Riemannian Conjugate Gradient export RCGSolver -struct SegreProjectionTransport <: ManifoldsBase.AbstractVectorTransportMethod end +# Vector transport selection +""" + _supports_vector_transport_to(M, p, vt, retraction_method) -function _uses_segre_projection_transport(M) - return _uses_segre_projection_transport_unwrapped(_unwrap_solver_manifold(M)) +Return `true` if `vt` can transport a zero tangent vector from `p` to the +corresponding retracted point and the result is accepted as a tangent vector. +This is a conservative compatibility probe for Manifolds.jl / ManifoldsBase +vector transports. +""" +function _supports_vector_transport_to(M, p, vt, retraction_method) + try + X = zero_vector(M, p) + q = retract(M, p, X, retraction_method) + Y = vector_transport_to(M, p, X, q, vt) + return isnothing(check_vector(M, q, Y)) + catch + return false + end end -_uses_segre_projection_transport_unwrapped(::Manifolds.Segre) = true -_uses_segre_projection_transport_unwrapped(M::ProductManifold) = - all(_uses_segre_projection_transport, M.manifolds) -_uses_segre_projection_transport_unwrapped(M) = - hasproperty(M, :native) ? _uses_segre_projection_transport(getproperty(M, :native)) : - false - -function ManifoldsBase.vector_transport_to( - M::Manifolds.Segre, - p, - X, - q, - ::SegreProjectionTransport, -) - xparts = point_parts(X) - qparts = point_parts(q) - length(xparts) == length(qparts) || throw( - DimensionMismatch( - "Segre tangent/point part count mismatch: $(length(xparts)) vs $(length(qparts)).", - ), - ) - T = promote_type(eltype(_unwrap_part(xparts[1])), eltype(_unwrap_part(qparts[1]))) - ν = T(_unwrap_part(xparts[1])[1]) - Udot = Vector{Vector{T}}(undef, length(xparts) - 1) - @inbounds for m in eachindex(Udot) - xm = Vector{T}(_unwrap_part(xparts[m+1])) - qm = _unwrap_part(qparts[m+1]) - length(xm) == length(qm) || - throw(DimensionMismatch("Segre mode $m transport length mismatch.")) - xm .-= dot(qm, xm) .* qm - Udot[m] = xm +""" + _default_vector_transport_method(M, p, retraction_method) + +Return the default vector transport method for the given manifold and point. +If the manifold and point layout support it, use `ProjectionTransport()`. +Otherwise, use the manifold's default vector transport method. +""" +function _default_vector_transport_method(M, p, retraction_method) + vt = ManifoldsBase.ProjectionTransport() + if _supports_vector_transport_to(M, p, vt, retraction_method) + return vt end - return pack_tangent_rank1_segre(ν, Udot) + + return ManifoldsBase.default_vector_transport_method(M, typeof(p)) end -function ManifoldsBase.vector_transport_to!( - M::Manifolds.Segre, - Y, - p, - X, - q, - m::SegreProjectionTransport, +# RCG coefficient and restart rule selection +function _rcg_coefficient_rule( + M, + coefficient::Symbol, + transport; + denom_threshold::Real = 1e-10, + beale_restart::Bool = false, + restart_threshold::Real = 0.2, ) - Ynew = vector_transport_to(M, p, X, q, m) - yparts = point_parts(Y) - newparts = point_parts(Ynew) - length(yparts) == length(newparts) || throw( - DimensionMismatch( - "Segre transport destination part count mismatch: $(length(yparts)) vs $(length(newparts)).", - ), - ) - @inbounds for k in eachindex(newparts) - yparts[k] = newparts[k] + rule = + coefficient in (:conjugate_descent, :cd) ? Manopt.ConjugateDescentCoefficient() : + coefficient in (:hager_zhang, :hz) ? Manopt.HagerZhangCoefficient( + M; + vector_transport_method = transport, + denom_threshold = denom_threshold, + ) : + coefficient in (:polak_ribiere, :pr, :prp) ? Manopt.PolakRibiereCoefficient( + M; + vector_transport_method = transport, + ) : + coefficient in (:fletcher_reeves, :fr) ? Manopt.FletcherReevesCoefficient() : + coefficient in (:dai_yuan, :dy) ? Manopt.DaiYuanCoefficient( + M; + vector_transport_method = transport, + ) : + coefficient in (:hestenes_stiefel, :hs) ? Manopt.HestenesStiefelCoefficient( + M; + vector_transport_method = transport, + ) : + coefficient in (:liu_storey, :ls) ? Manopt.LiuStoreyCoefficient( + M; + vector_transport_method = transport, + ) : + coefficient in (:steepest, :steepest_descent, :gd, :gradient_descent) ? + Manopt.SteepestDescentCoefficient() : + throw(ArgumentError("Unknown RCG coefficient=$(coefficient).")) + + if beale_restart + return Manopt.ConjugateGradientBealeRestart( + M, + rule; + threshold = restart_threshold, + vector_transport_method = transport, + ) end - return Y + + return rule end -""" - solve_rcg(model_cost, model_egrad, M, p0; maxiter, tol, verbose, return_stats, model_grad, vector_transport_method) - -Riemannian conjugate gradient. Uses a custom projection-style transport for -`Manifolds.Segre` / `ProductManifold(Manifolds.Segre(...), ...)`, and otherwise -prefers `ProjectionTransport()` when the current manifold/point layout -supports it, falling back to `SchildsLadderTransport()` as needed. Callers can -override this with `vector_transport_method=...` when they want an explicit -transport choice. -""" -function _default_vector_transport_method(M, p, retraction_method) - if _uses_segre_projection_transport(M) - return SegreProjectionTransport() - end - vt = ManifoldsBase.ProjectionTransport() - try - X = zero_vector(M, p) - q = retract(M, p, X, retraction_method) - Y = vector_transport_to(M, p, X, q, vt) - return isnothing(check_vector(M, q, Y)) ? vt : - ManifoldsBase.SchildsLadderTransport() - catch - return ManifoldsBase.SchildsLadderTransport() - end +function _rcg_restart_condition(restart::Symbol; κ::Real = 1e-4) + return restart in (:never, :none, :no_restart) ? Manopt.NeverRestart() : + restart in (:non_descent, :nondescent) ? Manopt.RestartOnNonDescent() : + restart in (:non_sufficient_descent, :sufficient_descent) ? + Manopt.RestartOnNonSufficientDescent(κ) : + throw(ArgumentError("Unknown RCG restart=$(restart).")) end function solve_rcg( @@ -107,6 +111,12 @@ function solve_rcg( iteration_callbacks = (), grad_tol = nothing, normalized_objective::Bool = true, + coefficient::Symbol = :hager_zhang, + restart::Symbol = :non_descent, + restart_threshold::Real = 0.2, + sufficient_descent_kappa::Real = 1e-4, + denom_threshold::Real = 1e-10, + beale_restart::Bool = false, ) setup = _prepare_manopt_solver_functions( model_cost, @@ -119,13 +129,29 @@ function solve_rcg( grad_tol, normalized_objective, ) + # Get the initial point and the tangent space type p0_local = setup.p0 T = setup.T + retraction_method = _solver_retraction_method(M, p0_local) + transport = isnothing(vector_transport_method) ? _default_vector_transport_method(M, p0_local, retraction_method) : vector_transport_method + coefficient_rule = _rcg_coefficient_rule( + M, + coefficient, + transport; + denom_threshold, + beale_restart, + restart_threshold, + ) + restart_rule = _rcg_restart_condition( + restart; + κ = sufficient_descent_kappa, + ) + tol_g = setup.dual_grad_tol dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) @@ -149,6 +175,8 @@ function solve_rcg( p0_local; retraction_method = retraction_method, vector_transport_method = transport, + coefficient = coefficient_rule, + restart_condition = restart_rule, stopping_criterion = stopping, debug = callbacks.debug_actions, count = [:Cost, :Gradient], @@ -171,20 +199,56 @@ function solve_rcg( return_stats, verbose, normalized_objective = setup.uses_relative_objective, + solver_info_extra = ( + rcg_coefficient = coefficient, + rcg_restart = restart, + rcg_beale_restart = beale_restart, + rcg_restart_threshold = Float64(restart_threshold), + rcg_sufficient_descent_kappa = Float64(sufficient_descent_kappa), + rcg_denom_threshold = Float64(denom_threshold), + rcg_transport = string(typeof(transport)), + rcg_coefficient_rule = string(typeof(coefficient_rule)), + rcg_restart_rule = string(typeof(restart_rule)), + ), ) end -# ========== RCGSolver (AbstractFirstOrderROSolver) ========== - +# RCGSolver object """ - RCGSolver + RCGSolver(; coefficient=:hager_zhang, restart=:non_descent, ...) + + Riemannian conjugate gradient solver. -Riemannian conjugate gradient. Call via -`solve(RCGSolver(), model; init=:random, gradient_mode=:riemannian, vector_transport_method=nothing)`. +Useful options: + +- `coefficient = :hager_zhang` +- `coefficient = :polak_ribiere` +- `coefficient = :fletcher_reeves` +- `coefficient = :dai_yuan` +- `coefficient = :hestenes_stiefel` +- `coefficient = :conjugate_descent` +- `coefficient = :steepest` + +Restart options: + +- `restart = :non_descent` +- `restart = :non_sufficient_descent` +- `restart = :never` + +The default is chosen for CPD swamp experiments: +RCGSolver(; coefficient=:hager_zhang, restart=:non_descent) """ -struct RCGSolver <: AbstractFirstOrderROSolver end +Base.@kwdef struct RCGSolver <: AbstractFirstOrderROSolver + coefficient::Symbol = :hager_zhang + restart::Symbol = :non_descent + restart_threshold::Float64 = 0.2 + sufficient_descent_kappa::Float64 = 1e-4 + denom_threshold::Float64 = 1e-10 + beale_restart::Bool = false +end solver_symbol(::RCGSolver) = :rcg + first_order_diagnostics_recorder(::RCGSolver) = _SolverDiagnosticsRecorder(line_search_enabled = true) @@ -219,5 +283,11 @@ function run_first_order_solver( iteration_callbacks, grad_tol, normalized_objective, + coefficient = solver.coefficient, + restart = solver.restart, + restart_threshold = solver.restart_threshold, + sufficient_descent_kappa = solver.sufficient_descent_kappa, + denom_threshold = solver.denom_threshold, + beale_restart = solver.beale_restart, ) end diff --git a/src/solvers/solve_dispatch.jl b/src/solvers/solve_dispatch.jl index 0f132bd..9bdc801 100644 --- a/src/solvers/solve_dispatch.jl +++ b/src/solvers/solve_dispatch.jl @@ -17,7 +17,17 @@ _solver_object(solver::AbstractSolver, ::Real; kwargs...) = solver _solver_object(::Val{:als}, ::Real; kwargs...) = ALSSolver() _solver_object(::Val{:rgd}, stepsize::Real; kwargs...) = RGDSolver(stepsize) _solver_object(::Val{:rgd_fixed}, stepsize::Real; kwargs...) = RGDFixedSolver(stepsize) -_solver_object(::Val{:rcg}, ::Real; kwargs...) = RCGSolver() + +function _solver_object(::Val{:rcg}, ::Real; kwargs...) + return RCGSolver(; + coefficient = get(kwargs, :coefficient, :hager_zhang), + restart = get(kwargs, :restart, :non_descent), + restart_threshold = Float64(get(kwargs, :restart_threshold, 0.2)), + sufficient_descent_kappa = Float64(get(kwargs, :sufficient_descent_kappa, 1e-4)), + denom_threshold = Float64(get(kwargs, :denom_threshold, 1e-10)), + beale_restart = Bool(get(kwargs, :beale_restart, false)), + ) +end function _solver_object(::Val{:lbfgs}, ::Real; kwargs...) return LBFGSSolver(; From ea4c9d2270683e27456490ff1d3b6db17192c831 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 18 Jun 2026 06:25:15 +0200 Subject: [PATCH 2/3] Run JuliaFormatter --- src/core/tensor_ops.jl | 4 ++-- src/solvers/rcg.jl | 36 +++++++++++++----------------------- 2 files changed, 15 insertions(+), 25 deletions(-) diff --git a/src/core/tensor_ops.jl b/src/core/tensor_ops.jl index 5ab2f89..e5c1427 100644 --- a/src/core/tensor_ops.jl +++ b/src/core/tensor_ops.jl @@ -470,8 +470,8 @@ gradU_column_cp( cp_reconstruction_norm2(components::Vector{RankOneTensor{T}}) where {T<:AbstractFloat} = sum( - cross_component(components[i], components[j]) for i in eachindex(components), - j in eachindex(components) + cross_component(components[i], components[j]) for + i in eachindex(components), j in eachindex(components) ) function cp_inner_AX( diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index 8cea02c..f58f039 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -48,30 +48,23 @@ function _rcg_coefficient_rule( ) rule = coefficient in (:conjugate_descent, :cd) ? Manopt.ConjugateDescentCoefficient() : - coefficient in (:hager_zhang, :hz) ? Manopt.HagerZhangCoefficient( + coefficient in (:hager_zhang, :hz) ? + Manopt.HagerZhangCoefficient( M; vector_transport_method = transport, denom_threshold = denom_threshold, ) : - coefficient in (:polak_ribiere, :pr, :prp) ? Manopt.PolakRibiereCoefficient( - M; - vector_transport_method = transport, - ) : + coefficient in (:polak_ribiere, :pr, :prp) ? + Manopt.PolakRibiereCoefficient(M; vector_transport_method = transport) : coefficient in (:fletcher_reeves, :fr) ? Manopt.FletcherReevesCoefficient() : - coefficient in (:dai_yuan, :dy) ? Manopt.DaiYuanCoefficient( - M; - vector_transport_method = transport, - ) : - coefficient in (:hestenes_stiefel, :hs) ? Manopt.HestenesStiefelCoefficient( - M; - vector_transport_method = transport, - ) : - coefficient in (:liu_storey, :ls) ? Manopt.LiuStoreyCoefficient( - M; - vector_transport_method = transport, - ) : + coefficient in (:dai_yuan, :dy) ? + Manopt.DaiYuanCoefficient(M; vector_transport_method = transport) : + coefficient in (:hestenes_stiefel, :hs) ? + Manopt.HestenesStiefelCoefficient(M; vector_transport_method = transport) : + coefficient in (:liu_storey, :ls) ? + Manopt.LiuStoreyCoefficient(M; vector_transport_method = transport) : coefficient in (:steepest, :steepest_descent, :gd, :gradient_descent) ? - Manopt.SteepestDescentCoefficient() : + Manopt.SteepestDescentCoefficient() : throw(ArgumentError("Unknown RCG coefficient=$(coefficient).")) if beale_restart @@ -90,7 +83,7 @@ function _rcg_restart_condition(restart::Symbol; κ::Real = 1e-4) return restart in (:never, :none, :no_restart) ? Manopt.NeverRestart() : restart in (:non_descent, :nondescent) ? Manopt.RestartOnNonDescent() : restart in (:non_sufficient_descent, :sufficient_descent) ? - Manopt.RestartOnNonSufficientDescent(κ) : + Manopt.RestartOnNonSufficientDescent(κ) : throw(ArgumentError("Unknown RCG restart=$(restart).")) end @@ -147,10 +140,7 @@ function solve_rcg( beale_restart, restart_threshold, ) - restart_rule = _rcg_restart_condition( - restart; - κ = sufficient_descent_kappa, - ) + restart_rule = _rcg_restart_condition(restart; κ = sufficient_descent_kappa) tol_g = setup.dual_grad_tol dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) From 7baa4f9aaf4855af9cc712d5b54233c63e202282 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 18 Jun 2026 06:28:44 +0200 Subject: [PATCH 3/3] Align format CI with project JuliaFormatter version --- .github/workflows/format_check.yml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/workflows/format_check.yml b/.github/workflows/format_check.yml index 0af4d85..2b61af9 100644 --- a/.github/workflows/format_check.yml +++ b/.github/workflows/format_check.yml @@ -14,12 +14,12 @@ jobs: steps: - uses: julia-actions/setup-julia@latest with: - version: "^1.4" + version: "1.10" - uses: actions/checkout@v6 - name: Install JuliaFormatter and format run: | - julia -e 'using Pkg; Pkg.add(PackageSpec(name="JuliaFormatter", version="1.0.33"))' - julia -e 'using JuliaFormatter; format(["./src", "./test"], verbose=true)' + julia --project=. -e 'using Pkg; Pkg.instantiate()' + julia --project=. -e 'using JuliaFormatter; format(["./src", "./test"], verbose=true)' - name: Format check run: | julia -e '