From 36622d0722abf7f04f4ed25e89793ae21d6b766a Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 17 Jun 2026 15:02:27 +0000 Subject: [PATCH 01/37] fix: apply JuliaFormatter output for cpd and rgd files --- src/api/cpd.jl | 3 ++- src/solvers/rgd.jl | 3 ++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/src/api/cpd.jl b/src/api/cpd.jl index e2a8de6..2d19de5 100644 --- a/src/api/cpd.jl +++ b/src/api/cpd.jl @@ -862,7 +862,8 @@ function _run_cpd_solver( refinement_verbose = verbose, vector_transport_method, grad_tol = _cpd_manifold_grad_tol(model, solver, tol), - normalized_objective = solver isa Union{RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver}, + normalized_objective = solver isa + Union{RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver}, iteration_callbacks, kwargs..., ) diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index 93b1d73..af180e9 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -787,7 +787,8 @@ function solve_rgd_fixed( tiny_grad_tol = isnothing(grad_tol) ? T(1e-5) : (uses_relative_objective ? T(grad_tol) * objective_scale : T(grad_tol)) - stopping = StopWhenAny(StopAfterIteration(maxiter), StopWhenGradientNormLess(grad_stop_tol)) + stopping = + StopWhenAny(StopAfterIteration(maxiter), StopWhenGradientNormLess(grad_stop_tol)) progress = maxiter > 0 ? make_rgd_fixed_progress(maxiter; enabled = verbose, phase = :refinement, dt = 0.2) : From 3e17731a685f2b3362521744fa2aec7d6650d9e9 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 20:36:40 +0200 Subject: [PATCH 02/37] make progress meter render correctly both phases --- src/core/progress.jl | 59 +++++++++++++++++++++++++++++++++++--------- 1 file changed, 48 insertions(+), 11 deletions(-) diff --git a/src/core/progress.jl b/src/core/progress.jl index 249deb4..30c514a 100644 --- a/src/core/progress.jl +++ b/src/core/progress.jl @@ -177,6 +177,32 @@ end return nothing end +@inline function _force_visible_phase_finish(tracker::PhaseProgress, progress) + return tracker.phase == :refinement && + _was_rendered(tracker.initialization) && + !_was_rendered(progress) +end + +function _render_unrendered_completion!(meter, showvalues) + # ProgressMeter does not render a meter that reaches 100% before its first + # visible update. Give it one display-only step before completion. + PM.update!( + meter, + meter.n; + showvalues, + force = true, + max_steps = meter.n + 1, + ) + return nothing +end + +function _force_visible_unrendered_progress!(progress, meter, showvalues) + _was_rendered(progress) && return nothing + _render_unrendered_completion!(meter, showvalues) + _mark_rendered!(progress) + return nothing +end + update_progress!(::NoMethodProgress, args...; kwargs...) = nothing function update_progress!( @@ -192,14 +218,17 @@ function update_progress!( set_phase!(tracker, progress.phase) end t = time() - if force || current >= meter.n || t > meter.tlast + meter.dt + renders_by_time = t > meter.tlast + meter.dt + if force || current >= meter.n || renders_by_time showvalues_with_method = if isnothing(showvalues) Any[("Method", _method_name(progress))] else Any[("Method", _method_name(progress)); showvalues] end - PM.update!(meter, current; showvalues = showvalues_with_method) - _mark_rendered!(progress) + PM.update!(meter, current; showvalues = showvalues_with_method, force) + if force || current < meter.n || _was_rendered(progress) + _mark_rendered!(progress) + end end return nothing end @@ -222,17 +251,14 @@ function finish_progress!( Any[("Method", _method_name(progress)); showvalues] end - force_refinement_finish = - tracker.phase == :refinement && - _was_rendered(tracker.initialization) && - !_was_rendered(progress) - - if force_refinement_finish - PM.update!(meter, meter.n; showvalues = showvalues_with_method) - _mark_rendered!(progress) + if _force_visible_phase_finish(tracker, progress) + _force_visible_unrendered_progress!(progress, meter, showvalues_with_method) + PM.finish!(meter; showvalues = showvalues_with_method) return nothing end + _was_rendered(progress) || return nothing + PM.finish!(meter; showvalues = showvalues_with_method) _mark_rendered!(progress) return nothing @@ -251,12 +277,23 @@ function finish_progress!( if active_progress(tracker) === progress return finish_progress!(tracker; current, showvalues) end + if _force_visible_phase_finish(tracker, progress) + showvalues_with_method = if isnothing(showvalues) + Any[("Method", _method_name(progress))] + else + Any[("Method", _method_name(progress)); showvalues] + end + _force_visible_unrendered_progress!(progress, meter, showvalues_with_method) + PM.finish!(meter; showvalues = showvalues_with_method) + return nothing + end end showvalues_with_method = if isnothing(showvalues) Any[("Method", _method_name(progress))] else Any[("Method", _method_name(progress)); showvalues] end + _was_rendered(progress) || return nothing PM.finish!(meter; showvalues = showvalues_with_method) _mark_rendered!(progress) return nothing From a1cffff4f08ceda6e1059f273ab77c3af70d0d1c Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 20:39:44 +0200 Subject: [PATCH 03/37] export manopt helpers to a separate file --- src/backend.jl | 1 + src/solvers/manopt_helpers.jl | 594 ++++++++++++++++++++++++++++++++++ src/solvers/rgd.jl | 594 ---------------------------------- 3 files changed, 595 insertions(+), 594 deletions(-) create mode 100644 src/solvers/manopt_helpers.jl diff --git a/src/backend.jl b/src/backend.jl index ab39e34..5a96538 100644 --- a/src/backend.jl +++ b/src/backend.jl @@ -3,6 +3,7 @@ include("solvers/nncp_updates.jl") include("solvers/cp_als.jl") include("solvers/btd_als.jl") include("solvers/rals.jl") +include("solvers/manopt_helpers.jl") include("solvers/rgd.jl") include("solvers/btd_tsd.jl") include("solvers/rcg.jl") diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl new file mode 100644 index 0000000..a107e27 --- /dev/null +++ b/src/solvers/manopt_helpers.jl @@ -0,0 +1,594 @@ + +struct _SolverDebugSink <: IO end +Base.isopen(::_SolverDebugSink) = true +Base.write(::_SolverDebugSink, ::UInt8) = 1 +Base.write(::_SolverDebugSink, s::Union{String,SubString{String}}) = sizeof(s) +Base.unsafe_write(::_SolverDebugSink, ::Ptr{UInt8}, n::UInt) = Int(n) + +const _SOLVER_DEBUG_SINK = _SolverDebugSink() + +mutable struct StopWhenCostRelChangeAndGradientLess{T<:Real} <: Manopt.StoppingCriterion + tol_cost::T + tol_grad::T + prev_cost::T + last_cost_rel_change::T + last_grad_norm::T + at_iteration::Int +end + +function StopWhenCostRelChangeAndGradientLess(tol_cost::T, tol_grad::T) where {T<:Real} + return StopWhenCostRelChangeAndGradientLess{T}( + tol_cost, + tol_grad, + T(Inf), + T(Inf), + T(Inf), + -1, + ) +end + +function (c::StopWhenCostRelChangeAndGradientLess)(problem, state, i) + if i == 0 + c.prev_cost = Manopt.get_cost(problem, Manopt.get_iterate(state)) + c.last_cost_rel_change = oftype(c.tol_cost, Inf) + c.last_grad_norm = oftype(c.tol_grad, Inf) + c.at_iteration = -1 + return false + end + M = Manopt.get_manifold(problem) + p = Manopt.get_iterate(state) + cost_val = Manopt.get_cost(problem, p) + grad_val = Manopt.get_gradient(problem, p) + grad_norm = norm(M, p, grad_val) + rel_change = abs(c.prev_cost - cost_val) / max(abs(c.prev_cost), one(cost_val)) + c.prev_cost = cost_val + c.last_cost_rel_change = rel_change + c.last_grad_norm = grad_norm + if rel_change < c.tol_cost && grad_norm < c.tol_grad + c.at_iteration = i + return true + end + return false +end + +function Manopt.get_reason(c::StopWhenCostRelChangeAndGradientLess) + if c.at_iteration >= 0 + return "At iteration $(c.at_iteration) the relative cost change ($(c.last_cost_rel_change)) " * + "is below $(c.tol_cost) and the gradient norm ($(c.last_grad_norm)) " * + "is below $(c.tol_grad).\n" + end + return "" +end + +function Manopt.status_summary(c::StopWhenCostRelChangeAndGradientLess) + has_stopped = c.at_iteration >= 0 + status = has_stopped ? "reached" : "not reached" + return "cost rel change < $(c.tol_cost) and |grad f| < $(c.tol_grad): $status" +end + +Manopt.indicates_convergence(::StopWhenCostRelChangeAndGradientLess) = true + +function Base.show(io::IO, c::StopWhenCostRelChangeAndGradientLess) + return print( + io, + "StopWhenCostRelChangeAndGradientLess($(c.tol_cost), $(c.tol_grad))\n $(Manopt.status_summary(c))", + ) +end + + +function _tk_get_solver_result(state) + try + return Manopt.get_solver_result(state) + catch + end + while hasproperty(state, :state) + state = state.state + end + for key in (:p, :x, :point) + hasproperty(state, key) && return getproperty(state, key) + end + throw( + ArgumentError( + "Cannot extract point from state $(typeof(state)). Properties: $(propertynames(state))", + ), + ) +end + + +@inline _align_layout_like_point(p, x) = + hasproperty(p, :x) ? + (hasproperty(x, :x) ? x : (x isa Tuple ? ArrayPartition(x...) : x)) : + (hasproperty(x, :x) ? Tuple(getproperty(x, :x)) : x) + + +function _to_array_partition(x) + if x isa ArrayPartition + return ArrayPartition(map(_to_array_partition, x.x)...) + elseif hasproperty(x, :x) + return ArrayPartition(map(_to_array_partition, getproperty(x, :x))...) + elseif x isa Tuple + return ArrayPartition(map(_to_array_partition, x)...) + end + return x +end + + +function _solver_point(M, p0) + M2 = _unwrap_solver_manifold(M) + return M2 isa ProductManifold ? _to_array_partition(p0) : p0 +end + + +function _contains_sqeuclidean_manifold(M) + M2 = _unwrap_solver_manifold(M) + if M2 isa SqEuclidean || M2 isa SoftplusEuclidean + return true + elseif M2 isa ProductManifold + return any(_contains_sqeuclidean_manifold, M2.manifolds) + elseif hasproperty(M2, :native) && ( + getproperty(M2, :native) isa SqEuclidean || + getproperty(M2, :native) isa SoftplusEuclidean + ) + return true + end + return false +end + +function _contains_strict_sqeuclidean_manifold(M) + M2 = _unwrap_solver_manifold(M) + if M2 isa SqEuclidean + return true + elseif M2 isa ProductManifold + return any(_contains_strict_sqeuclidean_manifold, M2.manifolds) + elseif hasproperty(M2, :native) && (getproperty(M2, :native) isa SqEuclidean) + return true + end + return false +end + + +function _armijo_max_decreases(initial_stepsize::Real, contraction::Real, alpha_min::Real) + initial_stepsize <= alpha_min && return 0 + (contraction <= 0 || contraction >= 1) && return 1000 + n = floor(Int, log(alpha_min / initial_stepsize) / log(contraction)) + return max(n, 0) +end + + +function _adaptive_initial_stepsize( + M, + p0, + model_grad, + retraction_method, + base_stepsize::T; + alpha_min::T, + scale_c::T = one(T), + clamp_low_factor::T = T(0.1), + clamp_high_factor::T = T(10), + delta_scale::T = T(1e-3), +) where {T<:AbstractFloat} + g0 = model_grad(M, p0) + d = -copy(g0) + dnorm = norm(M, p0, d) + (!isfinite(dnorm) || dnorm <= sqrt(eps(T))) && return base_stepsize + δ = delta_scale / max(dnorm, one(T)) + q = try + retract(M, p0, δ .* d, retraction_method) + catch + return base_stepsize + end + _all_finite(q) || return base_stepsize + gq = model_grad(M, q) + _all_finite(gq) || return base_stepsize + κ_num = inner(M, p0, gq .- g0, d) + κ_den = δ * dnorm^2 + (!isfinite(κ_num) || !isfinite(κ_den) || κ_den <= eps(T)) && return base_stepsize + κ = max(κ_num / κ_den, eps(T)) + α_raw = scale_c / κ + α_low = max(alpha_min, clamp_low_factor * base_stepsize) + α_high = clamp_high_factor * base_stepsize + α = clamp(α_raw, α_low, α_high) + if !isfinite(α) + return base_stepsize + end + return α +end + +function _all_finite(x) + if x isa Number + return isfinite(x) + elseif hasproperty(x, :x) + return all(_all_finite, getproperty(x, :x)) + elseif x isa AbstractArray + return all(isfinite, x) + elseif x isa Tuple + return all(_all_finite, x) + elseif x isa Manifolds.TuckerPoint + return _all_finite(x.hosvd.core) && all(_all_finite, x.hosvd.U) + elseif x isa Manifolds.TuckerTangentVector + return _all_finite(getproperty(x, :Ċ)) && all(_all_finite, getproperty(x, :U̇)) + end + try + return all(isfinite, x) + catch + return false + end +end + + +function _safe_cost_function(model_cost) + return function (M, p) + _all_finite(p) || return Inf + c = model_cost(M, p) + return isfinite(c) ? c : Inf + end +end + + +function _layout_adapt_gradient(model_grad) + return function (M, p) + g = model_grad(M, p) + return _align_layout_like_point(p, g) + end +end + +function _scalar_eltype(p) + if hasproperty(p, :x) || p isa AbstractVector || p isa Tuple + parts = point_parts(p) + isempty(parts) && return Float64 + return _scalar_eltype(first(parts)) + elseif p isa Manifolds.TuckerPoint + return eltype(p.hosvd.core) + elseif p isa Manifolds.TuckerTangentVector + return eltype(getproperty(p, :Ċ)) + elseif p isa Real + return typeof(p) + else + return eltype(p) + end +end + +@inline function _solver_gradient(state) + try + return Manopt.get_gradient(state) + catch + return nothing + end +end + +@inline _scale_solver_tangent(x::Number, scale::Real) = x * scale +_scale_solver_tangent(x::AbstractArray, scale::Real) = x .* scale +_scale_solver_tangent(x::ArrayPartition, scale::Real) = + ArrayPartition(map(part -> _scale_solver_tangent(part, scale), x.x)...) +_scale_solver_tangent(x::Tuple, scale::Real) = + map(part -> _scale_solver_tangent(part, scale), x) + +function _scale_solver_tangent(x, scale::Real) + try + return x .* scale + catch + return scale * x + end +end + +function _relative_solver_functions(model_cost, model_grad, scale::Real) + scale > 0 || return model_cost, model_grad, false + scale == one(scale) && return model_cost, model_grad, false + inv_scale = inv(scale) + return ( + (M, p) -> model_cost(M, p) * inv_scale, + (M, p) -> _scale_solver_tangent(model_grad(M, p), inv_scale), + true, + ) +end + +@inline _solver_has_converged(state) = Manopt.has_converged(state) + + +function _solver_iterations(state, maxiter::Int) + if isdefined(Manopt, :stopped_at) + try + k = Manopt.stopped_at(state) + return k > 0 ? Int(k) : maxiter + catch + end + end + while hasproperty(state, :state) + state = state.state + end + return hasproperty(state, :stop) && hasproperty(state.stop, :at_iteration) ? + state.stop.at_iteration : maxiter +end + +@inline function _solver_iteration_source() + return isdefined(Manopt, :stopped_at) ? :stopped_at : :stop_at_iteration_fallback +end + +function _solver_stats( + model_cost, + model_grad, + M, + p_opt, + state, + ::Nothing; + tol_T, + maxiter::Int, + solver::Symbol, + tiny_grad_tol = nothing, + solver_info = (;), + use_state_gradient::Bool = true, +) + T = typeof(tol_T) + final_cost = model_cost(M, p_opt) + cost_for_error = max(T(0), T(2) * final_cost) + rel_error = sqrt(cost_for_error) + grad_state = use_state_gradient ? _solver_gradient(state) : nothing + grad_from_state = !isnothing(grad_state) + grad_final = + isnothing(grad_state) ? model_grad(M, p_opt) : + _align_layout_like_point(p_opt, grad_state) + grad_norm = norm(M, p_opt, grad_final) + iterations = _solver_iterations(state, maxiter) + converged_grad = + grad_norm < tol_T || (!isnothing(tiny_grad_tol) && grad_norm < tiny_grad_tol) + converged_state = _solver_has_converged(state) + solver_info = merge( + solver_info, + ( + gradient_source = grad_from_state ? :state : :recomputed, + has_converged_state = converged_state, + converged_by_gradient_threshold = converged_grad, + iteration_source = _solver_iteration_source(), + ), + ) + return ( + point = p_opt, + cost = final_cost, + rel_error = rel_error, + grad_norm = grad_norm, + iterations = iterations, + converged = converged_state, + solver = solver, + solver_info = solver_info, + ) +end + +function _solver_stats( + model_cost, + model_grad, + M, + p_opt, + state, + normA2::Real; + tol_T, + maxiter::Int, + solver::Symbol, + tiny_grad_tol = nothing, + solver_info = (;), + use_state_gradient::Bool = true, +) + T = typeof(tol_T) + final_cost = model_cost(M, p_opt) + cost_for_error = max(T(0), T(2) * final_cost) + rel_error = _relative_error_frob_sq(cost_for_error, T(normA2)) + grad_state = use_state_gradient ? _solver_gradient(state) : nothing + grad_from_state = !isnothing(grad_state) + grad_final = + isnothing(grad_state) ? model_grad(M, p_opt) : + _align_layout_like_point(p_opt, grad_state) + grad_norm = norm(M, p_opt, grad_final) + iterations = _solver_iterations(state, maxiter) + converged_grad = + grad_norm < tol_T || (!isnothing(tiny_grad_tol) && grad_norm < tiny_grad_tol) + converged_state = _solver_has_converged(state) + solver_info = merge( + solver_info, + ( + gradient_source = grad_from_state ? :state : :recomputed, + has_converged_state = converged_state, + converged_by_gradient_threshold = converged_grad, + iteration_source = _solver_iteration_source(), + ), + ) + return ( + point = p_opt, + cost = final_cost, + rel_error = rel_error, + grad_norm = grad_norm, + iterations = iterations, + converged = converged_state, + solver = solver, + solver_info = solver_info, + ) +end + +_solver_debug_callbacks(callbacks...) = Any[cb for cb in callbacks if !isnothing(cb)] + +# Allow callers to pass `nothing` (e.g., when verbose/debug is omitted) +_solver_debug_actions(::Nothing, callbacks...) = _solver_debug_callbacks(callbacks...) + +function _solver_debug_actions(verbose::Bool, callbacks...) + callback_actions = _solver_debug_callbacks(callbacks...) + if verbose + io = _SOLVER_DEBUG_SINK + init_group = Manopt.DebugGroup([ + Manopt.DebugDivider("Initial "; io, at_init = true), + Manopt.DebugCost(; io, format = "f(x): %.6e", at_init = true), + Manopt.DebugGradientNorm(; io, format = "|grad f(p)|:%.6e", at_init = true), + Manopt.DebugDivider("\n"; io, at_init = true), + ]) + iter_group = Manopt.DebugEvery( + Manopt.DebugGroup([ + Manopt.DebugIteration(; io, format = "# %-6d"), + Manopt.DebugDivider(" "; io, at_init = true), + Manopt.DebugCost(; io, format = "f(x): %.6e", at_init = true), + Manopt.DebugGradientNorm(; io, format = "|grad f(p)|:%.6e", at_init = true), + Manopt.DebugDivider("\n"; io, at_init = true), + ]), + 100, + ) + iteration_actions = Any[iter_group] + append!(iteration_actions, callback_actions) + return Any[:Start=>Any[init_group], :Iteration=>iteration_actions] + end + return callback_actions +end + +function _solver_progress_callback( + progress, + model_cost, + model_grad, + M; + normA2 = nothing, + diagnostics_recorder = nothing, +) + progress isa NoMethodProgress && return nothing + has_relative_scale = !isnothing(normA2) && normA2 > 0 + target_norm = has_relative_scale ? sqrt(normA2) : nothing + return function (problem, state, k) + k <= 0 && return nothing + p = get_iterate(state) + c = model_cost(M, p) + g = model_grad(M, p) + gnorm = norm(M, p, g) + c_display = has_relative_scale ? sqrt(max(2 * c, zero(c))) : c + gnorm_display = has_relative_scale ? gnorm * target_norm : gnorm + showvalues = Any[("Iter", k), ("Cost", c_display), ("Grad norm", gnorm_display)] + if !isnothing(diagnostics_recorder) + step = diagnostics_recorder.accepted_stepsize_history + trials = diagnostics_recorder.line_search_trial_history + !isempty(step) && push!(showvalues, ("Accepted α", step[end])) + !isempty(trials) && push!(showvalues, ("Line-search trials", trials[end])) + end + update_progress!(progress, k; showvalues) + return nothing + end +end + +mutable struct _SolverDiagnosticsRecorder + first_accepted_stepsize::Float64 + min_accepted_stepsize::Float64 + first_line_search_trials::Int + line_search_trial_count::Int + function_evaluations::Int + gradient_evaluations::Int + prev_function_evaluations::Int + prev_gradient_evaluations::Int + line_search_enabled::Bool + fallback_stepsize::Float64 + accepted_stepsize_history::Vector{Float64} + line_search_trial_history::Vector{Int} +end + +function _SolverDiagnosticsRecorder(; + line_search_enabled::Bool, + fallback_stepsize::Real = NaN, +) + return _SolverDiagnosticsRecorder( + NaN, + Inf, + 0, + 0, + 0, + 0, + 0, + 0, + line_search_enabled, + Float64(fallback_stepsize), + Float64[], + Int[], + ) +end + +function _solver_eval_count(problem, sym::Symbol) + try + count = get_count(get_objective(problem), sym) + return count < 0 ? 0 : Int(count) + catch + return 0 + end +end + +function _solver_diagnostics_callback(recorder::_SolverDiagnosticsRecorder) + return function (problem, state, k) + fe = _solver_eval_count(problem, :Cost) + ge = _solver_eval_count(problem, :Gradient) + if k == 0 + recorder.function_evaluations = fe + recorder.gradient_evaluations = ge + recorder.prev_function_evaluations = fe + recorder.prev_gradient_evaluations = ge + return nothing + end + step = + recorder.line_search_enabled ? get_last_stepsize(problem, state, k) : + recorder.fallback_stepsize + delta_fe = max(fe - recorder.prev_function_evaluations, 0) + step_f = Float64(step) + ls_trials = recorder.line_search_enabled ? max(delta_fe - 1, 0) : 0 + if isnan(recorder.first_accepted_stepsize) + recorder.first_accepted_stepsize = step_f + recorder.first_line_search_trials = ls_trials + end + recorder.min_accepted_stepsize = min(recorder.min_accepted_stepsize, step_f) + recorder.line_search_trial_count += ls_trials + push!(recorder.accepted_stepsize_history, step_f) + push!(recorder.line_search_trial_history, ls_trials) + recorder.function_evaluations = fe + recorder.gradient_evaluations = ge + recorder.prev_function_evaluations = fe + recorder.prev_gradient_evaluations = ge + return nothing + end +end + +function _solver_info(recorder::_SolverDiagnosticsRecorder, iterations::Int) + return ( + total_iterations = iterations, + first_accepted_stepsize = recorder.first_accepted_stepsize, + min_accepted_stepsize = recorder.min_accepted_stepsize, + first_line_search_trials = recorder.first_line_search_trials, + line_search_trial_count = recorder.line_search_trial_count, + function_evaluations = recorder.function_evaluations, + gradient_evaluations = recorder.gradient_evaluations, + accepted_stepsize_history = recorder.accepted_stepsize_history, + line_search_trial_history = recorder.line_search_trial_history, + ) +end + +function _solver_post_step_callback( + model::AbstractDecompositionModel, + M, + normalization::AbstractNormalizationPolicy, + solver_sym::Symbol, +) + normalization isa NoNormalization && return nothing + return function (problem, state, k) + p_old = get_iterate(state) + # Backend postprocessing (for example normalization) runs in canonical + # CP coordinates and then packs back into the solver's current layout. + p_new = post_step!( + model, + p_old; + normalization, + solver = solver_sym, + problem, + state, + iteration = k, + ) + p_new = _align_layout_like_point(p_old, p_new) + p_new === p_old && return nothing + set_iterate!(state, M, p_new) + if solver_sym == :rcg && hasproperty(state, :X) + get_gradient!(problem, state.X, get_iterate(state)) + if hasproperty(state, :δ) + state.δ = -copy(M, get_iterate(state), state.X) + end + hasproperty(state, :β) && (state.β = zero(typeof(state.β))) + if hasproperty(state, :coefficient) && hasproperty(state.coefficient, :storage) + update_storage!(state.coefficient.storage, problem, state) + end + end + return nothing + end +end \ No newline at end of file diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index af180e9..be850a6 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -2,600 +2,6 @@ export RGDSolver, RGDFixedSolver using Manopt -struct _SolverDebugSink <: IO end -Base.isopen(::_SolverDebugSink) = true -Base.write(::_SolverDebugSink, ::UInt8) = 1 -Base.write(::_SolverDebugSink, s::Union{String,SubString{String}}) = sizeof(s) -Base.unsafe_write(::_SolverDebugSink, ::Ptr{UInt8}, n::UInt) = Int(n) - -const _SOLVER_DEBUG_SINK = _SolverDebugSink() - -mutable struct StopWhenCostRelChangeAndGradientLess{T<:Real} <: Manopt.StoppingCriterion - tol_cost::T - tol_grad::T - prev_cost::T - last_cost_rel_change::T - last_grad_norm::T - at_iteration::Int -end - -function StopWhenCostRelChangeAndGradientLess(tol_cost::T, tol_grad::T) where {T<:Real} - return StopWhenCostRelChangeAndGradientLess{T}( - tol_cost, - tol_grad, - T(Inf), - T(Inf), - T(Inf), - -1, - ) -end - -function (c::StopWhenCostRelChangeAndGradientLess)(problem, state, i) - if i == 0 - c.prev_cost = Manopt.get_cost(problem, Manopt.get_iterate(state)) - c.last_cost_rel_change = oftype(c.tol_cost, Inf) - c.last_grad_norm = oftype(c.tol_grad, Inf) - c.at_iteration = -1 - return false - end - M = Manopt.get_manifold(problem) - p = Manopt.get_iterate(state) - cost_val = Manopt.get_cost(problem, p) - grad_val = Manopt.get_gradient(problem, p) - grad_norm = norm(M, p, grad_val) - rel_change = abs(c.prev_cost - cost_val) / max(abs(c.prev_cost), one(cost_val)) - c.prev_cost = cost_val - c.last_cost_rel_change = rel_change - c.last_grad_norm = grad_norm - if rel_change < c.tol_cost && grad_norm < c.tol_grad - c.at_iteration = i - return true - end - return false -end - -function Manopt.get_reason(c::StopWhenCostRelChangeAndGradientLess) - if c.at_iteration >= 0 - return "At iteration $(c.at_iteration) the relative cost change ($(c.last_cost_rel_change)) " * - "is below $(c.tol_cost) and the gradient norm ($(c.last_grad_norm)) " * - "is below $(c.tol_grad).\n" - end - return "" -end - -function Manopt.status_summary(c::StopWhenCostRelChangeAndGradientLess) - has_stopped = c.at_iteration >= 0 - status = has_stopped ? "reached" : "not reached" - return "cost rel change < $(c.tol_cost) and |grad f| < $(c.tol_grad): $status" -end - -Manopt.indicates_convergence(::StopWhenCostRelChangeAndGradientLess) = true - -function Base.show(io::IO, c::StopWhenCostRelChangeAndGradientLess) - return print( - io, - "StopWhenCostRelChangeAndGradientLess($(c.tol_cost), $(c.tol_grad))\n $(Manopt.status_summary(c))", - ) -end - - -function _tk_get_solver_result(state) - try - return Manopt.get_solver_result(state) - catch - end - while hasproperty(state, :state) - state = state.state - end - for key in (:p, :x, :point) - hasproperty(state, key) && return getproperty(state, key) - end - throw( - ArgumentError( - "Cannot extract point from state $(typeof(state)). Properties: $(propertynames(state))", - ), - ) -end - - -@inline _align_layout_like_point(p, x) = - hasproperty(p, :x) ? - (hasproperty(x, :x) ? x : (x isa Tuple ? ArrayPartition(x...) : x)) : - (hasproperty(x, :x) ? Tuple(getproperty(x, :x)) : x) - - -function _to_array_partition(x) - if x isa ArrayPartition - return ArrayPartition(map(_to_array_partition, x.x)...) - elseif hasproperty(x, :x) - return ArrayPartition(map(_to_array_partition, getproperty(x, :x))...) - elseif x isa Tuple - return ArrayPartition(map(_to_array_partition, x)...) - end - return x -end - - -function _solver_point(M, p0) - M2 = _unwrap_solver_manifold(M) - return M2 isa ProductManifold ? _to_array_partition(p0) : p0 -end - - -function _contains_sqeuclidean_manifold(M) - M2 = _unwrap_solver_manifold(M) - if M2 isa SqEuclidean || M2 isa SoftplusEuclidean - return true - elseif M2 isa ProductManifold - return any(_contains_sqeuclidean_manifold, M2.manifolds) - elseif hasproperty(M2, :native) && ( - getproperty(M2, :native) isa SqEuclidean || - getproperty(M2, :native) isa SoftplusEuclidean - ) - return true - end - return false -end - -function _contains_strict_sqeuclidean_manifold(M) - M2 = _unwrap_solver_manifold(M) - if M2 isa SqEuclidean - return true - elseif M2 isa ProductManifold - return any(_contains_strict_sqeuclidean_manifold, M2.manifolds) - elseif hasproperty(M2, :native) && (getproperty(M2, :native) isa SqEuclidean) - return true - end - return false -end - - -function _armijo_max_decreases(initial_stepsize::Real, contraction::Real, alpha_min::Real) - initial_stepsize <= alpha_min && return 0 - (contraction <= 0 || contraction >= 1) && return 1000 - n = floor(Int, log(alpha_min / initial_stepsize) / log(contraction)) - return max(n, 0) -end - - -function _adaptive_initial_stepsize( - M, - p0, - model_grad, - retraction_method, - base_stepsize::T; - alpha_min::T, - scale_c::T = one(T), - clamp_low_factor::T = T(0.1), - clamp_high_factor::T = T(10), - delta_scale::T = T(1e-3), -) where {T<:AbstractFloat} - g0 = model_grad(M, p0) - d = -copy(g0) - dnorm = norm(M, p0, d) - (!isfinite(dnorm) || dnorm <= sqrt(eps(T))) && return base_stepsize - δ = delta_scale / max(dnorm, one(T)) - q = try - retract(M, p0, δ .* d, retraction_method) - catch - return base_stepsize - end - _all_finite(q) || return base_stepsize - gq = model_grad(M, q) - _all_finite(gq) || return base_stepsize - κ_num = inner(M, p0, gq .- g0, d) - κ_den = δ * dnorm^2 - (!isfinite(κ_num) || !isfinite(κ_den) || κ_den <= eps(T)) && return base_stepsize - κ = max(κ_num / κ_den, eps(T)) - α_raw = scale_c / κ - α_low = max(alpha_min, clamp_low_factor * base_stepsize) - α_high = clamp_high_factor * base_stepsize - α = clamp(α_raw, α_low, α_high) - if !isfinite(α) - return base_stepsize - end - return α -end - -function _all_finite(x) - if x isa Number - return isfinite(x) - elseif hasproperty(x, :x) - return all(_all_finite, getproperty(x, :x)) - elseif x isa AbstractArray - return all(isfinite, x) - elseif x isa Tuple - return all(_all_finite, x) - elseif x isa Manifolds.TuckerPoint - return _all_finite(x.hosvd.core) && all(_all_finite, x.hosvd.U) - elseif x isa Manifolds.TuckerTangentVector - return _all_finite(getproperty(x, :Ċ)) && all(_all_finite, getproperty(x, :U̇)) - end - try - return all(isfinite, x) - catch - return false - end -end - - -function _safe_cost_function(model_cost) - return function (M, p) - _all_finite(p) || return Inf - c = model_cost(M, p) - return isfinite(c) ? c : Inf - end -end - - -function _layout_adapt_gradient(model_grad) - return function (M, p) - g = model_grad(M, p) - return _align_layout_like_point(p, g) - end -end - -function _scalar_eltype(p) - if hasproperty(p, :x) || p isa AbstractVector || p isa Tuple - parts = point_parts(p) - isempty(parts) && return Float64 - return _scalar_eltype(first(parts)) - elseif p isa Manifolds.TuckerPoint - return eltype(p.hosvd.core) - elseif p isa Manifolds.TuckerTangentVector - return eltype(getproperty(p, :Ċ)) - elseif p isa Real - return typeof(p) - else - return eltype(p) - end -end - -@inline function _solver_gradient(state) - try - return Manopt.get_gradient(state) - catch - return nothing - end -end - -@inline _scale_solver_tangent(x::Number, scale::Real) = x * scale -_scale_solver_tangent(x::AbstractArray, scale::Real) = x .* scale -_scale_solver_tangent(x::ArrayPartition, scale::Real) = - ArrayPartition(map(part -> _scale_solver_tangent(part, scale), x.x)...) -_scale_solver_tangent(x::Tuple, scale::Real) = - map(part -> _scale_solver_tangent(part, scale), x) - -function _scale_solver_tangent(x, scale::Real) - try - return x .* scale - catch - return scale * x - end -end - -function _relative_solver_functions(model_cost, model_grad, scale::Real) - scale > 0 || return model_cost, model_grad, false - scale == one(scale) && return model_cost, model_grad, false - inv_scale = inv(scale) - return ( - (M, p) -> model_cost(M, p) * inv_scale, - (M, p) -> _scale_solver_tangent(model_grad(M, p), inv_scale), - true, - ) -end - -@inline _solver_has_converged(state) = Manopt.has_converged(state) - - -function _solver_iterations(state, maxiter::Int) - if isdefined(Manopt, :stopped_at) - try - k = Manopt.stopped_at(state) - return k > 0 ? Int(k) : maxiter - catch - end - end - while hasproperty(state, :state) - state = state.state - end - return hasproperty(state, :stop) && hasproperty(state.stop, :at_iteration) ? - state.stop.at_iteration : maxiter -end - -@inline function _solver_iteration_source() - return isdefined(Manopt, :stopped_at) ? :stopped_at : :stop_at_iteration_fallback -end - -function _solver_stats( - model_cost, - model_grad, - M, - p_opt, - state, - ::Nothing; - tol_T, - maxiter::Int, - solver::Symbol, - tiny_grad_tol = nothing, - solver_info = (;), - use_state_gradient::Bool = true, -) - T = typeof(tol_T) - final_cost = model_cost(M, p_opt) - cost_for_error = max(T(0), T(2) * final_cost) - rel_error = sqrt(cost_for_error) - grad_state = use_state_gradient ? _solver_gradient(state) : nothing - grad_from_state = !isnothing(grad_state) - grad_final = - isnothing(grad_state) ? model_grad(M, p_opt) : - _align_layout_like_point(p_opt, grad_state) - grad_norm = norm(M, p_opt, grad_final) - iterations = _solver_iterations(state, maxiter) - converged_grad = - grad_norm < tol_T || (!isnothing(tiny_grad_tol) && grad_norm < tiny_grad_tol) - converged_state = _solver_has_converged(state) - solver_info = merge( - solver_info, - ( - gradient_source = grad_from_state ? :state : :recomputed, - has_converged_state = converged_state, - converged_by_gradient_threshold = converged_grad, - iteration_source = _solver_iteration_source(), - ), - ) - return ( - point = p_opt, - cost = final_cost, - rel_error = rel_error, - grad_norm = grad_norm, - iterations = iterations, - converged = converged_state, - solver = solver, - solver_info = solver_info, - ) -end - -function _solver_stats( - model_cost, - model_grad, - M, - p_opt, - state, - normA2::Real; - tol_T, - maxiter::Int, - solver::Symbol, - tiny_grad_tol = nothing, - solver_info = (;), - use_state_gradient::Bool = true, -) - T = typeof(tol_T) - final_cost = model_cost(M, p_opt) - cost_for_error = max(T(0), T(2) * final_cost) - rel_error = _relative_error_frob_sq(cost_for_error, T(normA2)) - grad_state = use_state_gradient ? _solver_gradient(state) : nothing - grad_from_state = !isnothing(grad_state) - grad_final = - isnothing(grad_state) ? model_grad(M, p_opt) : - _align_layout_like_point(p_opt, grad_state) - grad_norm = norm(M, p_opt, grad_final) - iterations = _solver_iterations(state, maxiter) - converged_grad = - grad_norm < tol_T || (!isnothing(tiny_grad_tol) && grad_norm < tiny_grad_tol) - converged_state = _solver_has_converged(state) - solver_info = merge( - solver_info, - ( - gradient_source = grad_from_state ? :state : :recomputed, - has_converged_state = converged_state, - converged_by_gradient_threshold = converged_grad, - iteration_source = _solver_iteration_source(), - ), - ) - return ( - point = p_opt, - cost = final_cost, - rel_error = rel_error, - grad_norm = grad_norm, - iterations = iterations, - converged = converged_state, - solver = solver, - solver_info = solver_info, - ) -end - -_solver_debug_callbacks(callbacks...) = Any[cb for cb in callbacks if !isnothing(cb)] - -# Allow callers to pass `nothing` (e.g., when verbose/debug is omitted) -_solver_debug_actions(::Nothing, callbacks...) = _solver_debug_callbacks(callbacks...) - -function _solver_debug_actions(verbose::Bool, callbacks...) - callback_actions = _solver_debug_callbacks(callbacks...) - if verbose - io = _SOLVER_DEBUG_SINK - init_group = Manopt.DebugGroup([ - Manopt.DebugDivider("Initial "; io, at_init = true), - Manopt.DebugCost(; io, format = "f(x): %.6e", at_init = true), - Manopt.DebugGradientNorm(; io, format = "|grad f(p)|:%.6e", at_init = true), - Manopt.DebugDivider("\n"; io, at_init = true), - ]) - iter_group = Manopt.DebugEvery( - Manopt.DebugGroup([ - Manopt.DebugIteration(; io, format = "# %-6d"), - Manopt.DebugDivider(" "; io, at_init = true), - Manopt.DebugCost(; io, format = "f(x): %.6e", at_init = true), - Manopt.DebugGradientNorm(; io, format = "|grad f(p)|:%.6e", at_init = true), - Manopt.DebugDivider("\n"; io, at_init = true), - ]), - 100, - ) - iteration_actions = Any[iter_group] - append!(iteration_actions, callback_actions) - return Any[:Start=>Any[init_group], :Iteration=>iteration_actions] - end - return callback_actions -end - -function _solver_progress_callback( - progress, - model_cost, - model_grad, - M; - normA2 = nothing, - diagnostics_recorder = nothing, -) - progress isa NoMethodProgress && return nothing - has_relative_scale = !isnothing(normA2) && normA2 > 0 - target_norm = has_relative_scale ? sqrt(normA2) : nothing - return function (problem, state, k) - k <= 0 && return nothing - p = get_iterate(state) - c = model_cost(M, p) - g = model_grad(M, p) - gnorm = norm(M, p, g) - c_display = has_relative_scale ? sqrt(max(2 * c, zero(c))) : c - gnorm_display = has_relative_scale ? gnorm * target_norm : gnorm - showvalues = Any[("Iter", k), ("Cost", c_display), ("Grad norm", gnorm_display)] - if !isnothing(diagnostics_recorder) - step = diagnostics_recorder.accepted_stepsize_history - trials = diagnostics_recorder.line_search_trial_history - !isempty(step) && push!(showvalues, ("Accepted α", step[end])) - !isempty(trials) && push!(showvalues, ("Line-search trials", trials[end])) - end - update_progress!(progress, k; showvalues) - return nothing - end -end - -mutable struct _SolverDiagnosticsRecorder - first_accepted_stepsize::Float64 - min_accepted_stepsize::Float64 - first_line_search_trials::Int - line_search_trial_count::Int - function_evaluations::Int - gradient_evaluations::Int - prev_function_evaluations::Int - prev_gradient_evaluations::Int - line_search_enabled::Bool - fallback_stepsize::Float64 - accepted_stepsize_history::Vector{Float64} - line_search_trial_history::Vector{Int} -end - -function _SolverDiagnosticsRecorder(; - line_search_enabled::Bool, - fallback_stepsize::Real = NaN, -) - return _SolverDiagnosticsRecorder( - NaN, - Inf, - 0, - 0, - 0, - 0, - 0, - 0, - line_search_enabled, - Float64(fallback_stepsize), - Float64[], - Int[], - ) -end - -function _solver_eval_count(problem, sym::Symbol) - try - count = get_count(get_objective(problem), sym) - return count < 0 ? 0 : Int(count) - catch - return 0 - end -end - -function _solver_diagnostics_callback(recorder::_SolverDiagnosticsRecorder) - return function (problem, state, k) - fe = _solver_eval_count(problem, :Cost) - ge = _solver_eval_count(problem, :Gradient) - if k == 0 - recorder.function_evaluations = fe - recorder.gradient_evaluations = ge - recorder.prev_function_evaluations = fe - recorder.prev_gradient_evaluations = ge - return nothing - end - step = - recorder.line_search_enabled ? get_last_stepsize(problem, state, k) : - recorder.fallback_stepsize - delta_fe = max(fe - recorder.prev_function_evaluations, 0) - step_f = Float64(step) - ls_trials = recorder.line_search_enabled ? max(delta_fe - 1, 0) : 0 - if isnan(recorder.first_accepted_stepsize) - recorder.first_accepted_stepsize = step_f - recorder.first_line_search_trials = ls_trials - end - recorder.min_accepted_stepsize = min(recorder.min_accepted_stepsize, step_f) - recorder.line_search_trial_count += ls_trials - push!(recorder.accepted_stepsize_history, step_f) - push!(recorder.line_search_trial_history, ls_trials) - recorder.function_evaluations = fe - recorder.gradient_evaluations = ge - recorder.prev_function_evaluations = fe - recorder.prev_gradient_evaluations = ge - return nothing - end -end - -function _solver_info(recorder::_SolverDiagnosticsRecorder, iterations::Int) - return ( - total_iterations = iterations, - first_accepted_stepsize = recorder.first_accepted_stepsize, - min_accepted_stepsize = recorder.min_accepted_stepsize, - first_line_search_trials = recorder.first_line_search_trials, - line_search_trial_count = recorder.line_search_trial_count, - function_evaluations = recorder.function_evaluations, - gradient_evaluations = recorder.gradient_evaluations, - accepted_stepsize_history = recorder.accepted_stepsize_history, - line_search_trial_history = recorder.line_search_trial_history, - ) -end - -function _solver_post_step_callback( - model::AbstractDecompositionModel, - M, - normalization::AbstractNormalizationPolicy, - solver_sym::Symbol, -) - normalization isa NoNormalization && return nothing - return function (problem, state, k) - p_old = get_iterate(state) - # Backend postprocessing (for example normalization) runs in canonical - # CP coordinates and then packs back into the solver's current layout. - p_new = post_step!( - model, - p_old; - normalization, - solver = solver_sym, - problem, - state, - iteration = k, - ) - p_new = _align_layout_like_point(p_old, p_new) - p_new === p_old && return nothing - set_iterate!(state, M, p_new) - if solver_sym == :rcg && hasproperty(state, :X) - get_gradient!(problem, state.X, get_iterate(state)) - if hasproperty(state, :δ) - state.δ = -copy(M, get_iterate(state), state.X) - end - hasproperty(state, :β) && (state.β = zero(typeof(state.β))) - if hasproperty(state, :coefficient) && hasproperty(state.coefficient, :storage) - update_storage!(state.coefficient.storage, problem, state) - end - end - return nothing - end -end - function solve_rgd( model_cost, model_egrad, From b1c73cb85445fee64d8f679a0d83b18427e4a8d1 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 20:45:43 +0200 Subject: [PATCH 04/37] add helper documentation --- src/solvers/manopt_helpers.jl | 47 ++++++++++++++++++++++++++++++++--- 1 file changed, 43 insertions(+), 4 deletions(-) diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index a107e27..8c0f7bd 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -1,12 +1,19 @@ +######################## +# This file contains helper functions for communicating with Manopt +######################## + +# Sink Manopt's built-in debug text while still letting debug actions run. struct _SolverDebugSink <: IO end Base.isopen(::_SolverDebugSink) = true Base.write(::_SolverDebugSink, ::UInt8) = 1 Base.write(::_SolverDebugSink, s::Union{String,SubString{String}}) = sizeof(s) Base.unsafe_write(::_SolverDebugSink, ::Ptr{UInt8}, n::UInt) = Int(n) +# Shared no-op IO used by Manopt debug groups. const _SOLVER_DEBUG_SINK = _SolverDebugSink() +# Stop when both the relative cost change and Riemannian gradient norm are small. mutable struct StopWhenCostRelChangeAndGradientLess{T<:Real} <: Manopt.StoppingCriterion tol_cost::T tol_grad::T @@ -16,6 +23,7 @@ mutable struct StopWhenCostRelChangeAndGradientLess{T<:Real} <: Manopt.StoppingC at_iteration::Int end +# Initialize the dual cost-change/gradient stopping rule with empty history. function StopWhenCostRelChangeAndGradientLess(tol_cost::T, tol_grad::T) where {T<:Real} return StopWhenCostRelChangeAndGradientLess{T}( tol_cost, @@ -27,6 +35,7 @@ function StopWhenCostRelChangeAndGradientLess(tol_cost::T, tol_grad::T) where {T ) end +# Update the dual stopping rule from the current Manopt problem/state. function (c::StopWhenCostRelChangeAndGradientLess)(problem, state, i) if i == 0 c.prev_cost = Manopt.get_cost(problem, Manopt.get_iterate(state)) @@ -51,6 +60,7 @@ function (c::StopWhenCostRelChangeAndGradientLess)(problem, state, i) return false end +# Explain why the dual stopping rule stopped, for Manopt status reporting. function Manopt.get_reason(c::StopWhenCostRelChangeAndGradientLess) if c.at_iteration >= 0 return "At iteration $(c.at_iteration) the relative cost change ($(c.last_cost_rel_change)) " * @@ -60,14 +70,17 @@ function Manopt.get_reason(c::StopWhenCostRelChangeAndGradientLess) return "" end +# Summarize the current dual stopping-rule state for Manopt displays. function Manopt.status_summary(c::StopWhenCostRelChangeAndGradientLess) has_stopped = c.at_iteration >= 0 status = has_stopped ? "reached" : "not reached" return "cost rel change < $(c.tol_cost) and |grad f| < $(c.tol_grad): $status" end +# Mark the dual stopping rule as convergence, not failure or exhaustion. Manopt.indicates_convergence(::StopWhenCostRelChangeAndGradientLess) = true +# Print the dual stopping rule compactly in Manopt diagnostics. function Base.show(io::IO, c::StopWhenCostRelChangeAndGradientLess) return print( io, @@ -76,6 +89,7 @@ function Base.show(io::IO, c::StopWhenCostRelChangeAndGradientLess) end +# Extract the final iterate across Manopt versions and nested state wrappers. function _tk_get_solver_result(state) try return Manopt.get_solver_result(state) @@ -95,12 +109,14 @@ function _tk_get_solver_result(state) end +# Match gradients/tangents to the point container Manopt is currently using. @inline _align_layout_like_point(p, x) = hasproperty(p, :x) ? (hasproperty(x, :x) ? x : (x isa Tuple ? ArrayPartition(x...) : x)) : (hasproperty(x, :x) ? Tuple(getproperty(x, :x)) : x) +# Recursively convert tuple-like product points to ArrayPartition layout. function _to_array_partition(x) if x isa ArrayPartition return ArrayPartition(map(_to_array_partition, x.x)...) @@ -113,12 +129,14 @@ function _to_array_partition(x) end +# Adapt an initial point to the layout expected by the solver manifold. function _solver_point(M, p0) M2 = _unwrap_solver_manifold(M) return M2 isa ProductManifold ? _to_array_partition(p0) : p0 end +# Detect pullback nonnegative geometries that need conservative line search. function _contains_sqeuclidean_manifold(M) M2 = _unwrap_solver_manifold(M) if M2 isa SqEuclidean || M2 isa SoftplusEuclidean @@ -134,6 +152,7 @@ function _contains_sqeuclidean_manifold(M) return false end +# Detect strict squaring geometries that require extra Armijo safeguards. function _contains_strict_sqeuclidean_manifold(M) M2 = _unwrap_solver_manifold(M) if M2 isa SqEuclidean @@ -147,6 +166,7 @@ function _contains_strict_sqeuclidean_manifold(M) end +# Compute how many Armijo contractions are needed before alpha_min is reached. function _armijo_max_decreases(initial_stepsize::Real, contraction::Real, alpha_min::Real) initial_stepsize <= alpha_min && return 0 (contraction <= 0 || contraction >= 1) && return 1000 @@ -155,6 +175,7 @@ function _armijo_max_decreases(initial_stepsize::Real, contraction::Real, alpha_ end +# Estimate a guarded initial RGD step from a one-sided curvature probe. function _adaptive_initial_stepsize( M, p0, @@ -194,6 +215,7 @@ function _adaptive_initial_stepsize( return α end +# Check recursively that points, tangents, and nested manifold containers are finite. function _all_finite(x) if x isa Number return isfinite(x) @@ -216,6 +238,7 @@ function _all_finite(x) end +# Wrap a cost so invalid points return Inf instead of poisoning line search. function _safe_cost_function(model_cost) return function (M, p) _all_finite(p) || return Inf @@ -225,6 +248,7 @@ function _safe_cost_function(model_cost) end +# Wrap a gradient so its container layout follows the queried point layout. function _layout_adapt_gradient(model_grad) return function (M, p) g = model_grad(M, p) @@ -232,6 +256,7 @@ function _layout_adapt_gradient(model_grad) end end +# Infer the scalar element type from nested points/tangents used by solvers. function _scalar_eltype(p) if hasproperty(p, :x) || p isa AbstractVector || p isa Tuple parts = point_parts(p) @@ -248,6 +273,7 @@ function _scalar_eltype(p) end end +# Read Manopt's cached gradient when the current state exposes it. @inline function _solver_gradient(state) try return Manopt.get_gradient(state) @@ -256,13 +282,13 @@ end end end +# Scale cost and gradient when using relative error @inline _scale_solver_tangent(x::Number, scale::Real) = x * scale _scale_solver_tangent(x::AbstractArray, scale::Real) = x .* scale _scale_solver_tangent(x::ArrayPartition, scale::Real) = ArrayPartition(map(part -> _scale_solver_tangent(part, scale), x.x)...) _scale_solver_tangent(x::Tuple, scale::Real) = map(part -> _scale_solver_tangent(part, scale), x) - function _scale_solver_tangent(x, scale::Real) try return x .* scale @@ -270,7 +296,6 @@ function _scale_solver_tangent(x, scale::Real) return scale * x end end - function _relative_solver_functions(model_cost, model_grad, scale::Real) scale > 0 || return model_cost, model_grad, false scale == one(scale) && return model_cost, model_grad, false @@ -282,9 +307,10 @@ function _relative_solver_functions(model_cost, model_grad, scale::Real) ) end +# Query Manopt's convergence flag through one local compatibility point. @inline _solver_has_converged(state) = Manopt.has_converged(state) - +# Recover the stopping iteration across Manopt versions and wrapped states. function _solver_iterations(state, maxiter::Int) if isdefined(Manopt, :stopped_at) try @@ -300,10 +326,12 @@ function _solver_iterations(state, maxiter::Int) state.stop.at_iteration : maxiter end +# Report which Manopt iteration API was used for solver metadata. @inline function _solver_iteration_source() return isdefined(Manopt, :stopped_at) ? :stopped_at : :stop_at_iteration_fallback end +# Build common solver result stats when no target norm is available. function _solver_stats( model_cost, model_grad, @@ -353,6 +381,7 @@ function _solver_stats( ) end +# Build common solver result stats and relative error when ||A||^2 is available. function _solver_stats( model_cost, model_grad, @@ -402,11 +431,14 @@ function _solver_stats( ) end +# Collect only active Manopt debug callbacks, dropping omitted hooks. _solver_debug_callbacks(callbacks...) = Any[cb for cb in callbacks if !isnothing(cb)] # Allow callers to pass `nothing` (e.g., when verbose/debug is omitted) +# Build debug actions when the verbose flag was omitted. _solver_debug_actions(::Nothing, callbacks...) = _solver_debug_callbacks(callbacks...) +# Build Manopt debug actions and attach TensorKitchen callback hooks. function _solver_debug_actions(verbose::Bool, callbacks...) callback_actions = _solver_debug_callbacks(callbacks...) if verbose @@ -434,6 +466,7 @@ function _solver_debug_actions(verbose::Bool, callbacks...) return callback_actions end +# Create a Manopt iteration callback that updates TensorKitchen progress output. function _solver_progress_callback( progress, model_cost, @@ -465,6 +498,7 @@ function _solver_progress_callback( end end +# Accumulate line-search, step-size, and evaluation diagnostics during solving. mutable struct _SolverDiagnosticsRecorder first_accepted_stepsize::Float64 min_accepted_stepsize::Float64 @@ -480,6 +514,7 @@ mutable struct _SolverDiagnosticsRecorder line_search_trial_history::Vector{Int} end +# Initialize a diagnostics recorder for solvers with or without line search. function _SolverDiagnosticsRecorder(; line_search_enabled::Bool, fallback_stepsize::Real = NaN, @@ -500,6 +535,7 @@ function _SolverDiagnosticsRecorder(; ) end +# Read Manopt objective evaluation counters defensively across configurations. function _solver_eval_count(problem, sym::Symbol) try count = get_count(get_objective(problem), sym) @@ -509,6 +545,7 @@ function _solver_eval_count(problem, sym::Symbol) end end +# Create a Manopt callback that records per-iteration solver diagnostics. function _solver_diagnostics_callback(recorder::_SolverDiagnosticsRecorder) return function (problem, state, k) fe = _solver_eval_count(problem, :Cost) @@ -542,6 +579,7 @@ function _solver_diagnostics_callback(recorder::_SolverDiagnosticsRecorder) end end +# Convert recorded diagnostics to the public solver_info named tuple. function _solver_info(recorder::_SolverDiagnosticsRecorder, iterations::Int) return ( total_iterations = iterations, @@ -556,6 +594,7 @@ function _solver_info(recorder::_SolverDiagnosticsRecorder, iterations::Int) ) end +# Create a callback that normalizes/postprocesses iterates after each solver step. function _solver_post_step_callback( model::AbstractDecompositionModel, M, @@ -591,4 +630,4 @@ function _solver_post_step_callback( end return nothing end -end \ No newline at end of file +end From f76d069608bf6dab442c9391c4b8dc40b680aea7 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 20:49:00 +0200 Subject: [PATCH 05/37] clean code --- src/solvers/manopt_helpers.jl | 59 +++++------------------------------ 1 file changed, 7 insertions(+), 52 deletions(-) diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index 8c0f7bd..b997a24 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -331,64 +331,19 @@ end return isdefined(Manopt, :stopped_at) ? :stopped_at : :stop_at_iteration_fallback end -# Build common solver result stats when no target norm is available. -function _solver_stats( - model_cost, - model_grad, - M, - p_opt, - state, - ::Nothing; - tol_T, - maxiter::Int, - solver::Symbol, - tiny_grad_tol = nothing, - solver_info = (;), - use_state_gradient::Bool = true, -) - T = typeof(tol_T) - final_cost = model_cost(M, p_opt) - cost_for_error = max(T(0), T(2) * final_cost) - rel_error = sqrt(cost_for_error) - grad_state = use_state_gradient ? _solver_gradient(state) : nothing - grad_from_state = !isnothing(grad_state) - grad_final = - isnothing(grad_state) ? model_grad(M, p_opt) : - _align_layout_like_point(p_opt, grad_state) - grad_norm = norm(M, p_opt, grad_final) - iterations = _solver_iterations(state, maxiter) - converged_grad = - grad_norm < tol_T || (!isnothing(tiny_grad_tol) && grad_norm < tiny_grad_tol) - converged_state = _solver_has_converged(state) - solver_info = merge( - solver_info, - ( - gradient_source = grad_from_state ? :state : :recomputed, - has_converged_state = converged_state, - converged_by_gradient_threshold = converged_grad, - iteration_source = _solver_iteration_source(), - ), - ) - return ( - point = p_opt, - cost = final_cost, - rel_error = rel_error, - grad_norm = grad_norm, - iterations = iterations, - converged = converged_state, - solver = solver, - solver_info = solver_info, - ) -end +# Convert a squared residual value to the public relative-error diagnostic. +_solver_rel_error(cost_for_error, ::Nothing, ::Type) = sqrt(cost_for_error) +_solver_rel_error(cost_for_error, normA2::Real, ::Type{T}) where {T} = + _relative_error_frob_sq(cost_for_error, T(normA2)) -# Build common solver result stats and relative error when ||A||^2 is available. +# Build common solver result stats, with optional ||A||^2 for relative scaling. function _solver_stats( model_cost, model_grad, M, p_opt, state, - normA2::Real; + normA2::Union{Nothing,Real}; tol_T, maxiter::Int, solver::Symbol, @@ -399,7 +354,7 @@ function _solver_stats( T = typeof(tol_T) final_cost = model_cost(M, p_opt) cost_for_error = max(T(0), T(2) * final_cost) - rel_error = _relative_error_frob_sq(cost_for_error, T(normA2)) + rel_error = _solver_rel_error(cost_for_error, normA2, T) grad_state = use_state_gradient ? _solver_gradient(state) : nothing grad_from_state = !isnothing(grad_state) grad_final = From ec80afdc028211015cbb9841a2286c5ba76f8375 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 20:53:32 +0200 Subject: [PATCH 06/37] unify code --- src/api/approx.jl | 3 +-- src/solvers/lbfgs.jl | 17 +---------------- src/solvers/manopt_helpers.jl | 16 ++++++---------- src/solvers/rcg.jl | 17 +---------------- src/solvers/rgd.jl | 34 ++-------------------------------- 5 files changed, 11 insertions(+), 76 deletions(-) diff --git a/src/api/approx.jl b/src/api/approx.jl index 7765c18..3aa4866 100644 --- a/src/api/approx.jl +++ b/src/api/approx.jl @@ -324,8 +324,7 @@ function approx( return _approx_tucker_rank(approx_dispatch(dispatch), base, r, target; kwargs...) end -_reject_generic_rank_dispatch(::AutoApproxDispatch) = nothing -_reject_generic_rank_dispatch(::GenericApproxDispatch) = nothing +_reject_generic_rank_dispatch(::Union{AutoApproxDispatch,GenericApproxDispatch}) = nothing function _reject_generic_rank_dispatch(::CPDApproxDispatch) throw(ArgumentError("approx(...; dispatch=:cpd) requires Manifolds.Segre inputs.")) diff --git a/src/solvers/lbfgs.jl b/src/solvers/lbfgs.jl index fa0009c..626f48f 100644 --- a/src/solvers/lbfgs.jl +++ b/src/solvers/lbfgs.jl @@ -172,22 +172,7 @@ function solve_lbfgs( uses_nonpositive_curvature_behavior = false, ), ) - return isnothing(normA2) ? - _solver_stats( - model_cost, - model_grad_local, - M, - p_opt, - state, - nothing; - tol_T = T(tol), - maxiter, - solver = :lbfgs, - tiny_grad_tol = tol_g_raw, - solver_info, - use_state_gradient = !uses_relative_objective, - ) : - _solver_stats( + return _solver_stats( model_cost, model_grad_local, M, diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index b997a24..ac46e85 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -331,10 +331,10 @@ end return isdefined(Manopt, :stopped_at) ? :stopped_at : :stop_at_iteration_fallback end -# Convert a squared residual value to the public relative-error diagnostic. -_solver_rel_error(cost_for_error, ::Nothing, ::Type) = sqrt(cost_for_error) -_solver_rel_error(cost_for_error, normA2::Real, ::Type{T}) where {T} = - _relative_error_frob_sq(cost_for_error, T(normA2)) +function _solver_rel_error(cost_for_error, normA2::Union{Nothing,Real}, ::Type{T}) where {T} + return isnothing(normA2) ? sqrt(cost_for_error) : + _relative_error_frob_sq(cost_for_error, T(normA2)) +end # Build common solver result stats, with optional ||A||^2 for relative scaling. function _solver_stats( @@ -389,14 +389,10 @@ end # Collect only active Manopt debug callbacks, dropping omitted hooks. _solver_debug_callbacks(callbacks...) = Any[cb for cb in callbacks if !isnothing(cb)] -# Allow callers to pass `nothing` (e.g., when verbose/debug is omitted) -# Build debug actions when the verbose flag was omitted. -_solver_debug_actions(::Nothing, callbacks...) = _solver_debug_callbacks(callbacks...) - # Build Manopt debug actions and attach TensorKitchen callback hooks. -function _solver_debug_actions(verbose::Bool, callbacks...) +function _solver_debug_actions(verbose::Union{Nothing,Bool}, callbacks...) callback_actions = _solver_debug_callbacks(callbacks...) - if verbose + if verbose === true io = _SOLVER_DEBUG_SINK init_group = Manopt.DebugGroup([ Manopt.DebugDivider("Initial "; io, at_init = true), diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index a1ac2d9..f667379 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -181,22 +181,7 @@ function solve_rcg( if !return_stats return p_opt end - return isnothing(normA2) ? - _solver_stats( - model_cost, - model_grad_local, - M, - p_opt, - state, - nothing; - tol_T = T(tol), - maxiter, - solver = :rcg, - tiny_grad_tol = tol_g_raw, - solver_info, - use_state_gradient = !uses_relative_objective, - ) : - _solver_stats( + return _solver_stats( model_cost, model_grad_local, M, diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index be850a6..2490ae8 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -131,22 +131,7 @@ function solve_rgd( if !return_stats return p_opt end - return isnothing(normA2) ? - _solver_stats( - model_cost, - model_grad_local, - M, - p_opt, - state, - nothing; - tol_T = T(tol), - maxiter, - solver = :rgd, - tiny_grad_tol = tol_g_raw, - solver_info, - use_state_gradient = !uses_relative_objective, - ) : - _solver_stats( + return _solver_stats( model_cost, model_grad_local, M, @@ -244,22 +229,7 @@ function solve_rgd_fixed( if !return_stats return p_opt end - return isnothing(normA2) ? - _solver_stats( - model_cost, - model_grad_local, - M, - p_opt, - state, - nothing; - tol_T = T(tol), - maxiter, - solver = :rgd_fixed, - tiny_grad_tol = tiny_grad_tol, - solver_info, - use_state_gradient = !uses_relative_objective, - ) : - _solver_stats( + return _solver_stats( model_cost, model_grad_local, M, From 861e9e97dffdcaab49bc82ab53caad752144258f Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 20:55:56 +0200 Subject: [PATCH 07/37] make normalized_objective=true be the default --- src/api/cpd.jl | 2 -- src/solvers/abstract.jl | 4 ++-- src/solvers/lbfgs.jl | 4 ++-- src/solvers/rcg.jl | 4 ++-- src/solvers/rgd.jl | 8 ++++---- src/solvers/solve_dispatch.jl | 2 +- 6 files changed, 11 insertions(+), 13 deletions(-) diff --git a/src/api/cpd.jl b/src/api/cpd.jl index 2d19de5..c06fcf3 100644 --- a/src/api/cpd.jl +++ b/src/api/cpd.jl @@ -862,8 +862,6 @@ function _run_cpd_solver( refinement_verbose = verbose, vector_transport_method, grad_tol = _cpd_manifold_grad_tol(model, solver, tol), - normalized_objective = solver isa - Union{RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver}, iteration_callbacks, kwargs..., ) diff --git a/src/solvers/abstract.jl b/src/solvers/abstract.jl index dac468d..52afb35 100644 --- a/src/solvers/abstract.jl +++ b/src/solvers/abstract.jl @@ -225,7 +225,7 @@ function solve( return_stats::Bool = false, vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, iteration_callbacks = (), ) where {T<:AbstractFloat} setup = _prepare_solver_problem(model; init, p0, gradient_mode, verbose) @@ -268,7 +268,7 @@ function solve( return_stats::Bool = false, vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, iteration_callbacks = (), ) where {T<:AbstractFloat} setup = _prepare_solver_problem(model; init, p0, gradient_mode) diff --git a/src/solvers/lbfgs.jl b/src/solvers/lbfgs.jl index 626f48f..6c3badf 100644 --- a/src/solvers/lbfgs.jl +++ b/src/solvers/lbfgs.jl @@ -75,7 +75,7 @@ function solve_lbfgs( linesearch::Symbol = :wolfe, preconditioner = nothing, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -200,7 +200,7 @@ function run_second_order_solver( diagnostics_recorder, iteration_callbacks, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) return solve_lbfgs( setup.model_cost, diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index f667379..6d198a1 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -106,7 +106,7 @@ function solve_rcg( diagnostics_recorder = nothing, iteration_callbacks = (), grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -223,7 +223,7 @@ function run_first_order_solver( diagnostics_recorder, iteration_callbacks, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) return solve_rcg( setup.model_cost, diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index 2490ae8..14b7af7 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -19,7 +19,7 @@ function solve_rgd( diagnostics_recorder = nothing, iteration_callbacks = (), grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -163,7 +163,7 @@ function solve_rgd_fixed( diagnostics_recorder = nothing, iteration_callbacks = (), grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -274,7 +274,7 @@ function run_first_order_solver( diagnostics_recorder, iteration_callbacks, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) return solve_rgd( setup.model_cost, @@ -329,7 +329,7 @@ function run_first_order_solver( diagnostics_recorder, iteration_callbacks, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, ) return solve_rgd_fixed( setup.model_cost, diff --git a/src/solvers/solve_dispatch.jl b/src/solvers/solve_dispatch.jl index fcce422..0f132bd 100644 --- a/src/solvers/solve_dispatch.jl +++ b/src/solvers/solve_dispatch.jl @@ -65,7 +65,7 @@ function _solve_with_solver( verbose::Bool, vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, grad_tol = nothing, - normalized_objective::Bool = false, + normalized_objective::Bool = true, iteration_callbacks = (), kwargs..., ) From abe7003b3c9e2ddabf4ef5cb1008b8edbdb44bfc Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 21:15:51 +0200 Subject: [PATCH 08/37] wire in relative norm natively --- src/solvers/lbfgs.jl | 10 ++++---- src/solvers/manopt_helpers.jl | 27 +++++++++++---------- src/solvers/rcg.jl | 10 ++++---- src/solvers/rgd.jl | 21 +++++++---------- test/basic_tests.jl | 44 +++++++++++++++++++++++++++++++++-- 5 files changed, 73 insertions(+), 39 deletions(-) diff --git a/src/solvers/lbfgs.jl b/src/solvers/lbfgs.jl index 6c3badf..24ab188 100644 --- a/src/solvers/lbfgs.jl +++ b/src/solvers/lbfgs.jl @@ -92,7 +92,6 @@ function solve_lbfgs( vector_transport_method grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) tol_g = _dual_stop_grad_tol(T, tol, grad_tol) - tol_g_raw = uses_relative_objective ? tol_g * objective_scale : tol_g dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) stopping = StopWhenAny( StopAfterIteration(maxiter), @@ -116,7 +115,6 @@ function solve_lbfgs( solver_cost, solver_grad, M; - normA2 = uses_relative_objective ? normA2 : nothing, diagnostics_recorder, ) @@ -173,8 +171,8 @@ function solve_lbfgs( ), ) return _solver_stats( - model_cost, - model_grad_local, + solver_cost, + solver_grad, M, p_opt, state, @@ -182,9 +180,9 @@ function solve_lbfgs( tol_T = T(tol), maxiter, solver = :lbfgs, - tiny_grad_tol = tol_g_raw, + tiny_grad_tol = tol_g, solver_info, - use_state_gradient = !uses_relative_objective, + normalized_objective = uses_relative_objective, ) end diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index ac46e85..559cccc 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -296,6 +296,7 @@ function _scale_solver_tangent(x, scale::Real) return scale * x end end +# Build the Manopt objective as squared residual cost, optionally divided by ||target||^2. function _relative_solver_functions(model_cost, model_grad, scale::Real) scale > 0 || return model_cost, model_grad, false scale == one(scale) && return model_cost, model_grad, false @@ -331,12 +332,19 @@ end return isdefined(Manopt, :stopped_at) ? :stopped_at : :stop_at_iteration_fallback end -function _solver_rel_error(cost_for_error, normA2::Union{Nothing,Real}, ::Type{T}) where {T} - return isnothing(normA2) ? sqrt(cost_for_error) : - _relative_error_frob_sq(cost_for_error, T(normA2)) +function _solver_rel_error( + final_cost, + normA2::Union{Nothing,Real}, + normalized_objective::Bool, + ::Type{T}, +) where {T} + cost_for_error = max(T(0), T(2) * final_cost) + normalized_objective && return sqrt(cost_for_error) + return isnothing(normA2) || normA2 <= 0 ? sqrt(cost_for_error) : + sqrt(cost_for_error / T(normA2)) end -# Build common solver result stats, with optional ||A||^2 for relative scaling. +# Build common solver result stats function _solver_stats( model_cost, model_grad, @@ -350,11 +358,11 @@ function _solver_stats( tiny_grad_tol = nothing, solver_info = (;), use_state_gradient::Bool = true, + normalized_objective::Bool = false, ) T = typeof(tol_T) final_cost = model_cost(M, p_opt) - cost_for_error = max(T(0), T(2) * final_cost) - rel_error = _solver_rel_error(cost_for_error, normA2, T) + rel_error = _solver_rel_error(final_cost, normA2, normalized_objective, T) grad_state = use_state_gradient ? _solver_gradient(state) : nothing grad_from_state = !isnothing(grad_state) grad_final = @@ -423,21 +431,16 @@ function _solver_progress_callback( model_cost, model_grad, M; - normA2 = nothing, diagnostics_recorder = nothing, ) progress isa NoMethodProgress && return nothing - has_relative_scale = !isnothing(normA2) && normA2 > 0 - target_norm = has_relative_scale ? sqrt(normA2) : nothing return function (problem, state, k) k <= 0 && return nothing p = get_iterate(state) c = model_cost(M, p) g = model_grad(M, p) gnorm = norm(M, p, g) - c_display = has_relative_scale ? sqrt(max(2 * c, zero(c))) : c - gnorm_display = has_relative_scale ? gnorm * target_norm : gnorm - showvalues = Any[("Iter", k), ("Cost", c_display), ("Grad norm", gnorm_display)] + showvalues = Any[("Iter", k), ("Cost", c), ("Grad norm", gnorm)] if !isnothing(diagnostics_recorder) step = diagnostics_recorder.accepted_stepsize_history trials = diagnostics_recorder.line_search_trial_history diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index 6d198a1..afed20e 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -123,7 +123,6 @@ function solve_rcg( vector_transport_method grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) tol_g = _dual_stop_grad_tol(T, tol, grad_tol) - tol_g_raw = uses_relative_objective ? tol_g * objective_scale : tol_g dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) stopping = StopWhenAny( @@ -143,7 +142,6 @@ function solve_rcg( solver_cost, solver_grad, M; - normA2 = uses_relative_objective ? normA2 : nothing, diagnostics_recorder, ) @@ -182,8 +180,8 @@ function solve_rcg( return p_opt end return _solver_stats( - model_cost, - model_grad_local, + solver_cost, + solver_grad, M, p_opt, state, @@ -191,9 +189,9 @@ function solve_rcg( tol_T = T(tol), maxiter, solver = :rcg, - tiny_grad_tol = tol_g_raw, + tiny_grad_tol = tol_g, solver_info, - use_state_gradient = !uses_relative_objective, + normalized_objective = uses_relative_objective, ) end diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index 14b7af7..2503f61 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -34,7 +34,6 @@ function solve_rgd( armijo_alpha_min = T(1e-8) * objective_scale grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) tol_g = _dual_stop_grad_tol(T, tol, grad_tol) - tol_g_raw = uses_relative_objective ? tol_g * objective_scale : tol_g dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) stopping = StopWhenAny( @@ -91,7 +90,6 @@ function solve_rgd( solver_cost, solver_grad, M; - normA2 = uses_relative_objective ? normA2 : nothing, diagnostics_recorder, ) @@ -132,8 +130,8 @@ function solve_rgd( return p_opt end return _solver_stats( - model_cost, - model_grad_local, + solver_cost, + solver_grad, M, p_opt, state, @@ -141,9 +139,9 @@ function solve_rgd( tol_T = T(tol), maxiter, solver = :rgd, - tiny_grad_tol = tol_g_raw, + tiny_grad_tol = tol_g, solver_info, - use_state_gradient = !uses_relative_objective, + normalized_objective = uses_relative_objective, ) end @@ -175,9 +173,7 @@ function solve_rgd_fixed( _relative_solver_functions(model_cost, model_grad_local, objective_scale) retraction_method = _solver_retraction_method(M, p0_local) grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) - tiny_grad_tol = - isnothing(grad_tol) ? T(1e-5) : - (uses_relative_objective ? T(grad_tol) * objective_scale : T(grad_tol)) + tiny_grad_tol = isnothing(grad_tol) ? T(1e-5) : T(grad_tol) stopping = StopWhenAny(StopAfterIteration(maxiter), StopWhenGradientNormLess(grad_stop_tol)) progress = @@ -192,7 +188,6 @@ function solve_rgd_fixed( solver_cost, solver_grad, M; - normA2 = uses_relative_objective ? normA2 : nothing, diagnostics_recorder, ) state = gradient_descent( @@ -230,8 +225,8 @@ function solve_rgd_fixed( return p_opt end return _solver_stats( - model_cost, - model_grad_local, + solver_cost, + solver_grad, M, p_opt, state, @@ -241,7 +236,7 @@ function solve_rgd_fixed( solver = :rgd_fixed, tiny_grad_tol = tiny_grad_tol, solver_info, - use_state_gradient = !uses_relative_objective, + normalized_objective = uses_relative_objective, ) end diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 041a458..6a818c6 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -136,6 +136,45 @@ end @test out.solver_info.gradient_evaluations >= 1 end +@testset "Manopt normalized_objective controls objective units" begin + A = randn(6, 5, 4) + r = 2 + model = JoinModel(A, r; geometry = :canonical) + p0 = TensorKitchen.initial_point(model, TuckerInit(); verbose = false) + target_norm = norm(A) + + rel_out = solve( + RGDFixedSolver(0.0), + model; + p0, + maxiter = 1, + tol = 0.0, + verbose = false, + return_stats = true, + normalized_objective = true, + ) + abs_out = solve( + RGDFixedSolver(0.0), + model; + p0, + maxiter = 1, + tol = 0.0, + verbose = false, + return_stats = true, + normalized_objective = false, + ) + + @test isapprox(2 * rel_out.cost, rel_out.rel_error^2; rtol = 1e-12, atol = 1e-12) + @test isapprox(abs_out.rel_error, rel_out.rel_error; rtol = 1e-12, atol = 1e-12) + @test isapprox(abs_out.cost, rel_out.cost * target_norm^2; rtol = 1e-12, atol = 1e-12) + @test isapprox( + abs_out.grad_norm, + rel_out.grad_norm * target_norm^2; + rtol = 1e-10, + atol = 1e-10, + ) +end + # ========================================================================= # cpd/cp_rank.jl (cost/egrad functions) # ========================================================================= @@ -433,6 +472,7 @@ end rel = norm(A) > 0 ? norm(X .- A) / norm(A) : norm(X .- A) cost, rel end + expected_solver_cost(solver, cost, rel) = solver == :als ? cost : 0.5 * rel^2 public_columns_unit(res) = all( isapprox(norm(TensorKitchen.factors(res)[m][:, k]), 1; atol = 1e-8, rtol = 1e-8) for m in eachindex(TensorKitchen.factors(res)) for k in eachindex(TensorKitchen.weights(res)) @@ -443,7 +483,7 @@ end for solver in (:rgd, :rcg) res = cpd(A1, 1; solver = solver, nonnegative = true, maxiter = 4, verbose = false) cost, rel = explicit_stats(A1, res) - @test res.cost ≈ cost atol = 1e-8 rtol = 1e-8 + @test res.cost ≈ expected_solver_cost(solver, cost, rel) atol = 1e-8 rtol = 1e-8 @test res.rel_error ≈ rel atol = 1e-8 rtol = 1e-8 @test public_columns_unit(res) end @@ -453,7 +493,7 @@ end for solver in (:als, :rgd, :rcg) res = cpd(Ar, 3; solver = solver, nonnegative = true, maxiter = 4, verbose = false) cost, rel = explicit_stats(Ar, res) - @test res.cost ≈ cost atol = 1e-8 rtol = 1e-8 + @test res.cost ≈ expected_solver_cost(solver, cost, rel) atol = 1e-8 rtol = 1e-8 @test res.rel_error ≈ rel atol = 1e-8 rtol = 1e-8 @test public_columns_unit(res) end From 6a5aa55530b04916653cd69697c6bba4a383f573 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 21:16:32 +0200 Subject: [PATCH 09/37] compactify progress meter --- src/solvers/manopt_helpers.jl | 6 ------ 1 file changed, 6 deletions(-) diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index 559cccc..1ae3d2e 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -441,12 +441,6 @@ function _solver_progress_callback( g = model_grad(M, p) gnorm = norm(M, p, g) showvalues = Any[("Iter", k), ("Cost", c), ("Grad norm", gnorm)] - if !isnothing(diagnostics_recorder) - step = diagnostics_recorder.accepted_stepsize_history - trials = diagnostics_recorder.line_search_trial_history - !isempty(step) && push!(showvalues, ("Accepted α", step[end])) - !isempty(trials) && push!(showvalues, ("Line-search trials", trials[end])) - end update_progress!(progress, k; showvalues) return nothing end From 7fb5557937b6c0d8e77fa04a99d4c13f70546bbb Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 21:24:36 +0200 Subject: [PATCH 10/37] print progress meter correctly --- src/core/progress.jl | 20 +++++++++++++++++--- 1 file changed, 17 insertions(+), 3 deletions(-) diff --git a/src/core/progress.jl b/src/core/progress.jl index 30c514a..e77f00c 100644 --- a/src/core/progress.jl +++ b/src/core/progress.jl @@ -176,6 +176,20 @@ end p.was_rendered = true return nothing end +@inline _meter_was_printed(meter) = getproperty(getproperty(meter, :core), :printed) +@inline _sync_rendered!(::NoMethodProgress) = nothing +function _sync_rendered!(p::FamilyProgress) + meter = _meter(p) + if !isnothing(meter) && _meter_was_printed(meter) + _mark_rendered!(p) + end + return nothing +end +function _sync_tracker_rendered!(tracker::PhaseProgress) + _sync_rendered!(tracker.initialization) + _sync_rendered!(tracker.refinement) + return nothing +end @inline function _force_visible_phase_finish(tracker::PhaseProgress, progress) return tracker.phase == :refinement && @@ -226,9 +240,7 @@ function update_progress!( Any[("Method", _method_name(progress)); showvalues] end PM.update!(meter, current; showvalues = showvalues_with_method, force) - if force || current < meter.n || _was_rendered(progress) - _mark_rendered!(progress) - end + _sync_rendered!(progress) end return nothing end @@ -244,6 +256,7 @@ function finish_progress!( progress isa NoMethodProgress && return nothing meter = _meter(progress) isnothing(meter) && return nothing + _sync_tracker_rendered!(tracker) showvalues_with_method = if isnothing(showvalues) Any[("Method", _method_name(progress))] @@ -273,6 +286,7 @@ function finish_progress!( isnothing(meter) && return nothing tracker = _current_phase_tracker() if tracker isa PhaseProgress + _sync_tracker_rendered!(tracker) set_phase!(tracker, progress.phase) if active_progress(tracker) === progress return finish_progress!(tracker; current, showvalues) From 9efdcde9882a73ae85af5f969d302f72cd1d0f0c Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 21:24:52 +0200 Subject: [PATCH 11/37] run JuliaFormatter --- src/core/progress.jl | 8 +------- 1 file changed, 1 insertion(+), 7 deletions(-) diff --git a/src/core/progress.jl b/src/core/progress.jl index e77f00c..4136f1d 100644 --- a/src/core/progress.jl +++ b/src/core/progress.jl @@ -200,13 +200,7 @@ end function _render_unrendered_completion!(meter, showvalues) # ProgressMeter does not render a meter that reaches 100% before its first # visible update. Give it one display-only step before completion. - PM.update!( - meter, - meter.n; - showvalues, - force = true, - max_steps = meter.n + 1, - ) + PM.update!(meter, meter.n; showvalues, force = true, max_steps = meter.n + 1) return nothing end From 98c2d1e2c829741a338de180a46663ac77f555b8 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Wed, 17 Jun 2026 19:24:43 +0000 Subject: [PATCH 12/37] format code with JuliaFormatter --- Project.toml | 2 ++ src/core/tensor_ops.jl | 4 ++-- src/cpd/core/cpd_init.jl | 2 +- src/tucker/sthosvd.jl | 4 ++-- 4 files changed, 7 insertions(+), 5 deletions(-) diff --git a/Project.toml b/Project.toml index bd4414c..cd61afc 100644 --- a/Project.toml +++ b/Project.toml @@ -3,6 +3,7 @@ uuid = "3630a16b-0f2f-4d88-afbf-c7d59eccf553" version = "0.1.0" [deps] +JuliaFormatter = "98e50ef6-434e-11e9-1051-2b60c6c9e899" LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e" Manifolds = "1cead3c2-87b3-11e9-0ccd-23c62b72b94e" ManifoldsBase = "3362f125-f0bb-47a3-aa74-596ffd7ef2fb" @@ -13,6 +14,7 @@ RecursiveArrayTools = "731186ca-8d62-57ce-b412-fbd966d074cd" TensorOperations = "6aa20fa7-93e2-5fca-9bc0-fbd0db3c71a2" [compat] +JuliaFormatter = "2.8.5" Manifolds = "0.11.20" ManifoldsBase = "2.3.5" Manopt = "0.5.37" 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/cpd/core/cpd_init.jl b/src/cpd/core/cpd_init.jl index b6964e7..bff19c2 100644 --- a/src/cpd/core/cpd_init.jl +++ b/src/cpd/core/cpd_init.jl @@ -44,7 +44,7 @@ function _cp_core_diag_init(core::AbstractArray{T,N}, r::Int) where {T<:Abstract Um[k, k] = one(T) end if rm > 0 && r > n_eye - Um[:, n_eye+1:r] .= random_unit_matrix(rm, r - n_eye, T) + Um[:, (n_eye+1):r] .= random_unit_matrix(rm, r - n_eye, T) end U0[m] = Um end diff --git a/src/tucker/sthosvd.jl b/src/tucker/sthosvd.jl index 2b0e575..cf69651 100644 --- a/src/tucker/sthosvd.jl +++ b/src/tucker/sthosvd.jl @@ -213,7 +213,7 @@ function sthosvd( rk = max(rk, 1) # keep at least rank 1 if verbose - discarded = rk < length(sigma) ? sqrt(sum(sigma[rk+1:end] .^ 2)) : 0.0 + discarded = rk < length(sigma) ? sqrt(sum(sigma[(rk+1):end] .^ 2)) : 0.0 update_progress!( progress, step; @@ -322,7 +322,7 @@ function error_bound(td::TuckerResult{T,N}) where {T,N} rk = size(td.core, k) sigma = td.singular_values[k] if rk < length(sigma) - sq_error += sum(sigma[rk+1:end] .^ 2) + sq_error += sum(sigma[(rk+1):end] .^ 2) end end return sqrt(sq_error) From a5a65cd176a4e568e179d2235241b8ac69705606 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 21:43:52 +0200 Subject: [PATCH 13/37] unify code --- src/solvers/lbfgs.jl | 113 ++++++++------------ src/solvers/manopt_helpers.jl | 130 ++++++++++++++++++++++ src/solvers/rcg.jl | 93 +++++++--------- src/solvers/rgd.jl | 196 ++++++++++++++-------------------- 4 files changed, 292 insertions(+), 240 deletions(-) diff --git a/src/solvers/lbfgs.jl b/src/solvers/lbfgs.jl index 24ab188..1d4c17d 100644 --- a/src/solvers/lbfgs.jl +++ b/src/solvers/lbfgs.jl @@ -77,51 +77,49 @@ function solve_lbfgs( grad_tol = nothing, normalized_objective::Bool = true, ) - p0_local = _solver_point(M, p0) - T = _scalar_eltype(p0_local) - model_grad_raw = isnothing(model_grad) ? grad(model_egrad) : model_grad - model_grad_local = _layout_adapt_gradient(model_grad_raw) - objective_scale = - normalized_objective && !isnothing(normA2) && normA2 > 0 ? T(normA2) : one(T) - solver_cost, solver_grad, uses_relative_objective = - _relative_solver_functions(model_cost, model_grad_local, objective_scale) + setup = _prepare_manopt_solver_functions( + model_cost, + model_egrad, + M, + p0; + normA2, + model_grad, + tol, + grad_tol, + normalized_objective, + ) + 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 - grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) - tol_g = _dual_stop_grad_tol(T, tol, grad_tol) + tol_g = setup.dual_grad_tol dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) - stopping = StopWhenAny( - StopAfterIteration(maxiter), - StopWhenGradientNormLess(grad_stop_tol), - dual_stop, - ) - progress = - maxiter > 0 ? - make_manopt_family_progress( - maxiter; + stopping = _manopt_stopping(maxiter, setup.grad_stop_tol, dual_stop) + callbacks = _manopt_callbacks( + n -> make_manopt_family_progress( + n; enabled = verbose, phase = :refinement, method = "L-BFGS", dt = 0.2, - ) : NoMethodProgress() - diagnostics_callback = - isnothing(diagnostics_recorder) ? nothing : - _solver_diagnostics_callback(diagnostics_recorder) - progress_callback = _solver_progress_callback( - progress, - solver_cost, - solver_grad, + ), + maxiter, + verbose, + setup.solver_cost, + setup.solver_grad, M; diagnostics_recorder, + post_step_callback, + iteration_callbacks, ) state = Manopt.quasi_Newton( M, - solver_cost, - solver_grad, + setup.solver_cost, + setup.solver_grad, p0_local; cautious_update = cautious_update, direction_update = Manopt.InverseBFGS(), @@ -132,35 +130,28 @@ function solve_lbfgs( vector_transport_method = transport, stepsize = _lbfgs_linesearch(linesearch), stopping_criterion = stopping, - debug = _solver_debug_actions( - verbose, - post_step_callback, - diagnostics_callback, - progress_callback, - iteration_callbacks..., - ), + debug = callbacks.debug_actions, count = [:Cost, :Gradient], return_state = true, ) - p_opt = _tk_get_solver_result(state) - iterations_done = _solver_iterations(state, maxiter) - if verbose - finish_progress!( - progress; - current = iterations_done, - showvalues = Any[("Status", "Finished"), ("Iterations", iterations_done)], - ) - end - solver_info = - isnothing(diagnostics_recorder) ? (;) : - _solver_info(diagnostics_recorder, iterations_done) - if !return_stats - return p_opt - end - solver_info = merge( - solver_info, - ( + return _manopt_finish_result( + _tk_get_solver_result(state), + state, + callbacks.progress, + diagnostics_recorder, + setup.solver_cost, + setup.solver_grad, + M, + normA2; + tol_T = T(tol), + maxiter, + solver = :lbfgs, + tiny_grad_tol = tol_g, + return_stats, + verbose, + normalized_objective = setup.uses_relative_objective, + solver_info_extra = ( memory_size = memory_size, cautious_update = cautious_update, initial_scale = initial_scale, @@ -170,20 +161,6 @@ function solve_lbfgs( uses_nonpositive_curvature_behavior = false, ), ) - return _solver_stats( - solver_cost, - solver_grad, - M, - p_opt, - state, - normA2; - tol_T = T(tol), - maxiter, - solver = :lbfgs, - tiny_grad_tol = tol_g, - solver_info, - normalized_objective = uses_relative_objective, - ) end function run_second_order_solver( diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index 1ae3d2e..67b8ac8 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -308,6 +308,136 @@ function _relative_solver_functions(model_cost, model_grad, scale::Real) ) end +# Prepare the shared Manopt point, objective, gradient, and tolerance data. +function _prepare_manopt_solver_functions( + model_cost, + model_egrad, + M, + p0; + normA2 = nothing, + model_grad = nothing, + tol, + grad_tol = nothing, + normalized_objective::Bool, +) + p0_local = _solver_point(M, p0) + T = _scalar_eltype(p0_local) + model_grad_raw = isnothing(model_grad) ? grad(model_egrad) : model_grad + model_grad_local = _layout_adapt_gradient(model_grad_raw) + objective_scale = + normalized_objective && !isnothing(normA2) && normA2 > 0 ? T(normA2) : one(T) + solver_cost, solver_grad, uses_relative_objective = + _relative_solver_functions(model_cost, model_grad_local, objective_scale) + return ( + p0 = p0_local, + T = T, + solver_cost = solver_cost, + solver_grad = solver_grad, + uses_relative_objective = uses_relative_objective, + objective_scale = objective_scale, + grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol), + dual_grad_tol = _dual_stop_grad_tol(T, tol, grad_tol), + ) +end + +# Build the common Manopt stopping rule for iteration, gradient, and dual criteria. +function _manopt_stopping(maxiter::Int, grad_stop_tol, dual_stop; extra = ()) + return StopWhenAny( + StopAfterIteration(maxiter), + StopWhenGradientNormLess(grad_stop_tol), + extra..., + dual_stop, + ) +end + +# Create progress and debug callbacks shared by Manopt-backed solvers. +function _manopt_callbacks( + make_progress::Function, + maxiter::Int, + verbose::Bool, + solver_cost, + solver_grad, + M; + diagnostics_recorder = nothing, + post_step_callback = nothing, + iteration_callbacks = (), +) + progress = maxiter > 0 ? make_progress(maxiter) : NoMethodProgress() + diagnostics_callback = + isnothing(diagnostics_recorder) ? nothing : + _solver_diagnostics_callback(diagnostics_recorder) + progress_callback = _solver_progress_callback( + progress, + solver_cost, + solver_grad, + M; + diagnostics_recorder, + ) + debug_actions = _solver_debug_actions( + verbose, + post_step_callback, + diagnostics_callback, + progress_callback, + iteration_callbacks..., + ) + return ( + progress = progress, + diagnostics_callback = diagnostics_callback, + progress_callback = progress_callback, + debug_actions = debug_actions, + ) +end + +# Finish progress, collect diagnostics, and return common solver stats. +function _manopt_finish_result( + p_opt, + state, + progress, + diagnostics_recorder, + solver_cost, + solver_grad, + M, + normA2; + tol_T, + maxiter::Int, + solver::Symbol, + tiny_grad_tol, + return_stats::Bool, + verbose::Bool, + normalized_objective::Bool, + solver_info_extra = (;), +) + iterations_done = _solver_iterations(state, maxiter) + if verbose + finish_progress!( + progress; + current = iterations_done, + showvalues = Any[("Status", "Finished"), ("Iterations", iterations_done)], + ) + end + solver_info = + isnothing(diagnostics_recorder) ? (;) : + _solver_info(diagnostics_recorder, iterations_done) + solver_info = merge(solver_info, solver_info_extra) + if !return_stats + return p_opt + end + return _solver_stats( + solver_cost, + solver_grad, + M, + p_opt, + state, + normA2; + tol_T, + maxiter, + solver, + tiny_grad_tol, + solver_info, + normalized_objective, + ) +end + # Query Manopt's convergence flag through one local compatibility point. @inline _solver_has_converged(state) = Manopt.has_converged(state) diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index afed20e..7204a10 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -108,90 +108,69 @@ function solve_rcg( grad_tol = nothing, normalized_objective::Bool = true, ) - p0_local = _solver_point(M, p0) - T = _scalar_eltype(p0_local) - model_grad_raw = isnothing(model_grad) ? grad(model_egrad) : model_grad - model_grad_local = _layout_adapt_gradient(model_grad_raw) - objective_scale = - normalized_objective && !isnothing(normA2) && normA2 > 0 ? T(normA2) : one(T) - solver_cost, solver_grad, uses_relative_objective = - _relative_solver_functions(model_cost, model_grad_local, objective_scale) + setup = _prepare_manopt_solver_functions( + model_cost, + model_egrad, + M, + p0; + normA2, + model_grad, + tol, + grad_tol, + normalized_objective, + ) + 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 - grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) - tol_g = _dual_stop_grad_tol(T, tol, grad_tol) + tol_g = setup.dual_grad_tol dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) - stopping = StopWhenAny( - StopAfterIteration(maxiter), - StopWhenGradientNormLess(grad_stop_tol), - dual_stop, - ) - progress = - maxiter > 0 ? - make_rcg_progress(maxiter; enabled = verbose, phase = :refinement, dt = 0.2) : - NoMethodProgress() - diagnostics_callback = - isnothing(diagnostics_recorder) ? nothing : - _solver_diagnostics_callback(diagnostics_recorder) - progress_callback = _solver_progress_callback( - progress, - solver_cost, - solver_grad, + stopping = _manopt_stopping(maxiter, setup.grad_stop_tol, dual_stop) + callbacks = _manopt_callbacks( + n -> make_rcg_progress(n; enabled = verbose, phase = :refinement, dt = 0.2), + maxiter, + verbose, + setup.solver_cost, + setup.solver_grad, M; diagnostics_recorder, + post_step_callback, + iteration_callbacks, ) state = conjugate_gradient_descent( M, - solver_cost, - solver_grad, + setup.solver_cost, + setup.solver_grad, p0_local; retraction_method = retraction_method, vector_transport_method = transport, stopping_criterion = stopping, - debug = _solver_debug_actions( - verbose, - post_step_callback, - diagnostics_callback, - progress_callback, - iteration_callbacks..., - ), + debug = callbacks.debug_actions, count = [:Cost, :Gradient], return_state = true, ) - p_opt = get_solver_result(state) - iterations_done = _solver_iterations(state, maxiter) - if verbose - finish_progress!( - progress; - current = iterations_done, - showvalues = Any[("Status", "Finished"), ("Iterations", iterations_done)], - ) - end - solver_info = - isnothing(diagnostics_recorder) ? (;) : - _solver_info(diagnostics_recorder, iterations_done) - if !return_stats - return p_opt - end - return _solver_stats( - solver_cost, - solver_grad, - M, - p_opt, + return _manopt_finish_result( + get_solver_result(state), state, + callbacks.progress, + diagnostics_recorder, + setup.solver_cost, + setup.solver_grad, + M, normA2; tol_T = T(tol), maxiter, solver = :rcg, tiny_grad_tol = tol_g, - solver_info, - normalized_objective = uses_relative_objective, + return_stats, + verbose, + normalized_objective = setup.uses_relative_objective, ) end diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index 2503f61..e7517b5 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -21,26 +21,29 @@ function solve_rgd( grad_tol = nothing, normalized_objective::Bool = true, ) - p0_local = _solver_point(M, p0) - T = _scalar_eltype(p0_local) - model_grad_raw = isnothing(model_grad) ? grad(model_egrad) : model_grad - model_grad_local = _layout_adapt_gradient(model_grad_raw) - objective_scale = - normalized_objective && !isnothing(normA2) && normA2 > 0 ? T(normA2) : one(T) - solver_cost_base, solver_grad, uses_relative_objective = - _relative_solver_functions(model_cost, model_grad_local, objective_scale) + setup = _prepare_manopt_solver_functions( + model_cost, + model_egrad, + M, + p0; + normA2, + model_grad, + tol, + grad_tol, + normalized_objective, + ) + p0_local = setup.p0 + T = setup.T retraction_method = _solver_retraction_method(M, p0_local) - stepsize_eff_base = T(stepsize) * objective_scale - armijo_alpha_min = T(1e-8) * objective_scale - grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) - tol_g = _dual_stop_grad_tol(T, tol, grad_tol) + stepsize_eff_base = T(stepsize) * setup.objective_scale + armijo_alpha_min = T(1e-8) * setup.objective_scale + tol_g = setup.dual_grad_tol dual_stop = StopWhenCostRelChangeAndGradientLess(T(tol), tol_g) - - stopping = StopWhenAny( - StopAfterIteration(maxiter), - StopWhenGradientNormLess(grad_stop_tol), - StopWhenStepsizeLess(armijo_alpha_min), - dual_stop, + stopping = _manopt_stopping( + maxiter, + setup.grad_stop_tol, + dual_stop; + extra = (StopWhenStepsizeLess(armijo_alpha_min),), ) use_squaring_armijo = _contains_sqeuclidean_manifold(M) @@ -50,7 +53,7 @@ function solve_rgd( _adaptive_initial_stepsize( M, p0_local, - solver_grad, + setup.solver_grad, retraction_method, stepsize_eff_base; alpha_min = armijo_alpha_min, @@ -64,7 +67,7 @@ function solve_rgd( armijo_additional_decrease = use_strict_sqeuclidean ? ((M, q) -> _all_finite(q)) : ((M, q) -> true) solver_cost = - use_strict_sqeuclidean ? _safe_cost_function(solver_cost_base) : solver_cost_base + use_strict_sqeuclidean ? _safe_cost_function(setup.solver_cost) : setup.solver_cost armijo = Manopt.ArmijoLinesearch( M; retraction_method = retraction_method, @@ -78,70 +81,48 @@ function solve_rgd( additional_decrease_condition = armijo_additional_decrease, ) - progress = - maxiter > 0 ? - make_rgd_progress(maxiter; enabled = verbose, phase = :refinement, dt = 0.2) : - NoMethodProgress() - diagnostics_callback = - isnothing(diagnostics_recorder) ? nothing : - _solver_diagnostics_callback(diagnostics_recorder) - progress_callback = _solver_progress_callback( - progress, + callbacks = _manopt_callbacks( + n -> make_rgd_progress(n; enabled = verbose, phase = :refinement, dt = 0.2), + maxiter, + verbose, solver_cost, - solver_grad, + setup.solver_grad, M; diagnostics_recorder, + post_step_callback, + iteration_callbacks, ) state = gradient_descent( M, solver_cost, - solver_grad, + setup.solver_grad, p0_local; retraction_method = retraction_method, stepsize = armijo, stopping_criterion = stopping, - debug = _solver_debug_actions( - verbose, - post_step_callback, - diagnostics_callback, - progress_callback, - iteration_callbacks..., - ), + debug = callbacks.debug_actions, count = [:Cost, :Gradient], return_state = true, ) - p_opt = get_solver_result(state) - iterations_done = _solver_iterations(state, maxiter) - if verbose - finish_progress!( - progress; - current = iterations_done, - showvalues = Any[("Status", "Finished"), ("Iterations", iterations_done)], - ) - end - solver_info = - isnothing(diagnostics_recorder) ? (;) : - _solver_info(diagnostics_recorder, iterations_done) - solver_info = - merge(solver_info, (initial_stepsize_eff = Float64(initial_stepsize_eff),)) - if !return_stats - return p_opt - end - return _solver_stats( + return _manopt_finish_result( + get_solver_result(state), + state, + callbacks.progress, + diagnostics_recorder, solver_cost, - solver_grad, + setup.solver_grad, M, - p_opt, - state, normA2; tol_T = T(tol), maxiter, solver = :rgd, tiny_grad_tol = tol_g, - solver_info, - normalized_objective = uses_relative_objective, + return_stats, + verbose, + normalized_objective = setup.uses_relative_objective, + solver_info_extra = (initial_stepsize_eff = Float64(initial_stepsize_eff),), ) end @@ -163,80 +144,65 @@ function solve_rgd_fixed( grad_tol = nothing, normalized_objective::Bool = true, ) - p0_local = _solver_point(M, p0) - T = _scalar_eltype(p0_local) - model_grad_raw = isnothing(model_grad) ? grad(model_egrad) : model_grad - model_grad_local = _layout_adapt_gradient(model_grad_raw) - objective_scale = - normalized_objective && !isnothing(normA2) && normA2 > 0 ? T(normA2) : one(T) - solver_cost, solver_grad, uses_relative_objective = - _relative_solver_functions(model_cost, model_grad_local, objective_scale) + setup = _prepare_manopt_solver_functions( + model_cost, + model_egrad, + M, + p0; + normA2, + model_grad, + tol, + grad_tol, + normalized_objective, + ) + p0_local = setup.p0 + T = setup.T retraction_method = _solver_retraction_method(M, p0_local) - grad_stop_tol = isnothing(grad_tol) ? T(tol) : T(grad_tol) tiny_grad_tol = isnothing(grad_tol) ? T(1e-5) : T(grad_tol) - stopping = - StopWhenAny(StopAfterIteration(maxiter), StopWhenGradientNormLess(grad_stop_tol)) - progress = - maxiter > 0 ? - make_rgd_fixed_progress(maxiter; enabled = verbose, phase = :refinement, dt = 0.2) : - NoMethodProgress() - diagnostics_callback = - isnothing(diagnostics_recorder) ? nothing : - _solver_diagnostics_callback(diagnostics_recorder) - progress_callback = _solver_progress_callback( - progress, - solver_cost, - solver_grad, + stopping = StopWhenAny( + StopAfterIteration(maxiter), + StopWhenGradientNormLess(setup.grad_stop_tol), + ) + callbacks = _manopt_callbacks( + n -> make_rgd_fixed_progress(n; enabled = verbose, phase = :refinement, dt = 0.2), + maxiter, + verbose, + setup.solver_cost, + setup.solver_grad, M; diagnostics_recorder, + post_step_callback, + iteration_callbacks, ) state = gradient_descent( M, - solver_cost, - solver_grad, + setup.solver_cost, + setup.solver_grad, p0_local; retraction_method = retraction_method, - stepsize = Manopt.ConstantStepsize(M, T(stepsize) * objective_scale), + stepsize = Manopt.ConstantStepsize(M, T(stepsize) * setup.objective_scale), stopping_criterion = stopping, - debug = _solver_debug_actions( - verbose, - post_step_callback, - diagnostics_callback, - progress_callback, - iteration_callbacks..., - ), + debug = callbacks.debug_actions, count = [:Cost, :Gradient], return_state = true, ) - p_opt = get_solver_result(state) - iterations_done = _solver_iterations(state, maxiter) - if verbose - finish_progress!( - progress; - current = iterations_done, - showvalues = Any[("Status", "Finished"), ("Iterations", iterations_done)], - ) - end - solver_info = - isnothing(diagnostics_recorder) ? (;) : - _solver_info(diagnostics_recorder, iterations_done) - if !return_stats - return p_opt - end - return _solver_stats( - solver_cost, - solver_grad, - M, - p_opt, + return _manopt_finish_result( + get_solver_result(state), state, + callbacks.progress, + diagnostics_recorder, + setup.solver_cost, + setup.solver_grad, + M, normA2; tol_T = T(tol), maxiter, solver = :rgd_fixed, tiny_grad_tol = tiny_grad_tol, - solver_info, - normalized_objective = uses_relative_objective, + return_stats, + verbose, + normalized_objective = setup.uses_relative_objective, ) end From a9656a6a3923ae03a008a50d0eb4532886e31f27 Mon Sep 17 00:00:00 2001 From: PBrdng Date: Wed, 17 Jun 2026 22:00:57 +0200 Subject: [PATCH 14/37] unify more code --- src/api/btd.jl | 4 +- src/api/cpd.jl | 2 +- src/core/types.jl | 4 +- src/join/cpd_backend.jl | 2 +- src/results/conversion.jl | 32 +++++--------- src/results/rel_error.jl | 14 +----- src/solvers/abstract.jl | 90 ++++++++++++++++++++++++++------------- 7 files changed, 77 insertions(+), 71 deletions(-) diff --git a/src/api/btd.jl b/src/api/btd.jl index aa06ded..97d6f02 100644 --- a/src/api/btd.jl +++ b/src/api/btd.jl @@ -24,7 +24,7 @@ function _polish_btd_with_als( max_stagnation_restarts = 0, ) rel_error(als_res) < rel_error(result) || return result - si0 = hasproperty(result, :solver_info) ? solver_info(result) : (;) + si0 = _result_solver_info(result) si = merge( si0, ( @@ -127,7 +127,7 @@ function _btd_warm_start_result( end function _merge_btd_solver_info(result, extra::NamedTuple) - si0 = hasproperty(result, :solver_info) ? result.solver_info : (;) + si0 = _result_solver_info(result) return ( point = result.point, cost = result.cost, diff --git a/src/api/cpd.jl b/src/api/cpd.jl index c06fcf3..e03ffab 100644 --- a/src/api/cpd.jl +++ b/src/api/cpd.jl @@ -9,7 +9,7 @@ function _pullback_eps_value(::Type{T}, pullback_eps) where {T<:AbstractFloat} end function _merge_res_solver_info(res, patch::NamedTuple) - si0 = hasproperty(res, :solver_info) ? solver_info(res) : (;) + si0 = _result_solver_info(res) return ( point = point(res), cost = cost(res), diff --git a/src/core/types.jl b/src/core/types.jl index 8c781a3..d39ee6a 100644 --- a/src/core/types.jl +++ b/src/core/types.jl @@ -363,9 +363,7 @@ solver_info(r::Union{CPDResult,ApproxResult,BTDResult}) = r.solver_info Return the decoded components stored in a decomposition result. """ -components(r::CPDResult) = r.components -components(r::ApproxResult) = r.components -components(r::BTDResult) = r.components +components(r::Union{CPDResult,ApproxResult,BTDResult}) = r.components components(r::NamedTuple) = getproperty(r, :components) """ diff --git a/src/join/cpd_backend.jl b/src/join/cpd_backend.jl index 03493f9..123a6b5 100644 --- a/src/join/cpd_backend.jl +++ b/src/join/cpd_backend.jl @@ -109,7 +109,7 @@ end function _cpd_result(model::JoinModel{<:AbstractFloat,<:CPDBackend}, result, dims, r) m = cpd_model(model) solver_sym = _result_solver_symbol(solver(result)) - si = hasproperty(result, :solver_info) ? solver_info(result) : (;) + si = _result_solver_info(result) als_family = solver_sym in _CP_ALS_FAMILY_SOLVERS if r == 1 diff --git a/src/results/conversion.jl b/src/results/conversion.jl index 7c5795c..15755f8 100644 --- a/src/results/conversion.jl +++ b/src/results/conversion.jl @@ -2,12 +2,13 @@ _result_solver_symbol(solver::Symbol) = solver _result_solver_symbol(solver) = :unknown +_result_solver_info(result) = hasproperty(result, :solver_info) ? solver_info(result) : (;) -function _to_approx_result(model::JoinModel{T}, result) where {T<:AbstractFloat} +function _to_join_result(result_type, model::JoinModel{T}, result) where {T<:AbstractFloat} comps = extract_components(model, result.point) solver_sym = _result_solver_symbol(result.solver) - solver_info = hasproperty(result, :solver_info) ? result.solver_info : (;) - return ApproxResult( + si = _result_solver_info(result) + return result_type( result.point, comps, result.cost, @@ -16,27 +17,16 @@ function _to_approx_result(model::JoinModel{T}, result) where {T<:AbstractFloat} result.iterations, result.converged, solver_sym, - solver_info, + si, ) end -function _to_btd_result(model::JoinModel{T}, result) where {T<:AbstractFloat} - comps = extract_components(model, result.point) - solver_sym = _result_solver_symbol(result.solver) - solver_info = hasproperty(result, :solver_info) ? result.solver_info : (;) - # Reuse solver-reported cost/error instead of reconstructing the full BTD residual again. - return BTDResult( - result.point, - comps, - result.cost, - result.rel_error, - result.grad_norm, - result.iterations, - result.converged, - solver_sym, - solver_info, - ) -end +_to_approx_result(model::JoinModel{T}, result) where {T<:AbstractFloat} = + _to_join_result(ApproxResult, model, result) + +# Reuse solver-reported cost/error instead of reconstructing the full BTD residual again. +_to_btd_result(model::JoinModel{T}, result) where {T<:AbstractFloat} = + _to_join_result(BTDResult, model, result) _to_cpd_result(model, result, dims, r) = throw(ArgumentError("No CPD result converter for model $(typeof(model)).")) diff --git a/src/results/rel_error.jl b/src/results/rel_error.jl index c6ea060..567be1c 100644 --- a/src/results/rel_error.jl +++ b/src/results/rel_error.jl @@ -18,19 +18,7 @@ function rel_error(A::AbstractArray, Ahat::AbstractArray) return relative_frobenius_error(A, Ahat) end -function rel_error(A::AbstractArray, res::CPDResult) - return relative_frobenius_error(A, reconstruct(res)) -end - -function rel_error(A::AbstractArray, tucker_res::TuckerResult) - return relative_frobenius_error(A, reconstruct(tucker_res)) -end - -function rel_error(A::AbstractArray, res::ApproxResult) - return relative_frobenius_error(A, reconstruct(res)) -end - -function rel_error(A::AbstractArray, res::BTDResult) +function rel_error(A::AbstractArray, res::Union{CPDResult,TuckerResult,ApproxResult,BTDResult}) return relative_frobenius_error(A, reconstruct(res)) end diff --git a/src/solvers/abstract.jl b/src/solvers/abstract.jl index 52afb35..5aa7608 100644 --- a/src/solvers/abstract.jl +++ b/src/solvers/abstract.jl @@ -212,6 +212,50 @@ function solve(solver::AbstractSolver, model::AbstractDecompositionModel; kwargs error("solve not implemented for $(typeof(solver))") end +function _solve_ro_solver( + solver::AbstractROSolver, + model::AbstractDecompositionModel; + init, + p0, + maxiter::Int, + tol::Real, + gradient_mode, + normalization::Union{AbstractNormalizationPolicy,Symbol,Nothing}, + verbose::Bool, + return_stats::Bool, + vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing}, + grad_tol, + normalized_objective::Bool, + iteration_callbacks, + diagnostics_recorder, + run_solver::Function, +) + setup = _prepare_solver_problem(model; init, p0, gradient_mode, verbose) + normalization_policy = _normalization_policy(normalization) + supports_normalization_policy(model, normalization_policy) || throw( + ArgumentError( + "Normalization policy $(typeof(normalization_policy)) is not supported for model $(typeof(model)).", + ), + ) + solver_sym = solver_symbol(solver) + post_step_callback = + _solver_post_step_callback(model, setup.M, normalization_policy, solver_sym) + return run_solver( + solver, + setup; + maxiter, + tol, + verbose, + return_stats, + vector_transport_method, + grad_tol, + normalized_objective, + post_step_callback, + diagnostics_recorder, + iteration_callbacks, + ) +end + function solve( solver::AbstractFirstOrderROSolver, model::AbstractDecompositionModel{T}; @@ -228,30 +272,23 @@ function solve( normalized_objective::Bool = true, iteration_callbacks = (), ) where {T<:AbstractFloat} - setup = _prepare_solver_problem(model; init, p0, gradient_mode, verbose) - normalization_policy = _normalization_policy(normalization) - supports_normalization_policy(model, normalization_policy) || throw( - ArgumentError( - "Normalization policy $(typeof(normalization_policy)) is not supported for model $(typeof(model)).", - ), - ) - solver_sym = solver_symbol(solver) - post_step_callback = - _solver_post_step_callback(model, setup.M, normalization_policy, solver_sym) - diagnostics_recorder = first_order_diagnostics_recorder(solver) - return run_first_order_solver( + return _solve_ro_solver( solver, - setup; + model; + init, + p0, maxiter, tol, + gradient_mode, + normalization, verbose, return_stats, vector_transport_method, grad_tol, normalized_objective, - post_step_callback, - diagnostics_recorder, iteration_callbacks, + diagnostics_recorder = first_order_diagnostics_recorder(solver), + run_solver = run_first_order_solver, ) end @@ -271,30 +308,23 @@ function solve( normalized_objective::Bool = true, iteration_callbacks = (), ) where {T<:AbstractFloat} - setup = _prepare_solver_problem(model; init, p0, gradient_mode) - normalization_policy = _normalization_policy(normalization) - supports_normalization_policy(model, normalization_policy) || throw( - ArgumentError( - "Normalization policy $(typeof(normalization_policy)) is not supported for model $(typeof(model)).", - ), - ) - solver_sym = solver_symbol(solver) - post_step_callback = - _solver_post_step_callback(model, setup.M, normalization_policy, solver_sym) - diagnostics_recorder = second_order_diagnostics_recorder(solver) - return run_second_order_solver( + return _solve_ro_solver( solver, - setup; + model; + init, + p0, maxiter, tol, + gradient_mode, + normalization, verbose, return_stats, vector_transport_method, grad_tol, normalized_objective, - post_step_callback, - diagnostics_recorder, iteration_callbacks, + diagnostics_recorder = second_order_diagnostics_recorder(solver), + run_solver = run_second_order_solver, ) end From 85a86d9c046580f6379025aed05c0fa969cdf664 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 18 Jun 2026 00:26:25 +0000 Subject: [PATCH 15/37] run JuliaFormatter --- src/results/rel_error.jl | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/results/rel_error.jl b/src/results/rel_error.jl index 567be1c..b46ed02 100644 --- a/src/results/rel_error.jl +++ b/src/results/rel_error.jl @@ -18,7 +18,10 @@ function rel_error(A::AbstractArray, Ahat::AbstractArray) return relative_frobenius_error(A, Ahat) end -function rel_error(A::AbstractArray, res::Union{CPDResult,TuckerResult,ApproxResult,BTDResult}) +function rel_error( + A::AbstractArray, + res::Union{CPDResult,TuckerResult,ApproxResult,BTDResult}, +) return relative_frobenius_error(A, reconstruct(res)) end From 718dcbf9fded5b5981f7a9afa168f560e93c3805 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 18 Jun 2026 00:32:22 +0000 Subject: [PATCH 16/37] Format: apply JuliaFormatter --- src/core/tensor_ops.jl | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/core/tensor_ops.jl b/src/core/tensor_ops.jl index e5c1427..5ab2f89 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( From ee58c0bb5ba373ca1275004daa4f62767d2a0b45 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 18 Jun 2026 06:12:48 +0200 Subject: [PATCH 17/37] 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 18/37] 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 19/37] 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 ' From eef8c4a44b8f3a7b00e88a52a068c5124f0ea14a Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 19 Jun 2026 04:03:36 +0200 Subject: [PATCH 20/37] Add Riemannian Levenberg-Marquardt solver via Manopt. Wire LMSolver into solve dispatch and cpd/approx/btd APIs, add orthonormal coordinate support on pullback metrics for Jacobian assembly, and cover residual/Jacobian consistency with tests. --- src/api/approx.jl | 3 +- src/api/btd.jl | 1 + src/api/cpd.jl | 9 +- src/backend.jl | 1 + src/core/types.jl | 2 +- src/cpd/core/cp_cost.jl | 2 +- src/manifolds/softplus_metric.jl | 45 +++- src/manifolds/squaring_metric.jl | 43 ++- src/solvers/abstract.jl | 1 + src/solvers/lm.jl | 446 +++++++++++++++++++++++++++++++ src/solvers/manopt_helpers.jl | 2 +- src/solvers/solve_dispatch.jl | 14 +- test/basic_tests.jl | 77 ++++++ 13 files changed, 631 insertions(+), 15 deletions(-) create mode 100644 src/solvers/lm.jl diff --git a/src/api/approx.jl b/src/api/approx.jl index 3aa4866..e380061 100644 --- a/src/api/approx.jl +++ b/src/api/approx.jl @@ -127,11 +127,12 @@ For the generic join path: - `rgd_fixed`: Riemannian gradient descent with fixed step size - `rcg`: Riemannian conjugate gradient - `lbfgs`: Limited-memory quasi-Newton + - `lm`: Levenberg-Marquardt on residual/Jacobian least squares ##Notes## * `:als` is not a solver option for `approx(...)`. However, if `approx(...)` auto-routes to `cpd(...)` or `btd(...)`, then those specialized pipelines may support ALS separately. * `warm_steps` and `warm_init` are not part of the generic `approx(...)` path. Generic joins start from random initial point and then use manifold solvers for refinement. -* For generic mixed joins, use manifold solvers such as `:rgd`, `:rcg`, or `:lbfgs`. +* For generic mixed joins, use manifold solvers such as `:rgd`, `:rcg`, `:lbfgs`, or `:lm`. """ function _approx_manifold_collection( dispatch::AutoApproxDispatch, diff --git a/src/api/btd.jl b/src/api/btd.jl index 97d6f02..1127881 100644 --- a/src/api/btd.jl +++ b/src/api/btd.jl @@ -159,6 +159,7 @@ refines it. Returns a [`BTDResult`](@ref). - `:als`: Alternating least squares. - `:rcg`: Riemannian conjugate gradient. - `:lbfgs`: Limited-memory quasi-Newton refinement. + - `:lm`: Levenberg-Marquardt refinement. ## Extended Options diff --git a/src/api/cpd.jl b/src/api/cpd.jl index e03ffab..1d07f56 100644 --- a/src/api/cpd.jl +++ b/src/api/cpd.jl @@ -134,7 +134,7 @@ function _component_energy_summary(energies::AbstractVector{<:Real}) top1 = ordered[1], top2 = sum(@view ordered[1:min(2, r)]), top3 = sum(@view ordered[1:min(3, r)]), - effective = hhi > 0 ? inv(hhi) : NaN, + effective = hhi > 0 ? one(hhi) / hhi : NaN, argmax_component = argmax(shares), ) end @@ -579,13 +579,13 @@ end function _validate_cpd_solver_supported(solver::AbstractSolver) throw( ArgumentError( - "Unsupported CPD solver $(typeof(solver)). Use :als, :rgd, :rgd_fixed, :rcg, or :lbfgs.", + "Unsupported CPD solver $(typeof(solver)). Use :als, :rgd, :rgd_fixed, :rcg, :lbfgs, or :lm.", ), ) end _validate_cpd_solver_supported( - ::Union{ALSSolver,RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver}, + ::Union{ALSSolver,RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver,LMSolver}, ) = nothing function _validate_cpd_solver_options( @@ -750,7 +750,7 @@ end function _cpd_manifold_grad_tol( model::JoinModel{<:AbstractFloat,<:CPDBackend}, - solver::Union{RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver}, + solver::Union{RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver,LMSolver}, tol::Real, ) return tol @@ -992,6 +992,7 @@ If `r` is omitted, uses the smallest tensor mode as a heuristic rank. - `rgd_fixed`: Riemannian gradient descent with fixed step size - `rcg`: Riemannian conjugate gradient - `lbfgs`: Limited-memory Riemannian quasi-Newton + - `lm`: Levenberg-Marquardt using residual/Jacobian least squares - `als`: Alternating Least Squares ## Extended Options diff --git a/src/backend.jl b/src/backend.jl index 5a96538..3f497fc 100644 --- a/src/backend.jl +++ b/src/backend.jl @@ -8,6 +8,7 @@ include("solvers/rgd.jl") include("solvers/btd_tsd.jl") include("solvers/rcg.jl") include("solvers/lbfgs.jl") +include("solvers/lm.jl") include("results/reconstruct.jl") include("results/rel_error.jl") include("solvers/solve_dispatch.jl") diff --git a/src/core/types.jl b/src/core/types.jl index d39ee6a..9549dcc 100644 --- a/src/core/types.jl +++ b/src/core/types.jl @@ -446,7 +446,7 @@ function LinearAlgebra.cond(R::CPDResult{T}) where {T<:AbstractFloat} U = hcat(Us...) s = svdvals(U) smin = minimum(s) - return iszero(smin) ? T(Inf) : inv(smin) + return iszero(smin) ? T(Inf) : one(T) / smin end """ diff --git a/src/cpd/core/cp_cost.jl b/src/cpd/core/cp_cost.jl index c99547b..a10ef67 100644 --- a/src/cpd/core/cp_cost.jl +++ b/src/cpd/core/cp_cost.jl @@ -69,7 +69,7 @@ end end @inline function _softplus_derivative(x::Real) - x >= 0 ? inv(one(x) + exp(-x)) : begin + x >= 0 ? one(x) / (one(x) + exp(-x)) : begin ex = exp(x) ex / (one(x) + ex) end diff --git a/src/manifolds/softplus_metric.jl b/src/manifolds/softplus_metric.jl index 6edfb39..ed85a5f 100644 --- a/src/manifolds/softplus_metric.jl +++ b/src/manifolds/softplus_metric.jl @@ -7,7 +7,7 @@ export SoftplusEuclidean, softplus_metric_inverse -@inline _sp_sigmoid(x::Real) = x >= 0 ? inv(one(x) + exp(-x)) : begin +@inline _sp_sigmoid(x::Real) = x >= 0 ? one(x) / (one(x) + exp(-x)) : begin ex = exp(x) ex / (one(x) + ex) end @@ -57,8 +57,7 @@ function softplus_metric_diag(M::SoftplusEuclidean, p::AbstractVector) end function softplus_metric_inverse(M::SoftplusEuclidean, p::AbstractVector, X::AbstractVector) - g_inv = inv.(softplus_metric_diag(M, p)) - return g_inv .* X + return X ./ softplus_metric_diag(M, p) end pullback_metric_inverse(M::SoftplusEuclidean, p::AbstractVector, X::AbstractVector) = @@ -68,3 +67,43 @@ function ManifoldsBase.inner(M::SoftplusEuclidean, p, X::AbstractVector, Y::Abst g = softplus_metric_diag(M, p) return dot(X, g .* Y) end + +function ManifoldsBase.get_coordinates_orthonormal( + M::SoftplusEuclidean, + p::AbstractVector, + X::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + return sqrt.(softplus_metric_diag(M, p)) .* X +end + +function ManifoldsBase.get_coordinates_orthonormal!( + M::SoftplusEuclidean, + c, + p::AbstractVector, + X::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + c .= sqrt.(softplus_metric_diag(M, p)) .* X + return c +end + +function ManifoldsBase.get_vector_orthonormal( + M::SoftplusEuclidean, + p::AbstractVector, + c::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + return c ./ sqrt.(softplus_metric_diag(M, p)) +end + +function ManifoldsBase.get_vector_orthonormal!( + M::SoftplusEuclidean, + X, + p::AbstractVector, + c::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + X .= c ./ sqrt.(softplus_metric_diag(M, p)) + return X +end diff --git a/src/manifolds/squaring_metric.jl b/src/manifolds/squaring_metric.jl index 696d1c3..b1d61a3 100644 --- a/src/manifolds/squaring_metric.jl +++ b/src/manifolds/squaring_metric.jl @@ -72,8 +72,7 @@ function pullback_metric_diag(M::SqEuclidean, p::AbstractVector) end function pullback_metric_inverse(M::SqEuclidean, p::AbstractVector, X::AbstractVector) - g_inv = 1.0 ./ pullback_metric_diag(M, p) - return g_inv .* X + return X ./ pullback_metric_diag(M, p) end function ManifoldsBase.inner(M::SqEuclidean, p, X::AbstractVector, Y::AbstractVector) @@ -81,6 +80,46 @@ function ManifoldsBase.inner(M::SqEuclidean, p, X::AbstractVector, Y::AbstractVe return dot(X, g .* Y) end +function ManifoldsBase.get_coordinates_orthonormal( + M::SqEuclidean, + p::AbstractVector, + X::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + return sqrt.(pullback_metric_diag(M, p)) .* X +end + +function ManifoldsBase.get_coordinates_orthonormal!( + M::SqEuclidean, + c, + p::AbstractVector, + X::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + c .= sqrt.(pullback_metric_diag(M, p)) .* X + return c +end + +function ManifoldsBase.get_vector_orthonormal( + M::SqEuclidean, + p::AbstractVector, + c::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + return c ./ sqrt.(pullback_metric_diag(M, p)) +end + +function ManifoldsBase.get_vector_orthonormal!( + M::SqEuclidean, + X, + p::AbstractVector, + c::AbstractVector, + ::ManifoldsBase.RealNumbers, +) + X .= c ./ sqrt.(pullback_metric_diag(M, p)) + return X +end + # ----------------------------- # 2D benchmark objectives for squaring/nonnegative experiments # p = [x, y] diff --git a/src/solvers/abstract.jl b/src/solvers/abstract.jl index 5aa7608..bab4b9e 100644 --- a/src/solvers/abstract.jl +++ b/src/solvers/abstract.jl @@ -471,6 +471,7 @@ function _prepare_solver_problem( ) return ( + model = model, M = M, p0 = p0_local, model_cost = model_cost, diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl new file mode 100644 index 0000000..6568729 --- /dev/null +++ b/src/solvers/lm.jl @@ -0,0 +1,446 @@ +# solvers/lm.jl — Riemannian Levenberg-Marquardt via Manopt nonlinear least squares +export LMSolver + +struct LMSolver <: AbstractSecondOrderROSolver + η::Float64 + damping_term_min::Float64 + β::Float64 + expect_zero_residual::Bool + linear_subsolver::Any +end + +function LMSolver(; + η::Real = 0.2, + damping_term_min::Real = 0.1, + β::Real = 5.0, + expect_zero_residual::Bool = false, + linear_subsolver = Manopt.default_lm_lin_solve!, +) + 0 < η < 1 || throw(ArgumentError("η must satisfy 0 < η < 1, got $η")) + damping_term_min > 0 || + throw(ArgumentError("damping_term_min must be > 0, got $damping_term_min")) + β > 1 || throw(ArgumentError("β must be > 1, got $β")) + return LMSolver( + Float64(η), + Float64(damping_term_min), + Float64(β), + expect_zero_residual, + linear_subsolver, + ) +end + +solver_symbol(::LMSolver) = :lm + +@inline function _lm_objective_scale( + ::Type{T}, + normA2, + normalized_objective::Bool, +) where {T} + return normalized_objective && !isnothing(normA2) && normA2 > 0 ? + one(T) / sqrt(T(normA2)) : one(T) +end + +function _ambient_tangent_vector!(out::AbstractVector, M, p, X) + emb = ManifoldsBase.embed(M, p, X) + length(emb) == length(out) || throw( + DimensionMismatch( + "Tangent embedding length $(length(emb)) does not match output length $(length(out)).", + ), + ) + copyto!(out, vec(emb)) + return out +end + +function _ambient_tangent_vector!(out::AbstractVector, M::Manifolds.Segre, p, X) + copyto!(out, _segre_tangent_tensorvec(p, X)) + return out +end + +function _join_tangent_ambient_vector!( + out::AbstractVector, + backend::Union{JoinBackend,BTDBackend}, + p, + X, +) + parts = point_parts(p) + xparts = point_parts(X) + _check_parts_len(parts, backend.r, "_join_tangent_ambient_vector!") + _check_parts_len(xparts, backend.r, "_join_tangent_ambient_vector!") + fill!(out, zero(eltype(out))) + @inbounds for k = 1:backend.r + _ambient_tangent_vector!( + backend.component_bufs[k], + backend.manifolds[k], + parts[k], + xparts[k], + ) + out .+= backend.component_bufs[k] + end + return out +end + +function _lm_raw_residual_vector( + model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, + p, +) + return copy(_join_residual!(model.backend, p)) +end + +function _lm_raw_jacobian_matrix( + model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, + M, + p; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) + p_work = _solver_point(M, p) + T = _scalar_eltype(p_work) + ambient_dim = length(tensor(model)) + d = manifold_dimension(M) + J = Matrix{T}(undef, ambient_dim, d) + coeff = zeros(T, d) + column = similar(model.backend.work_rec, T, ambient_dim) + @inbounds for j = 1:d + fill!(coeff, zero(T)) + coeff[j] = one(T) + Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + _join_tangent_ambient_vector!(column, model.backend, p_work, Xj) + J[:, j] .= column + end + return J +end + +@inline function _cp_scaled_tangent_factors(λ, U, λ̇, U̇, ::Val{:identity}) + return λ, U, λ̇, U̇ +end + +@inline function _cp_scaled_tangent_factors(λ̃, Ũ, λ̇̃, U̇̃, ::Val{:square}) + λ = λ̃ .^ 2 + U = [Ũ[m] .^ 2 for m in eachindex(Ũ)] + λ̇ = 2 .* λ̃ .* λ̇̃ + U̇ = [2 .* Ũ[m] .* U̇̃[m] for m in eachindex(Ũ)] + return λ, U, λ̇, U̇ +end + +@inline function _cp_scaled_tangent_factors(λ̃, Ũ, λ̇̃, U̇̃, ::Val{:softplus}) + λ = _softplus_value.(λ̃) + U = [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] + λ̇ = _softplus_derivative.(λ̃) .* λ̇̃ + U̇ = [_softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(Ũ)] + return λ, U, λ̇, U̇ +end + +function _cp_rankr_tangent_tensorvec!( + out::AbstractVector{T}, + λ::AbstractVector{T}, + U::Vector{<:AbstractMatrix{T}}, + λ̇::AbstractVector{T}, + U̇::Vector{<:AbstractMatrix{T}}, +) where {T<:AbstractFloat} + fill!(out, zero(T)) + r = length(λ) + @inbounds for k = 1:r + comp = ([λ[k]], [Vector(@view U[m][:, k]) for m in eachindex(U)]...) + xcomp = ([λ̇[k]], [Vector(@view U̇[m][:, k]) for m in eachindex(U̇)]...) + out .+= _segre_tangent_tensorvec(comp, xcomp) + end + return out +end + +function _cp_rank1_tangent_tensorvec!( + out::AbstractVector{T}, + λ::T, + U::Vector{<:AbstractVector{T}}, + λ̇::T, + U̇::Vector{<:AbstractVector{T}}, +) where {T<:AbstractFloat} + comp = ([λ], U...) + xcomp = ([λ̇], U̇...) + copyto!(out, _segre_tangent_tensorvec(comp, xcomp)) + return out +end + +function _lm_raw_residual_vector(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) + return _lm_raw_residual_vector(cpd_model(model), p) +end + +function _lm_raw_jacobian_matrix( + model::JoinModel{<:AbstractFloat,<:CPDBackend}, + M, + p; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) + return _lm_raw_jacobian_matrix(cpd_model(model), M, p; basis) +end + +function _lm_raw_residual_vector(model::Rank1CPDModel{T}, p) where {T<:AbstractFloat} + return vec(embed_point(model, p)) .- vec(model.A) +end + +function _lm_raw_jacobian_matrix( + model::Rank1CPDModel{T}, + M, + p; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) where {T<:AbstractFloat} + p_work = _solver_point(M, p) + ambient_dim = length(model.A) + d = manifold_dimension(M) + J = Matrix{T}(undef, ambient_dim, d) + coeff = zeros(T, d) + λp, Up = unpack_point_rank1(p_work, model.dims) + column = Vector{T}(undef, ambient_dim) + @inbounds for j = 1:d + fill!(coeff, zero(T)) + coeff[j] = one(T) + Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + if model.nonnegative + λ̇p, U̇p = unpack_point_rank1(Xj, model.dims) + kind = _rank1_uses_softplus_metric(model.M) ? Val(:softplus) : Val(:square) + λ, U, λ̇, U̇ = _cp_scaled_tangent_factors( + [λp], + [reshape(u, :, 1) for u in Up], + [λ̇p], + [reshape(u, :, 1) for u in U̇p], + kind, + ) + U_vec = [Vector(@view U[m][:, 1]) for m in eachindex(U)] + U̇_vec = [Vector(@view U̇[m][:, 1]) for m in eachindex(U̇)] + _cp_rank1_tangent_tensorvec!(column, λ[1], U_vec, λ̇[1], U̇_vec) + else + copyto!(column, vec(ManifoldsBase.embed(M, p_work, Xj))) + end + J[:, j] .= column + end + return J +end + +function _lm_raw_residual_vector(model::RankRCPDModel{T}, p) where {T<:AbstractFloat} + return vec(embed_point(model, p)) .- vec(model.A) +end + +function _lm_raw_jacobian_matrix( + model::RankRCPDModel{T}, + M, + p; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) where {T<:AbstractFloat} + p_work = _solver_point(M, p) + ambient_dim = length(model.A) + d = manifold_dimension(M) + J = Matrix{T}(undef, ambient_dim, d) + coeff = zeros(T, d) + column = Vector{T}(undef, ambient_dim) + if model.geometry == :native && !model.nonnegative + pparts = point_parts(p_work) + @inbounds for j = 1:d + fill!(coeff, zero(T)) + coeff[j] = one(T) + Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + xparts = point_parts(Xj) + fill!(column, zero(T)) + for k = 1:model.r + column .+= _segre_tangent_tensorvec(pparts[k], xparts[k]) + end + J[:, j] .= column + end + return J + end + + λp, Up = + model.nonnegative ? unpack_point_rankr(p_work, model.dims, model.r) : + unpack_rankr_canonical(p_work, model.dims, model.r) + kind = + model.nonnegative ? + (model.geometry == :softplus_metric ? Val(:softplus) : Val(:square)) : + Val(:identity) + @inbounds for j = 1:d + fill!(coeff, zero(T)) + coeff[j] = one(T) + Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + λ̇p, U̇p = + model.nonnegative ? unpack_point_rankr(Xj, model.dims, model.r) : + unpack_rankr_canonical(Xj, model.dims, model.r) + λ, U, λ̇, U̇ = _cp_scaled_tangent_factors(λp, Up, λ̇p, U̇p, kind) + _cp_rankr_tangent_tensorvec!(column, λ, U, λ̇, U̇) + J[:, j] .= column + end + return J +end + +function _lm_raw_residual_vector(model::AbstractDecompositionModel, p) + throw(ArgumentError("LMSolver residual is not implemented for model $(typeof(model)).")) +end + +function _lm_raw_jacobian_matrix(model::AbstractDecompositionModel, M, p; basis) + throw(ArgumentError("LMSolver Jacobian is not implemented for model $(typeof(model)).")) +end + +function _lm_residual_function( + model::AbstractDecompositionModel, + ::Type{T}, + normA2, + normalized_objective::Bool, +) where {T<:AbstractFloat} + scale = _lm_objective_scale(T, normA2, normalized_objective) + return (M, p) -> scale .* _lm_raw_residual_vector(model, p) +end + +function _lm_jacobian_function( + model::AbstractDecompositionModel, + ::Type{T}, + normA2, + normalized_objective::Bool; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) where {T<:AbstractFloat} + scale = _lm_objective_scale(T, normA2, normalized_objective) + return (M, p) -> scale .* _lm_raw_jacobian_matrix(model, M, p; basis) +end + +function solve_lm( + model, + model_cost, + model_egrad, + M, + p0; + maxiter::Int = 1000, + tol::Real = 1e-6, + verbose::Bool = true, + return_stats::Bool = false, + normA2 = nothing, + model_grad = nothing, + vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, + post_step_callback = nothing, + diagnostics_recorder = nothing, + iteration_callbacks = (), + η::Real = 0.2, + damping_term_min::Real = 0.1, + β::Real = 5.0, + expect_zero_residual::Bool = false, + linear_subsolver = Manopt.default_lm_lin_solve!, + grad_tol = nothing, + normalized_objective::Bool = true, +) + setup = _prepare_manopt_solver_functions( + model_cost, + model_egrad, + M, + p0; + normA2, + model_grad, + tol, + grad_tol, + normalized_objective, + ) + p0_local = setup.p0 + T = setup.T + basis = ManifoldsBase.DefaultOrthonormalBasis() + residual = _lm_residual_function(model, T, normA2, setup.uses_relative_objective) + jacobian = _lm_jacobian_function(model, T, normA2, setup.uses_relative_objective; basis) + retraction_method = _solver_retraction_method(M, p0_local) + stopping = StopWhenAny( + StopAfterIteration(maxiter), + StopWhenGradientNormLess(setup.grad_stop_tol), + StopWhenStepsizeLess(T(tol)), + StopWhenCostRelChangeAndGradientLess(T(tol), setup.dual_grad_tol), + ) + callbacks = _manopt_callbacks( + n -> make_manopt_family_progress( + n; + enabled = verbose, + phase = :refinement, + method = "LM", + dt = 0.2, + ), + maxiter, + verbose, + setup.solver_cost, + setup.solver_grad, + M; + diagnostics_recorder, + post_step_callback, + iteration_callbacks, + ) + state = Manopt.LevenbergMarquardt( + M, + residual, + jacobian, + p0_local; + evaluation = Manopt.AllocatingEvaluation(), + function_type = Manopt.FunctionVectorialType(), + jacobian_type = Manopt.CoordinateVectorialType(basis), + retraction_method = retraction_method, + stopping_criterion = stopping, + η = η, + damping_term_min = damping_term_min, + β = β, + expect_zero_residual = expect_zero_residual, + linear_subsolver! = linear_subsolver, + debug = callbacks.debug_actions, + return_state = true, + ) + + return _manopt_finish_result( + _tk_get_solver_result(state), + state, + callbacks.progress, + diagnostics_recorder, + setup.solver_cost, + setup.solver_grad, + M, + normA2; + tol_T = T(tol), + maxiter, + solver = :lm, + tiny_grad_tol = setup.dual_grad_tol, + return_stats, + verbose, + normalized_objective = setup.uses_relative_objective, + solver_info_extra = ( + η = Float64(η), + damping_term_min = Float64(damping_term_min), + β = Float64(β), + expect_zero_residual = expect_zero_residual, + uses_vector_transport = !isnothing(vector_transport_method), + ), + ) +end + +function run_second_order_solver( + solver::LMSolver, + setup; + maxiter::Int, + tol::Real, + verbose::Bool, + return_stats::Bool, + vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, + post_step_callback, + diagnostics_recorder, + iteration_callbacks, + grad_tol = nothing, + normalized_objective::Bool = true, +) + return solve_lm( + setup.model, + setup.model_cost, + setup.model_egrad, + setup.M, + setup.p0; + maxiter, + tol, + verbose, + return_stats, + normA2 = setup.normA2, + model_grad = setup.model_grad, + vector_transport_method, + post_step_callback, + diagnostics_recorder, + iteration_callbacks, + η = solver.η, + damping_term_min = solver.damping_term_min, + β = solver.β, + expect_zero_residual = solver.expect_zero_residual, + linear_subsolver = solver.linear_subsolver, + grad_tol, + normalized_objective, + ) +end diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index 67b8ac8..e6bc167 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -300,7 +300,7 @@ end function _relative_solver_functions(model_cost, model_grad, scale::Real) scale > 0 || return model_cost, model_grad, false scale == one(scale) && return model_cost, model_grad, false - inv_scale = inv(scale) + inv_scale = one(scale) / scale return ( (M, p) -> model_cost(M, p) * inv_scale, (M, p) -> _scale_solver_tangent(model_grad(M, p), inv_scale), diff --git a/src/solvers/solve_dispatch.jl b/src/solvers/solve_dispatch.jl index 9bdc801..70e1bf6 100644 --- a/src/solvers/solve_dispatch.jl +++ b/src/solvers/solve_dispatch.jl @@ -3,7 +3,7 @@ function _solver_object(solver, ::Real; kwargs...) throw( ArgumentError( - "Unsupported solver specification $(typeof(solver)). Use a solver symbol such as :als, :rgd, :rgd_fixed, :rcg, :lbfgs, or :btd_tsd, or pass an AbstractSolver object.", + "Unsupported solver specification $(typeof(solver)). Use a solver symbol such as :als, :rgd, :rgd_fixed, :rcg, :lbfgs, :lm, or :btd_tsd, or pass an AbstractSolver object.", ), ) end @@ -44,6 +44,16 @@ function _solver_object(::Val{:lbfgs}, ::Real; kwargs...) ) end +function _solver_object(::Val{:lm}, ::Real; kwargs...) + return LMSolver(; + η = get(kwargs, :η, 0.2), + damping_term_min = get(kwargs, :damping_term_min, 0.1), + β = get(kwargs, :β, 5.0), + expect_zero_residual = get(kwargs, :expect_zero_residual, false), + linear_subsolver = get(kwargs, :linear_subsolver, Manopt.default_lm_lin_solve!), + ) +end + function _solver_object(::Val{:btd_tsd}, stepsize::Real; kwargs...) return BTDTSDSolver(; stepsize, @@ -58,7 +68,7 @@ end function _solver_object(::Val{S}, ::Real; kwargs...) where {S} throw( ArgumentError( - "Unknown solver=$S. Use :als, :rgd, :rgd_fixed, :rcg, :lbfgs, or :btd_tsd.", + "Unknown solver=$S. Use :als, :rgd, :rgd_fixed, :rcg, :lbfgs, :lm, or :btd_tsd.", ), ) end diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 6a818c6..fccb9c9 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -175,6 +175,83 @@ end ) end +@testset "LM residual/Jacobian smoke check matches gradient" begin + cases = ( + JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), randn(5, 4, 3)), + JoinModel(abs.(randn(5, 4, 3)), 2; geometry = :softplus_metric, nonnegative = true), + ) + for model in cases + M = TensorKitchen.manifold(model) + p = TensorKitchen._solver_point( + M, + TensorKitchen.initial_point(model, :random; verbose = false), + ) + basis = ManifoldsBase.DefaultOrthonormalBasis() + r = TensorKitchen._lm_raw_residual_vector(model, p) + J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) + g_coord = transpose(J) * r + g_from_J = ManifoldsBase.get_vector(M, p, g_coord, basis) + g_model = TensorKitchen.rgrad(model, p) + @test norm(M, p, g_from_J - g_model) ≤ 1e-7 * max(1.0, norm(M, p, g_model)) + end +end + +@testset "LM normalized and unnormalized objectives take the same step" begin + A = randn(6, 5, 4) + model = JoinModel((Manifolds.Segre((6, 5, 4)), Manifolds.Segre((6, 5, 4))), A) + p0 = TensorKitchen._solver_point( + TensorKitchen.manifold(model), + TensorKitchen.initial_point(model, :random; verbose = false), + ) + res_rel = solve( + LMSolver(), + model; + p0, + maxiter = 1, + tol = 0.0, + verbose = false, + return_stats = true, + normalized_objective = true, + ) + res_abs = solve( + LMSolver(), + model; + p0, + maxiter = 1, + tol = 0.0, + verbose = false, + return_stats = true, + normalized_objective = false, + ) + buf_rel = similar(model.backend.work_rec) + buf_abs = similar(model.backend.work_rec) + TensorKitchen._join_reconstruct!(buf_rel, model.backend, TensorKitchen.point(res_rel)) + TensorKitchen._join_reconstruct!(buf_abs, model.backend, TensorKitchen.point(res_abs)) + @test maximum(abs.(buf_rel .- buf_abs)) ≤ 1e-10 + @test isapprox(res_rel.rel_error, res_abs.rel_error; rtol = 1e-10, atol = 1e-10) +end + +@testset "cpd/approx accept LMSolver" begin + A = randn(5, 4, 3) + res_cpd_symbol = cpd(A, 2; solver = :lm, maxiter = 2, tol = 1e-6, verbose = false) + @test res_cpd_symbol.solver == :lm + + res_cpd_object = cpd( + A, + 2; + solver = LMSolver(damping_term_min = 1e-2), + maxiter = 2, + tol = 1e-6, + verbose = false, + ) + @test res_cpd_object.solver == :lm + + target = [1.2, -0.4, 0.8] + res_approx = + approx(Manifolds.Sphere(2), target; solver = :lm, maxiter = 2, verbose = false) + @test res_approx.solver == :lm +end + # ========================================================================= # cpd/cp_rank.jl (cost/egrad functions) # ========================================================================= From 6d47cff37ec22758b3a9ddefb24d596422dcd528 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 19 Jun 2026 08:41:56 +0200 Subject: [PATCH 21/37] LM solver accepts nested arrays for BTD, temporarily fix, eventually raise the issue to manopt --- src/solvers/lm.jl | 45 +++++++++++++++++---------------------------- test/basic_tests.jl | 42 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 59 insertions(+), 28 deletions(-) diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 6568729..eb9b8a8 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -31,13 +31,9 @@ end solver_symbol(::LMSolver) = :lm -@inline function _lm_objective_scale( - ::Type{T}, - normA2, - normalized_objective::Bool, -) where {T} - return normalized_objective && !isnothing(normA2) && normA2 > 0 ? - one(T) / sqrt(T(normA2)) : one(T) +@inline function _lm_scaling_factor(::Type{T}, normA2, normalized_objective::Bool) where {T} + return normalized_objective && !isnothing(normA2) && normA2 > 0 ? inv(sqrt(T(normA2))) : + one(T) end function _ambient_tangent_vector!(out::AbstractVector, M, p, X) @@ -51,11 +47,6 @@ function _ambient_tangent_vector!(out::AbstractVector, M, p, X) return out end -function _ambient_tangent_vector!(out::AbstractVector, M::Manifolds.Segre, p, X) - copyto!(out, _segre_tangent_tensorvec(p, X)) - return out -end - function _join_tangent_ambient_vector!( out::AbstractVector, backend::Union{JoinBackend,BTDBackend}, @@ -92,8 +83,7 @@ function _lm_raw_jacobian_matrix( p; basis = ManifoldsBase.DefaultOrthonormalBasis(), ) - p_work = _solver_point(M, p) - T = _scalar_eltype(p_work) + T = _scalar_eltype(p) ambient_dim = length(tensor(model)) d = manifold_dimension(M) J = Matrix{T}(undef, ambient_dim, d) @@ -102,8 +92,8 @@ function _lm_raw_jacobian_matrix( @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) - _join_tangent_ambient_vector!(column, model.backend, p_work, Xj) + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) + _join_tangent_ambient_vector!(column, model.backend, p, Xj) J[:, j] .= column end return J @@ -182,17 +172,16 @@ function _lm_raw_jacobian_matrix( p; basis = ManifoldsBase.DefaultOrthonormalBasis(), ) where {T<:AbstractFloat} - p_work = _solver_point(M, p) ambient_dim = length(model.A) d = manifold_dimension(M) J = Matrix{T}(undef, ambient_dim, d) coeff = zeros(T, d) - λp, Up = unpack_point_rank1(p_work, model.dims) + λp, Up = unpack_point_rank1(p, model.dims) column = Vector{T}(undef, ambient_dim) @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) if model.nonnegative λ̇p, U̇p = unpack_point_rank1(Xj, model.dims) kind = _rank1_uses_softplus_metric(model.M) ? Val(:softplus) : Val(:square) @@ -207,7 +196,7 @@ function _lm_raw_jacobian_matrix( U̇_vec = [Vector(@view U̇[m][:, 1]) for m in eachindex(U̇)] _cp_rank1_tangent_tensorvec!(column, λ[1], U_vec, λ̇[1], U̇_vec) else - copyto!(column, vec(ManifoldsBase.embed(M, p_work, Xj))) + copyto!(column, vec(ManifoldsBase.embed(M, p, Xj))) end J[:, j] .= column end @@ -224,18 +213,17 @@ function _lm_raw_jacobian_matrix( p; basis = ManifoldsBase.DefaultOrthonormalBasis(), ) where {T<:AbstractFloat} - p_work = _solver_point(M, p) ambient_dim = length(model.A) d = manifold_dimension(M) J = Matrix{T}(undef, ambient_dim, d) coeff = zeros(T, d) column = Vector{T}(undef, ambient_dim) if model.geometry == :native && !model.nonnegative - pparts = point_parts(p_work) + pparts = point_parts(p) @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) xparts = point_parts(Xj) fill!(column, zero(T)) for k = 1:model.r @@ -247,8 +235,8 @@ function _lm_raw_jacobian_matrix( end λp, Up = - model.nonnegative ? unpack_point_rankr(p_work, model.dims, model.r) : - unpack_rankr_canonical(p_work, model.dims, model.r) + model.nonnegative ? unpack_point_rankr(p, model.dims, model.r) : + unpack_rankr_canonical(p, model.dims, model.r) kind = model.nonnegative ? (model.geometry == :softplus_metric ? Val(:softplus) : Val(:square)) : @@ -256,7 +244,7 @@ function _lm_raw_jacobian_matrix( @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p_work, coeff, basis) + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) λ̇p, U̇p = model.nonnegative ? unpack_point_rankr(Xj, model.dims, model.r) : unpack_rankr_canonical(Xj, model.dims, model.r) @@ -281,7 +269,7 @@ function _lm_residual_function( normA2, normalized_objective::Bool, ) where {T<:AbstractFloat} - scale = _lm_objective_scale(T, normA2, normalized_objective) + scale = _lm_scaling_factor(T, normA2, normalized_objective) return (M, p) -> scale .* _lm_raw_residual_vector(model, p) end @@ -292,7 +280,7 @@ function _lm_jacobian_function( normalized_objective::Bool; basis = ManifoldsBase.DefaultOrthonormalBasis(), ) where {T<:AbstractFloat} - scale = _lm_objective_scale(T, normA2, normalized_objective) + scale = _lm_scaling_factor(T, normA2, normalized_objective) return (M, p) -> scale .* _lm_raw_jacobian_matrix(model, M, p; basis) end @@ -376,6 +364,7 @@ function solve_lm( expect_zero_residual = expect_zero_residual, linear_subsolver! = linear_subsolver, debug = callbacks.debug_actions, + count = [:Cost, :Gradient], return_state = true, ) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index fccb9c9..487de48 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -252,6 +252,48 @@ end @test res_approx.solver == :lm end +@testset "BTD accepts LMSolver on nested Tucker layouts" begin + A = randn(7, 6, 5) + ranks = (2, 2, 2) + manifolds = TensorKitchen._as_join_manifold_tuple(TuckerJoin(size(A), ranks, 2)) + backend = TensorKitchen._sum_backend_instance(TensorKitchen.BTDBackend, manifolds, A) + model = TensorKitchen.JoinModel{Float64,typeof(backend)}(backend) + p0 = TensorKitchen.initial_point(model, :random; verbose = false) + + @test p0 isa ArrayPartition + @test TensorKitchen.point_parts(p0)[1] isa Manifolds.TuckerPoint + + low = solve( + LMSolver(), + model; + p0, + maxiter = 2, + tol = 1e-6, + verbose = false, + return_stats = true, + ) + low_parts = TensorKitchen.point_parts(low.point) + @test low.solver == :lm + @test low.point isa ArrayPartition + @test length(low_parts) == 2 + @test low_parts[1] isa Manifolds.TuckerPoint + + res_btd = btd( + A, + 2, + ranks; + solver = :lm, + warm_rel_error_gate = nothing, + maxiter = 2, + tol = 1e-6, + verbose = false, + ) + @test res_btd isa BTDResult + @test res_btd.solver == :lm + @test length(res_btd.components) == 2 + @test !get(res_btd.solver_info, :btd_skipped_manifold_polish, false) +end + # ========================================================================= # cpd/cp_rank.jl (cost/egrad functions) # ========================================================================= From 51e1a2dcf5b71449649780b3c6c2e12d09df851d Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 08:54:48 +0200 Subject: [PATCH 22/37] refactor abstract joinModel for RLM --- src/join/join_backend.jl | 54 ++++++++++++++++++++++++++++++++++++---- src/solvers/lm.jl | 27 ++++++-------------- test/basic_tests.jl | 40 +++++++++++++++++++++++------ 3 files changed, 89 insertions(+), 32 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index e8d88d7..d7abd0f 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -448,12 +448,19 @@ function rgrad(model::JoinModel{<:AbstractFloat,<:JoinBackend}, p) return wrap_like_point(p, vals) end -function _ambient_vector!(out::AbstractVector, M, p) +""" + _component_ambient_embedding!(out, M, p) + +Write the ambient embedding of one join component into `out`. +Component manifolds may specialize this hook when their native point structure +admits a more direct embedding than the generic `embed!` path. +""" +function _component_ambient_embedding!(out::AbstractVector, M, p) ManifoldsBase.embed!(M, out, p) return out end -function _ambient_vector!( +function _component_ambient_embedding!( out::AbstractVector, M::Manifolds.Tucker, p::Manifolds.TuckerPoint, @@ -464,7 +471,16 @@ function _ambient_vector!( return out end -function _ambient_vector!(out::AbstractVector, M::Manifolds.Tucker, p) +function _component_ambient_embedding!( + out::AbstractVector, + M::Manifolds.Segre, + p, +) + copyto!(out, _segre_component_tensorvec(p)) + return out +end + +function _component_ambient_embedding!(out::AbstractVector, M::Manifolds.Tucker, p) throw( ArgumentError( "Expected native TuckerPoint for Manifolds.Tucker, got $(typeof(p)).", @@ -472,6 +488,34 @@ function _ambient_vector!(out::AbstractVector, M::Manifolds.Tucker, p) ) end +""" + _component_ambient_pushforward!(out, M, p, X) + +Write the ambient pushforward `DΦ(p)[X]` of one join component into `out`. +This is the component-level differential used by LM Jacobian assembly on +generic `JoinModel`s. +""" +function _component_ambient_pushforward!(out::AbstractVector, M, p, X) + emb = ManifoldsBase.embed(M, p, X) + length(emb) == length(out) || throw( + DimensionMismatch( + "Tangent embedding length $(length(emb)) does not match output length $(length(out)).", + ), + ) + copyto!(out, vec(emb)) + return out +end + +function _component_ambient_pushforward!( + out::AbstractVector, + M::Manifolds.Segre, + p, + X, +) + copyto!(out, _segre_tangent_tensorvec(p, X)) + return out +end + function _subtract_ambient_tensor!( residual::AbstractArray{T,N}, M, @@ -483,7 +527,7 @@ function _subtract_ambient_tensor!( "_subtract_ambient_tensor!: work length $(length(work_vec)) != residual length $(length(residual))", ), ) - _ambient_vector!(work_vec, M, p) + _component_ambient_embedding!(work_vec, M, p) residual_vec = vec(residual) @inbounds for i in eachindex(residual_vec, work_vec) residual_vec[i] -= work_vec[i] @@ -531,7 +575,7 @@ function _join_reconstruct!(out::AbstractArray, backend::Union{JoinBackend,BTDBa @inbounds for k = 1:r # Reconstruct each component into its preallocated workspace. - _ambient_vector!(bufs[k], manifolds[k], parts[k]) + _component_ambient_embedding!(bufs[k], manifolds[k], parts[k]) # Accumulate into the output tensor without allocating a Khatri-Rao-sized object. out .+= bufs[k] diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index eb9b8a8..5625a4f 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -32,19 +32,8 @@ end solver_symbol(::LMSolver) = :lm @inline function _lm_scaling_factor(::Type{T}, normA2, normalized_objective::Bool) where {T} - return normalized_objective && !isnothing(normA2) && normA2 > 0 ? inv(sqrt(T(normA2))) : - one(T) -end - -function _ambient_tangent_vector!(out::AbstractVector, M, p, X) - emb = ManifoldsBase.embed(M, p, X) - length(emb) == length(out) || throw( - DimensionMismatch( - "Tangent embedding length $(length(emb)) does not match output length $(length(out)).", - ), - ) - copyto!(out, vec(emb)) - return out + return normalized_objective && !isnothing(normA2) && normA2 > 0 ? + one(T) / sqrt(T(normA2)) : one(T) end function _join_tangent_ambient_vector!( @@ -59,12 +48,7 @@ function _join_tangent_ambient_vector!( _check_parts_len(xparts, backend.r, "_join_tangent_ambient_vector!") fill!(out, zero(eltype(out))) @inbounds for k = 1:backend.r - _ambient_tangent_vector!( - backend.component_bufs[k], - backend.manifolds[k], - parts[k], - xparts[k], - ) + _component_ambient_pushforward!(backend.component_bufs[k], backend.manifolds[k], parts[k], xparts[k]) out .+= backend.component_bufs[k] end return out @@ -324,6 +308,8 @@ function solve_lm( basis = ManifoldsBase.DefaultOrthonormalBasis() residual = _lm_residual_function(model, T, normA2, setup.uses_relative_objective) jacobian = _lm_jacobian_function(model, T, normA2, setup.uses_relative_objective; basis) + initial_residual_values = residual(M, p0_local) + initial_jacobian_f = jacobian(M, p0_local) retraction_method = _solver_retraction_method(M, p0_local) stopping = StopWhenAny( StopAfterIteration(maxiter), @@ -358,13 +344,14 @@ function solve_lm( jacobian_type = Manopt.CoordinateVectorialType(basis), retraction_method = retraction_method, stopping_criterion = stopping, + initial_residual_values = initial_residual_values, + initial_jacobian_f = initial_jacobian_f, η = η, damping_term_min = damping_term_min, β = β, expect_zero_residual = expect_zero_residual, linear_subsolver! = linear_subsolver, debug = callbacks.debug_actions, - count = [:Cost, :Gradient], return_state = true, ) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 487de48..496705b 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -177,7 +177,7 @@ end @testset "LM residual/Jacobian smoke check matches gradient" begin cases = ( - JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), randn(5, 4, 3)), + JoinModel(randn(5, 4, 3), 2; geometry = :canonical), JoinModel(abs.(randn(5, 4, 3)), 2; geometry = :softplus_metric, nonnegative = true), ) for model in cases @@ -196,9 +196,36 @@ end end end +@testset "LM generic join Jacobian uses component pushforwards" begin + A = randn(5, 4, 3) + model = JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), A) + M = TensorKitchen.manifold(model) + p = TensorKitchen._solver_point(M, TensorKitchen.initial_point(model, :random; verbose = false)) + basis = ManifoldsBase.DefaultOrthonormalBasis() + J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) + @test all(isfinite, J) + + retraction_method = TensorKitchen._solver_retraction_method(M, p) + ϵ = 1e-6 + buf_plus = similar(model.backend.work_rec) + buf_minus = similar(model.backend.work_rec) + d = manifold_dimension(M) + for j = 1:min(d, 3) + coeff = zeros(Float64, d) + coeff[j] = 1.0 + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) + p_plus = ManifoldsBase.retract(M, p, ϵ * Xj, retraction_method) + p_minus = ManifoldsBase.retract(M, p, -ϵ * Xj, retraction_method) + TensorKitchen._join_reconstruct!(buf_plus, model.backend, p_plus) + TensorKitchen._join_reconstruct!(buf_minus, model.backend, p_minus) + fd = (buf_plus .- buf_minus) ./ (2 * ϵ) + @test maximum(abs.(fd .- J[:, j])) ≤ 1e-7 + end +end + @testset "LM normalized and unnormalized objectives take the same step" begin A = randn(6, 5, 4) - model = JoinModel((Manifolds.Segre((6, 5, 4)), Manifolds.Segre((6, 5, 4))), A) + model = JoinModel(A, 2; geometry = :canonical) p0 = TensorKitchen._solver_point( TensorKitchen.manifold(model), TensorKitchen.initial_point(model, :random; verbose = false), @@ -223,11 +250,10 @@ end return_stats = true, normalized_objective = false, ) - buf_rel = similar(model.backend.work_rec) - buf_abs = similar(model.backend.work_rec) - TensorKitchen._join_reconstruct!(buf_rel, model.backend, TensorKitchen.point(res_rel)) - TensorKitchen._join_reconstruct!(buf_abs, model.backend, TensorKitchen.point(res_abs)) - @test maximum(abs.(buf_rel .- buf_abs)) ≤ 1e-10 + rawmodel = TensorKitchen.cpd_model(model) + X_rel = TensorKitchen.embed_point(rawmodel, TensorKitchen.point(res_rel)) + X_abs = TensorKitchen.embed_point(rawmodel, TensorKitchen.point(res_abs)) + @test maximum(abs.(X_rel .- X_abs)) ≤ 1e-10 @test isapprox(res_rel.rel_error, res_abs.rel_error; rtol = 1e-10, atol = 1e-10) end From be39355aa6716db88ca30fb5d9c0c28249f99d00 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 08:56:30 +0200 Subject: [PATCH 23/37] apply formatter --- src/join/join_backend.jl | 13 ++----------- src/solvers/lm.jl | 7 ++++++- test/basic_tests.jl | 5 ++++- 3 files changed, 12 insertions(+), 13 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index d7abd0f..89fc48e 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -471,11 +471,7 @@ function _component_ambient_embedding!( return out end -function _component_ambient_embedding!( - out::AbstractVector, - M::Manifolds.Segre, - p, -) +function _component_ambient_embedding!(out::AbstractVector, M::Manifolds.Segre, p) copyto!(out, _segre_component_tensorvec(p)) return out end @@ -506,12 +502,7 @@ function _component_ambient_pushforward!(out::AbstractVector, M, p, X) return out end -function _component_ambient_pushforward!( - out::AbstractVector, - M::Manifolds.Segre, - p, - X, -) +function _component_ambient_pushforward!(out::AbstractVector, M::Manifolds.Segre, p, X) copyto!(out, _segre_tangent_tensorvec(p, X)) return out end diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 5625a4f..77391f4 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -48,7 +48,12 @@ function _join_tangent_ambient_vector!( _check_parts_len(xparts, backend.r, "_join_tangent_ambient_vector!") fill!(out, zero(eltype(out))) @inbounds for k = 1:backend.r - _component_ambient_pushforward!(backend.component_bufs[k], backend.manifolds[k], parts[k], xparts[k]) + _component_ambient_pushforward!( + backend.component_bufs[k], + backend.manifolds[k], + parts[k], + xparts[k], + ) out .+= backend.component_bufs[k] end return out diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 496705b..2d88071 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -200,7 +200,10 @@ end A = randn(5, 4, 3) model = JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), A) M = TensorKitchen.manifold(model) - p = TensorKitchen._solver_point(M, TensorKitchen.initial_point(model, :random; verbose = false)) + p = TensorKitchen._solver_point( + M, + TensorKitchen.initial_point(model, :random; verbose = false), + ) basis = ManifoldsBase.DefaultOrthonormalBasis() J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) @test all(isfinite, J) From aebd1a3eb8f91bcb5544302c81e77e1e2f9cd441 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 09:14:03 +0200 Subject: [PATCH 24/37] Keep the RLM path for canonical CP (ALS) --- src/core/model.jl | 3 +- src/core/unpack_points.jl | 44 +++++++++++++++++++- src/join/join_backend.jl | 33 +++++++++++++++ test/basic_tests.jl | 85 +++++++++++++++++++++++++++++++++++++++ 4 files changed, 163 insertions(+), 2 deletions(-) diff --git a/src/core/model.jl b/src/core/model.jl index bf2be86..f598072 100644 --- a/src/core/model.jl +++ b/src/core/model.jl @@ -1,5 +1,6 @@ # core/model.jl — Top-level decomposition model interface -export AbstractDecompositionModel, rgrad, supports_rgrad, tensor, cost, post_step! +export AbstractDecompositionModel, + manifold, initial_point, egrad, rgrad, supports_rgrad, tensor, cost, post_step! """ AbstractDecompositionModel{T} diff --git a/src/core/unpack_points.jl b/src/core/unpack_points.jl index 8828e86..e8d3d02 100644 --- a/src/core/unpack_points.jl +++ b/src/core/unpack_points.jl @@ -1,6 +1,11 @@ # core/unpack_points.jl — Point unpacking and legacy/vector interop export pack_point_rank1, - unpack_point_rank1, pack_point_rankr, unpack_point_rankr, unpack_point_rankr_components + unpack_point_rank1, + pack_point_rankr, + unpack_point_rankr, + unpack_point_rankr_components, + canonical_to_joinpoint, + joinpoint_to_canonical function unpack_rankr_native(p, dims::NTuple{N,Int}, r::Int) where {N} parts = normalize_rankr_native_point(p, dims, r) @@ -72,6 +77,43 @@ function unpack_rankr_join(p, dims::NTuple{N,Int}, r::Int) where {N} return λ, U end +""" + canonical_to_joinpoint(λ, U) + canonical_to_joinpoint(p_canonical, dims, r) + +Convert a CPD point from canonical factor-matrix storage to the native +rank-`r` Segre join point layout used by generic `JoinModel((Segre, ...), A)`. + +The conversion preserves the represented tensor but may renormalize component +gauges the same way `pack_rankr_native` does. +""" +function canonical_to_joinpoint( + λ::AbstractVector{T}, + U::Vector{<:AbstractMatrix{T}}, +) where {T<:AbstractFloat} + r = length(λ) + return pack_rankr_native(λ, U, r) +end + +function canonical_to_joinpoint(p, dims::NTuple{N,Int}, r::Int) where {N} + λ, U = unpack_rankr_canonical(p, dims, r) + return canonical_to_joinpoint(λ, U) +end + +""" + joinpoint_to_canonical(p_join, dims, r) + +Convert a native Segre join point layout back to the canonical CPD point +layout `(λ, (u₁¹, …, uᵣ¹), …, (u₁ᴺ, …, uᵣᴺ))`. + +The conversion preserves the represented tensor but may renormalize component +gauges the same way `pack_rankr_canonical` does. +""" +function joinpoint_to_canonical(p, dims::NTuple{N,Int}, r::Int) where {N} + λ, U = unpack_rankr_native(p, dims, r) + return pack_rankr_canonical(λ, U, r) +end + function pack_point_rank1(λ::T, U::Vector{Vector{T}}) where {T<:AbstractFloat} parts = Vector{Vector{T}}(undef, length(U) + 1) parts[1] = T[λ] diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 89fc48e..cf8417e 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -22,6 +22,16 @@ function _as_join_manifold_tuple(manifolds::AbstractVector) end _as_join_manifold_tuple(M::ProductManifold) = Tuple(M.manifolds) + +function _uniform_segre_dims(manifolds::Tuple) + isempty(manifolds) && return nothing + first_manifold = first(manifolds) + first_manifold isa Manifolds.Segre || return nothing + dims = factor_dims(first_manifold) + all(M -> M isa Manifolds.Segre && factor_dims(M) == dims, manifolds) || return nothing + return dims +end + @inline function _check_parts_len(parts, expected::Int, where_fn::AbstractString) length(parts) == expected || throw( DimensionMismatch( @@ -344,6 +354,29 @@ function initial_point( return ArrayPartition(parts...) end +function initial_point( + model::JoinModel{<:AbstractFloat,<:JoinBackend}, + init::ALSWarmStartInit; + verbose::Bool = false, + kwargs..., +) + backend = model.backend + dims = _uniform_segre_dims(backend.manifolds) + isnothing(dims) && throw( + ArgumentError( + "ALSWarmStartInit for a generic JoinModel requires all component manifolds to be Manifolds.Segre with identical factor_dims.", + ), + ) + dims == backend.target_shape || throw( + DimensionMismatch( + "Uniform Segre factor_dims $dims must match target size $(backend.target_shape) for ALS warm start.", + ), + ) + warm_model = JoinModel(backend.target, backend.r; geometry = :canonical) + p_canonical = initial_point(warm_model, init; verbose, kwargs...) + return canonical_to_joinpoint(p_canonical, backend.target_shape, backend.r) +end + # Gradient path: always recomputes the ambient reconstruction and marks the # WORO cache fresh so that the immediately following cost evaluation can reuse it. function _join_residual_grad!(backend::JoinBackend, p) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 2d88071..45d55a4 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -175,6 +175,59 @@ end ) end +@testset "JoinModel: generic (Segre, Segre, ...) backend" begin + dims = (5, 4, 3) + rng = MersenneTwister(4242) + A = randn(rng, dims...) + segres = (Manifolds.Segre(dims), Manifolds.Segre(dims), Manifolds.Segre(dims)) + + model = JoinModel(segres, A) + @test model isa JoinModel + @test model.backend isa TensorKitchen.JoinBackend + @test model.backend.r == 3 + @test tensor(model) == A + @test model.backend.manifolds == segres + + M = TensorKitchen.manifold(model) + @test M isa ProductManifold + @test length(M.manifolds) == 3 + @test all(Mk -> Mk isa Manifolds.Segre && factor_dims(Mk) == dims, M.manifolds) + + model_repeat = JoinModel(Manifolds.Segre(dims), 3, A) + @test model_repeat.backend.r == 3 + @test model_repeat.backend.manifolds == segres + + p = TensorKitchen.initial_point(model, :random; verbose = false) + @test length(TensorKitchen.point_parts(p)) == 3 + + f = cost(model, p) + g = TensorKitchen.egrad(model, p) + rg = rgrad(model, p) + @test isfinite(f) && f >= 0 + @test isfinite(norm(M, p, g)) + @test isfinite(norm(M, p, rg)) + + rec = similar(model.backend.work_rec) + TensorKitchen._join_reconstruct!(rec, model.backend, p) + @test f ≈ 0.5 * sum(abs2, rec .- vec(A)) + + comps = TensorKitchen.extract_components(model, p) + @test length(comps) == 3 + @test all(c -> c.manifold isa Manifolds.Segre, comps) + @test all(c -> size(c.tensor) == dims, comps) + + out = solve( + RGDSolver(1.0), + model; + p0 = p, + maxiter = 2, + tol = 1e-6, + verbose = false, + return_stats = true, + ) + @test isfinite(out.cost) && isfinite(out.rel_error) +end + @testset "LM residual/Jacobian smoke check matches gradient" begin cases = ( JoinModel(randn(5, 4, 3), 2; geometry = :canonical), @@ -672,6 +725,28 @@ end p_warm_sym = TensorKitchen.initial_point(model, :alswarm) @test TensorKitchen.cost(model, p_warm) <= TensorKitchen.cost(model, p_base) + 1e-10 @test isfinite(TensorKitchen.cost(model, p_warm_sym)) + generic_join = JoinModel((Manifolds.Segre(dims), Manifolds.Segre(dims)), A) + p_warm_canonical_match = TensorKitchen.initial_point( + model, + ALSWarmStartInit(2; base_init = TuckerInit()); + verbose = false, + ) + p_join_warm = TensorKitchen.initial_point( + generic_join, + ALSWarmStartInit(2; base_init = TuckerInit()); + verbose = false, + ) + p_join_from_canonical = + TensorKitchen.canonical_to_joinpoint(p_warm_canonical_match, dims, r) + A_join_warm = reconstruct_cpd_rankr( + components_from_factors(TensorKitchen.unpack_rankr_native(p_join_warm, dims, r)...), + ) + A_join_from_canonical = reconstruct_cpd_rankr( + components_from_factors( + TensorKitchen.unpack_rankr_native(p_join_from_canonical, dims, r)..., + ), + ) + @test A_join_warm ≈ A_join_from_canonical res_p0 = cpd(A, r; solver = :rgd, p0 = p0, maxiter = 5, tol = 1e-6, verbose = false) @test res_p0 isa CPDResult @@ -2618,6 +2693,16 @@ end A_native = reconstruct_cpd_rankr(components_from_factors(λn, Un)) @test all(λn .>= 0) @test A_native ≈ A_in + p_canonical = TensorKitchen.pack_rankr_canonical(λ, U, r) + p_join = TensorKitchen.canonical_to_joinpoint(p_canonical, dims, r) + p_canonical_roundtrip = TensorKitchen.joinpoint_to_canonical(p_join, dims, r) + λ_join, U_join = TensorKitchen.unpack_rankr_native(p_join, dims, r) + λ_canon_rt, U_canon_rt = + TensorKitchen.unpack_rankr_canonical(p_canonical_roundtrip, dims, r) + A_join = reconstruct_cpd_rankr(components_from_factors(λ_join, U_join)) + A_canon_rt = reconstruct_cpd_rankr(components_from_factors(λ_canon_rt, U_canon_rt)) + @test A_join ≈ A_in + @test A_canon_rt ≈ A_in A = randn(8, 6, 5) λ0, U0 = cp_init_tucker(A, 3) From 55ebe67ec5005c05c7d75106567383271ecdc5f5 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 11:40:24 +0200 Subject: [PATCH 25/37] add generic JoinModel fitness checks for RLM in basic_tests --- test/basic_tests.jl | 33 +++++++++++++++++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 45d55a4..0f5ab1f 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -328,6 +328,20 @@ end ) @test res_cpd_object.solver == :lm + res_cpd_alswarm_lm = cpd( + A, + 2; + solver = :lm, + init = :alswarm, + warm_steps = 2, + warm_init = :tucker, + maxiter = 2, + tol = 1e-6, + verbose = false, + ) + @test res_cpd_alswarm_lm.solver == :lm + @test isfinite(res_cpd_alswarm_lm.rel_error) + target = [1.2, -0.4, 0.8] res_approx = approx(Manifolds.Sphere(2), target; solver = :lm, maxiter = 2, verbose = false) @@ -374,6 +388,25 @@ end @test res_btd.solver == :lm @test length(res_btd.components) == 2 @test !get(res_btd.solver_info, :btd_skipped_manifold_polish, false) + + res_btd_alswarm_lm = btd( + A, + 2, + ranks; + solver = :lm, + init = :alswarm, + warm_init = BTDHOSVDMultistartInit(2; screening_steps = 0, block_maxiter = 1), + warm_steps = 1, + warm_block_maxiter = 1, + warm_rel_error_gate = nothing, + maxiter = 2, + tol = 1e-6, + verbose = false, + ) + @test res_btd_alswarm_lm isa BTDResult + @test res_btd_alswarm_lm.solver == :lm + @test hasproperty(res_btd_alswarm_lm.solver_info, :btd_als_warm_start_iters) + @test res_btd_alswarm_lm.solver_info.btd_als_warm_start_requested_solver == :lm end # ========================================================================= From caa638a579b94a4d5d2daca7d15f5d1a07cd813e Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 12:50:11 +0200 Subject: [PATCH 26/37] refactor JoinComponent Abstraction --- src/join/join_backend.jl | 220 +++++++++++++++++++++++++++++++-------- src/join/join_model.jl | 18 +++- src/solvers/lm.jl | 3 +- test/basic_tests.jl | 10 +- 4 files changed, 200 insertions(+), 51 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index cf8417e..5e828b4 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -2,6 +2,16 @@ _is_manifold_like(::AbstractManifold) = true _is_manifold_like(_) = false +_is_join_component_like(::JoinComponent) = true +_is_join_component_like(x) = _is_manifold_like(x) + +_wrap_join_component(component::JoinComponent) = component +_wrap_join_component(manifold::AbstractManifold) = JoinComponent(manifold) + +_component_manifold(component::JoinComponent) = component.manifold +_component_manifold(manifold::AbstractManifold) = manifold +_backend_components(backend::JoinBackend) = backend.components +_backend_components(backend::BTDBackend) = backend.manifolds function _as_join_manifold_tuple(manifolds::Tuple) all(_is_manifold_like, manifolds) || throw( @@ -23,12 +33,35 @@ end _as_join_manifold_tuple(M::ProductManifold) = Tuple(M.manifolds) -function _uniform_segre_dims(manifolds::Tuple) - isempty(manifolds) && return nothing - first_manifold = first(manifolds) +function _as_join_component_tuple(components::Tuple) + all(_is_join_component_like, components) || throw( + ArgumentError( + "All join components must be AbstractManifold or JoinComponent. Got types: $(map(typeof, components)).", + ), + ) + return ntuple(k -> _wrap_join_component(components[k]), length(components)) +end + +function _as_join_component_tuple(components::AbstractVector) + all(_is_join_component_like, components) || throw( + ArgumentError( + "All join components must be AbstractManifold or JoinComponent. Got types: $(map(typeof, components)).", + ), + ) + return Tuple(_wrap_join_component(c) for c in components) +end + +_as_join_component_tuple(M::ProductManifold) = _as_join_component_tuple(Tuple(M.manifolds)) + +function _uniform_segre_dims(components::Tuple) + isempty(components) && return nothing + first_manifold = _component_manifold(first(components)) first_manifold isa Manifolds.Segre || return nothing dims = factor_dims(first_manifold) - all(M -> M isa Manifolds.Segre && factor_dims(M) == dims, manifolds) || return nothing + all(c -> begin + M = _component_manifold(c) + M isa Manifolds.Segre && factor_dims(M) == dims + end, components) || return nothing return dims end @@ -48,6 +81,9 @@ _manifold_init(M::Manifolds.Sphere, target, init::Symbol) = _sphere_init(M, targ _manifold_init(M::Manifolds.Segre, target, init::Symbol) = _segre_init(M, target, init) _manifold_init(M::Manifolds.Tucker, target, init::Symbol) = _tucker_init(M, target, init) +_component_init(component, target, init) = + _manifold_init(_component_manifold(component), target, init) + function _manifold_init(M, target, init_sym::Symbol) init_sym == :random && return rand(M) throw( @@ -58,15 +94,16 @@ function _manifold_init(M, target, init_sym::Symbol) end """ - _manifold_egrad(M, p, residual) + _component_egrad(component, p, residual) Compute the component Euclidean gradient induced by a join residual. Tucker components use their native tensor gradient; other components copy the residual. """ -_manifold_egrad(M, p, residual) = copy(residual) -_manifold_egrad(M::Manifolds.Tucker, p, residual) = _tucker_egrad(M, p, residual) +_component_egrad(::DefaultJoinEmbedding, M, p, residual) = copy(residual) +_component_egrad(::DefaultJoinEmbedding, M::Manifolds.Tucker, p, residual) = + _tucker_egrad(M, p, residual) -function _manifold_egrad(M::Manifolds.Segre, p, residual) +function _component_egrad(::DefaultJoinEmbedding, M::Manifolds.Segre, p, residual) dims = factor_dims(M) R = reshape(residual, dims) parts = point_parts(p) @@ -81,6 +118,12 @@ function _manifold_egrad(M::Manifolds.Segre, p, residual) return pack_tangent_rank1_segre(grad_λ, grad_U) end +_component_egrad(component::JoinComponent, p, residual) = + _component_egrad(component.embedding, component.manifold, p, residual) +_component_egrad(M, p, residual) = _component_egrad(DefaultJoinEmbedding(), M, p, residual) + +_manifold_egrad(M, p, residual) = _component_egrad(M, p, residual) + """ _ambient_vector(M, p, target_len) returns AbstractVector @@ -146,6 +189,7 @@ end ambient_length(M::Manifolds.Segre) = prod(factor_dims(M)) ambient_length(M::Manifolds.Tucker) = prod(factor_dims(M)) +ambient_length(component::JoinComponent) = ambient_length(component.manifold) """ _join_vector_workspace_like(target, n) returns AbstractVector @@ -167,15 +211,15 @@ end Ensure every join component embeds into the same flattened ambient space as the target tensor. """ -function _validate_join_ambient_compatibility(manifolds::Tuple, target::AbstractArray) +function _validate_join_ambient_compatibility(components::Tuple, target::AbstractArray) target_len = length(target) failures = String[] - @inbounds for k in eachindex(manifolds) - mk = ambient_length(manifolds[k]) + @inbounds for k in eachindex(components) + mk = ambient_length(components[k]) if mk != target_len push!( failures, - "[$k] $(typeof(manifolds[k])) has ambient length $mk but target has length $target_len", + "[$k] $(typeof(components[k])) has ambient length $mk but target has length $target_len", ) end end @@ -203,14 +247,14 @@ function _sum_backend_instance( end function _sum_backend_parts( - manifolds, + components, target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - r = length(manifolds) + r = length(components) # Keep the original target representation instead of eagerly materializing Array. tgt = target - _validate_join_ambient_compatibility(manifolds, tgt) + _validate_join_ambient_compatibility(components, tgt) tflat = vec(tgt) tgt_len = length(tgt) @@ -218,8 +262,10 @@ function _sum_backend_parts( component_bufs = [_join_vector_workspace_like(tgt, tgt_len) for _ = 1:r] work_rec = _join_vector_workspace_like(tgt, tgt_len) work_residual = _join_vector_workspace_like(tgt, tgt_len) + manifolds = ntuple(k -> _component_manifold(components[k]), r) return (; + components, manifolds, r, target = tgt, @@ -236,13 +282,13 @@ end function _sum_backend_instance( ::Type{JoinBackend}, - manifolds, + components, target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - parts = _sum_backend_parts(manifolds, target; init_point) + parts = _sum_backend_parts(components, target; init_point) return JoinBackend( - parts.manifolds, + parts.components, parts.r, parts.target, parts.target_size, @@ -258,11 +304,11 @@ end function _sum_backend_instance( ::Type{BTDBackend}, - manifolds, + components, target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - parts = _sum_backend_parts(manifolds, target; init_point) + parts = _sum_backend_parts(components, target; init_point) return BTDBackend( parts.manifolds, parts.r, @@ -286,13 +332,14 @@ Construct a generic sum-of-manifolds approximation model from explicit component manifolds and a target tensor. """ function JoinModel( - manifolds::Tuple{Vararg{AbstractManifold}}, + components::Tuple, target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - r = length(manifolds) + components_tuple = _as_join_component_tuple(components) + r = length(components_tuple) r >= 1 || throw(ArgumentError("JoinModel(manifolds, target) needs r >= 1, got r=$r")) - b = _sum_backend_instance(JoinBackend, manifolds, target; init_point) + b = _sum_backend_instance(JoinBackend, components_tuple, target; init_point) return JoinModel{T,typeof(b)}(b) end @@ -302,11 +349,11 @@ end Construct a join model by repeating a base manifold `r` times. """ function JoinModel( - manifolds::AbstractVector, + components::AbstractVector, target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - return JoinModel(_as_join_manifold_tuple(manifolds), target; init_point) + return JoinModel(_as_join_component_tuple(components), target; init_point) end function JoinModel( @@ -314,7 +361,7 @@ function JoinModel( target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - return JoinModel(_as_join_manifold_tuple(M), target; init_point) + return JoinModel(_as_join_component_tuple(M), target; init_point) end function JoinModel( @@ -326,6 +373,15 @@ function JoinModel( return JoinModel(ntuple(_ -> base, r), target; init_point) end +function JoinModel( + base::JoinComponent, + r::Int, + target::AbstractArray{T,N}; + init_point = nothing, +) where {T<:AbstractFloat,N} + return JoinModel(ntuple(_ -> base, r), target; init_point) +end + function JoinModel( base::AbstractManifold, target::AbstractArray{T,N}; @@ -334,6 +390,14 @@ function JoinModel( return JoinModel((base,), target; init_point) end +function JoinModel( + base::JoinComponent, + target::AbstractArray{T,N}; + init_point = nothing, +) where {T<:AbstractFloat,N} + return JoinModel((base,), target; init_point) +end + manifold(model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}) = model.backend.M_product tensor(model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}) = @@ -350,7 +414,7 @@ function initial_point( return backend.init_point(M, init) end parts = - ntuple(k -> _manifold_init(backend.manifolds[k], backend.target, init), backend.r) + ntuple(k -> _component_init(backend.components[k], backend.target, init), backend.r) return ArrayPartition(parts...) end @@ -361,7 +425,7 @@ function initial_point( kwargs..., ) backend = model.backend - dims = _uniform_segre_dims(backend.manifolds) + dims = _uniform_segre_dims(backend.components) isnothing(dims) && throw( ArgumentError( "ALSWarmStartInit for a generic JoinModel requires all component manifolds to be Manifolds.Segre with identical factor_dims.", @@ -407,19 +471,21 @@ function egrad(model::JoinModel{<:AbstractFloat,<:JoinBackend}, p) backend = model.backend residual = _join_residual_grad!(backend, p) parts = point_parts(p) - vals = ntuple(k -> _manifold_egrad(backend.manifolds[k], parts[k], residual), backend.r) + vals = + ntuple(k -> _component_egrad(backend.components[k], parts[k], residual), backend.r) return wrap_like_point(p, vals) end -function _join_basis_project(manifolds::Tuple, p, residual) +function _join_basis_project(components::Tuple, p, residual) parts = point_parts(p) - _check_parts_len(parts, length(manifolds), "_join_basis_project") + _check_parts_len(parts, length(components), "_join_basis_project") vals = ntuple(k -> begin - Mk = manifolds[k] + ck = components[k] + Mk = _component_manifold(ck) pk = parts[k] - eg = _manifold_egrad(Mk, pk, residual) + eg = _component_egrad(ck, pk, residual) egrad_to_rgrad(Mk, pk, eg) - end, length(manifolds)) + end, length(components)) return wrap_like_point(p, vals) end @@ -437,7 +503,7 @@ model_exact_join_basis_function(model::JoinModel{<:AbstractFloat,<:JoinBackend}) (M, p) -> begin backend = model.backend residual = _join_residual!(backend, p) - _join_basis_project(backend.manifolds, p, residual) + _join_basis_project(backend.components, p, residual) end function extract_components( @@ -445,6 +511,7 @@ function extract_components( p, ) backend = model.backend + components = _backend_components(backend) parts = point_parts(p) _check_parts_len(parts, backend.r, "extract_components") T = eltype(backend.target) @@ -453,14 +520,14 @@ function extract_components( manifold_type = Union{} @inbounds for k = 1:backend.r point_type = typejoin(point_type, typeof(parts[k])) - manifold_type = typejoin(manifold_type, typeof(backend.manifolds[k])) + manifold_type = typejoin(manifold_type, typeof(_component_manifold(components[k]))) end comps = Vector{DecompositionComponent{T,N,point_type,manifold_type}}(undef, backend.r) @inbounds for k = 1:backend.r # Components keep only point/manifold metadata and reconstruct derived tensors on demand. comps[k] = DecompositionComponent{T,N,point_type,manifold_type}( parts[k], - backend.manifolds[k], + _component_manifold(components[k]), backend.target_shape, ) end @@ -473,28 +540,34 @@ function rgrad(model::JoinModel{<:AbstractFloat,<:JoinBackend}, p) _check_parts_len(parts, backend.r, "rgrad") residual = _join_residual_grad!(backend, p) vals = ntuple(k -> begin - Mk = backend.manifolds[k] + ck = backend.components[k] + Mk = _component_manifold(ck) pk = parts[k] - eg = _manifold_egrad(Mk, pk, residual) + eg = _component_egrad(ck, pk, residual) egrad_to_rgrad(Mk, pk, eg) end, backend.r) return wrap_like_point(p, vals) end """ - _component_ambient_embedding!(out, M, p) + _component_ambient_embedding!(out, component, p) Write the ambient embedding of one join component into `out`. Component manifolds may specialize this hook when their native point structure admits a more direct embedding than the generic `embed!` path. """ function _component_ambient_embedding!(out::AbstractVector, M, p) + return _component_ambient_embedding!(out, DefaultJoinEmbedding(), M, p) +end + +function _component_ambient_embedding!(out::AbstractVector, ::DefaultJoinEmbedding, M, p) ManifoldsBase.embed!(M, out, p) return out end function _component_ambient_embedding!( out::AbstractVector, + ::DefaultJoinEmbedding, M::Manifolds.Tucker, p::Manifolds.TuckerPoint, ) @@ -505,11 +578,29 @@ function _component_ambient_embedding!( end function _component_ambient_embedding!(out::AbstractVector, M::Manifolds.Segre, p) + return _component_ambient_embedding!(out, DefaultJoinEmbedding(), M, p) +end + +function _component_ambient_embedding!( + out::AbstractVector, + ::DefaultJoinEmbedding, + M::Manifolds.Segre, + p, +) copyto!(out, _segre_component_tensorvec(p)) return out end function _component_ambient_embedding!(out::AbstractVector, M::Manifolds.Tucker, p) + return _component_ambient_embedding!(out, DefaultJoinEmbedding(), M, p) +end + +function _component_ambient_embedding!( + out::AbstractVector, + ::DefaultJoinEmbedding, + M::Manifolds.Tucker, + p, +) throw( ArgumentError( "Expected native TuckerPoint for Manifolds.Tucker, got $(typeof(p)).", @@ -517,14 +608,28 @@ function _component_ambient_embedding!(out::AbstractVector, M::Manifolds.Tucker, ) end +function _component_ambient_embedding!(out::AbstractVector, component::JoinComponent, p) + return _component_ambient_embedding!(out, component.embedding, component.manifold, p) +end + """ - _component_ambient_pushforward!(out, M, p, X) + _component_ambient_pushforward!(out, component, p, X) Write the ambient pushforward `DΦ(p)[X]` of one join component into `out`. This is the component-level differential used by LM Jacobian assembly on generic `JoinModel`s. """ function _component_ambient_pushforward!(out::AbstractVector, M, p, X) + return _component_ambient_pushforward!(out, DefaultJoinEmbedding(), M, p, X) +end + +function _component_ambient_pushforward!( + out::AbstractVector, + ::DefaultJoinEmbedding, + M, + p, + X, +) emb = ManifoldsBase.embed(M, p, X) length(emb) == length(out) || throw( DimensionMismatch( @@ -536,13 +641,38 @@ function _component_ambient_pushforward!(out::AbstractVector, M, p, X) end function _component_ambient_pushforward!(out::AbstractVector, M::Manifolds.Segre, p, X) + return _component_ambient_pushforward!(out, DefaultJoinEmbedding(), M, p, X) +end + +function _component_ambient_pushforward!( + out::AbstractVector, + ::DefaultJoinEmbedding, + M::Manifolds.Segre, + p, + X, +) copyto!(out, _segre_tangent_tensorvec(p, X)) return out end +function _component_ambient_pushforward!( + out::AbstractVector, + component::JoinComponent, + p, + X, +) + return _component_ambient_pushforward!( + out, + component.embedding, + component.manifold, + p, + X, + ) +end + function _subtract_ambient_tensor!( residual::AbstractArray{T,N}, - M, + component, p, work_vec::AbstractVector{T}, ) where {T<:AbstractFloat,N} @@ -551,7 +681,7 @@ function _subtract_ambient_tensor!( "_subtract_ambient_tensor!: work length $(length(work_vec)) != residual length $(length(residual))", ), ) - _component_ambient_embedding!(work_vec, M, p) + _component_ambient_embedding!(work_vec, component, p) residual_vec = vec(residual) @inbounds for i in eachindex(residual_vec, work_vec) residual_vec[i] -= work_vec[i] @@ -588,7 +718,7 @@ then accumulates those buffers into `out`. This avoids allocating one dense ambient tensor per component during solver iterations. """ function _join_reconstruct!(out::AbstractArray, backend::Union{JoinBackend,BTDBackend}, p) - manifolds = backend.manifolds + components = _backend_components(backend) r = backend.r bufs = backend.component_bufs @@ -599,7 +729,7 @@ function _join_reconstruct!(out::AbstractArray, backend::Union{JoinBackend,BTDBa @inbounds for k = 1:r # Reconstruct each component into its preallocated workspace. - _component_ambient_embedding!(bufs[k], manifolds[k], parts[k]) + _component_ambient_embedding!(bufs[k], components[k], parts[k]) # Accumulate into the output tensor without allocating a Khatri-Rao-sized object. out .+= bufs[k] diff --git a/src/join/join_model.jl b/src/join/join_model.jl index c6218a1..1b92e2c 100644 --- a/src/join/join_model.jl +++ b/src/join/join_model.jl @@ -1,10 +1,22 @@ # join/join_model.jl — Join-front-end model and backend type definitions -export AbstractJoinBackend, JoinModel, CPDBackend, JoinBackend, BTDBackend +export AbstractJoinBackend, JoinComponent, JoinModel, CPDBackend, JoinBackend, BTDBackend # BTDBackend is defined in `btd/model.jl` (includes contraction workspace). abstract type AbstractJoinBackend end +struct JoinComponent{M,E} + manifold::M + embedding::E +end + +struct DefaultJoinEmbedding end + +JoinComponent(manifold::M) where {M} = + JoinComponent{M,DefaultJoinEmbedding}(manifold, DefaultJoinEmbedding()) + +manifold(component::JoinComponent) = component.manifold + struct JoinModel{T<:AbstractFloat,B<:AbstractJoinBackend} <: AbstractDecompositionModel{T} backend::B end @@ -24,7 +36,7 @@ _JoinResidualWORO(residual::V) where {V} = struct JoinBackend{ T, N, - MT<:Tuple, + CT<:Tuple, A<:AbstractArray{T,N}, V, MP<:ProductManifold, @@ -32,7 +44,7 @@ struct JoinBackend{ W<:_JoinResidualWORO{T}, C, } <: AbstractJoinBackend - manifolds::MT + components::CT r::Int target::A target_shape::NTuple{N,Int} diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 77391f4..d0f6d34 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -47,10 +47,11 @@ function _join_tangent_ambient_vector!( _check_parts_len(parts, backend.r, "_join_tangent_ambient_vector!") _check_parts_len(xparts, backend.r, "_join_tangent_ambient_vector!") fill!(out, zero(eltype(out))) + components = _backend_components(backend) @inbounds for k = 1:backend.r _component_ambient_pushforward!( backend.component_bufs[k], - backend.manifolds[k], + components[k], parts[k], xparts[k], ) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 0f5ab1f..ab7913b 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -186,7 +186,8 @@ end @test model.backend isa TensorKitchen.JoinBackend @test model.backend.r == 3 @test tensor(model) == A - @test model.backend.manifolds == segres + @test length(model.backend.components) == 3 + @test map(TensorKitchen.manifold, model.backend.components) == segres M = TensorKitchen.manifold(model) @test M isa ProductManifold @@ -195,7 +196,12 @@ end model_repeat = JoinModel(Manifolds.Segre(dims), 3, A) @test model_repeat.backend.r == 3 - @test model_repeat.backend.manifolds == segres + @test map(TensorKitchen.manifold, model_repeat.backend.components) == segres + + model_component = + JoinModel((TensorKitchen.JoinComponent(segres[1]), segres[2], segres[3]), A) + @test model_component.backend.components[1] isa TensorKitchen.JoinComponent + @test map(TensorKitchen.manifold, model_component.backend.components) == segres p = TensorKitchen.initial_point(model, :random; verbose = false) @test length(TensorKitchen.point_parts(p)) == 3 From 58000f2cf50ccb2b2ad0c28b36caa856ce4b555b Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 14:15:28 +0200 Subject: [PATCH 27/37] refactor JoinComponent Abstraction --- src/btd/model.jl | 38 +++-- src/cpd/model/parameterizations.jl | 251 +++++++++++++++++++++++++++++ src/cpd/model/rank1.jl | 57 ++----- src/cpd/model/rankr.jl | 67 ++------ src/decompositions.jl | 1 + src/join/cpd_backend.jl | 83 +++------- src/join/join_backend.jl | 14 +- src/solvers/btd_als.jl | 2 +- src/solvers/btd_tsd.jl | 4 +- src/solvers/lm.jl | 100 +----------- test/basic_tests.jl | 171 +++++++++++++++++++- 11 files changed, 508 insertions(+), 280 deletions(-) create mode 100644 src/cpd/model/parameterizations.jl diff --git a/src/btd/model.jl b/src/btd/model.jl index 509905d..400ef95 100644 --- a/src/btd/model.jl +++ b/src/btd/model.jl @@ -31,8 +31,18 @@ Backend state for block-term decomposition as a sum of Tucker blocks. Stores the target tensor, product manifold, reusable work buffers, and per-block ambient reconstruction buffers used by cost, gradient, and ALS routines. """ -struct BTDBackend{T,N,MT<:Tuple,A<:AbstractArray{T,N},V,MP<:ProductManifold,I,C} <: - AbstractJoinBackend +struct BTDBackend{ + T, + N, + CT<:Tuple, + MT<:Tuple, + A<:AbstractArray{T,N}, + V, + MP<:ProductManifold, + I, + C, +} <: AbstractJoinBackend + components::CT manifolds::MT r::Int # Preserve the target array/backend so BTD shares the generic join storage behavior. @@ -60,7 +70,7 @@ model_exact_join_basis_function(model::JoinModel{<:AbstractFloat,<:BTDBackend}) (M, p) -> begin backend = model.backend residual = _join_residual!(backend, p) - _join_basis_project(backend.manifolds, p, residual) + _join_basis_project(backend.components, p, residual) end @@ -71,14 +81,11 @@ function _btd_sequential_tucker_init(model::JoinModel{<:AbstractFloat,<:BTDBacke residual = copy(backend.target) parts = Vector{Manifolds.TuckerPoint{eltype(backend.target)}}(undef, backend.r) for k = 1:backend.r - pk = _manifold_init(backend.manifolds[k], residual, init) + component = _backend_component(backend, k) + Mk = _backend_manifold(backend, k) + pk = _component_init(component, residual, init) parts[k] = pk - _subtract_ambient_tensor!( - residual, - backend.manifolds[k], - pk, - backend.component_bufs[k], - ) + _subtract_ambient_tensor!(residual, Mk, pk, backend.component_bufs[k]) end return ArrayPartition(parts...) end @@ -86,7 +93,7 @@ end function _btd_block_ranks_by_mode(backend::BTDBackend{T,N}) where {T,N} ranks = Vector{NTuple{N,Int}}(undef, backend.r) for b = 1:backend.r - M = backend.manifolds[b] + M = _backend_manifold(backend, b) M isa Manifolds.Tucker || throw( ArgumentError( "BTD HOSVD multistart expects Tucker manifolds, got $(typeof(M)) at block $b.", @@ -160,7 +167,7 @@ function _btd_hosvd_split_candidate( parts[b] = pk _subtract_ambient_tensor!( residual, - backend.manifolds[b], + _backend_manifold(backend, b), pk, backend.component_bufs[b], ) @@ -190,7 +197,7 @@ function initial_point( init_sym = _builtin_initializer_symbol(init) if init_sym == :random parts = ntuple( - k -> _manifold_init(backend.manifolds[k], backend.target, init), + k -> _component_init(_backend_component(backend, k), backend.target, init), backend.r, ) return ArrayPartition(parts...) @@ -282,6 +289,9 @@ function rgrad(model::JoinModel{<:AbstractFloat,<:BTDBackend}, p) parts = point_parts(p) _check_parts_len(parts, backend.r, "BTD rgrad") eg = point_parts(_btd_egrad(backend, p)) - vals = ntuple(k -> egrad_to_rgrad(backend.manifolds[k], parts[k], eg[k]), backend.r) + vals = ntuple( + k -> egrad_to_rgrad(_backend_manifold(backend, k), parts[k], eg[k]), + backend.r, + ) return wrap_like_point(p, vals) end diff --git a/src/cpd/model/parameterizations.jl b/src/cpd/model/parameterizations.jl new file mode 100644 index 0000000..6d2f35f --- /dev/null +++ b/src/cpd/model/parameterizations.jl @@ -0,0 +1,251 @@ +# cpd/model/parameterizations.jl — CP point parameterization helpers + +abstract type AbstractCPParameterization end + +struct NativeCPEmbedding <: AbstractCPParameterization end +struct CanonicalCPEmbedding <: AbstractCPParameterization end +struct SquaredNonnegativeCPEmbedding <: AbstractCPParameterization end +struct SoftplusNonnegativeCPEmbedding <: AbstractCPParameterization end + +@inline _cp_softplus_encode_value(x::T) where {T<:AbstractFloat} = + _invsoftplus(max(x, eps(T))) + +function _cp_rank1_decode_factors(::NativeCPEmbedding, dims, p) + return unpack_point_rank1(p, dims) +end + +function _cp_rank1_decode_factors(::SquaredNonnegativeCPEmbedding, dims, p) + λ̃, Ũ = unpack_point_rank1(p, dims) + return λ̃^2, [Ũ[m] .^ 2 for m in eachindex(Ũ)] +end + +function _cp_rank1_decode_factors(::SoftplusNonnegativeCPEmbedding, dims, p) + λ̃, Ũ = unpack_point_rank1(p, dims) + return _softplus_value(λ̃), [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] +end + +function _cp_rank1_encode_point(::NativeCPEmbedding, λ, U) + return pack_point_rank1_segre(λ, U) +end + +function _cp_rank1_encode_point(::SquaredNonnegativeCPEmbedding, λ, U) + return pack_point_rank1( + sqrt(max(λ, zero(λ))), + [sqrt.(max.(u, zero(eltype(u)))) for u in U], + ) +end + +function _cp_rank1_encode_point(::SoftplusNonnegativeCPEmbedding, λ, U) + T = typeof(λ) + return pack_point_rank1( + _cp_softplus_encode_value(max(λ, zero(T))), + [_cp_softplus_encode_value.(max.(u, zero(eltype(u)))) for u in U], + ) +end + +function _cp_rank1_seed_point(::NativeCPEmbedding, λ, U) + return pack_point_rank1_segre(λ, U) +end + +function _cp_rank1_seed_point(::SquaredNonnegativeCPEmbedding, λ, U) + T = typeof(λ) + return pack_point_rank1( + sqrt(max(abs(λ), eps(T))), + [sqrt.(max.(abs.(u), eps(eltype(u)))) for u in U], + ) +end + +function _cp_rank1_seed_point(::SoftplusNonnegativeCPEmbedding, λ, U) + T = typeof(λ) + return pack_point_rank1( + _invsoftplus(max(abs(λ), eps(T))), + [_invsoftplus.(max.(abs.(u), eps(eltype(u)))) for u in U], + ) +end + +function _cp_rank1_embed_tensor(embedding::AbstractCPParameterization, dims, p) + λ, U = _cp_rank1_decode_factors(embedding, dims, p) + return reconstruct_cp_rank1(λ, U) +end + +function _cp_rank1_tangent_tensorvec!( + out::AbstractVector{T}, + λ::T, + U::AbstractVector{<:AbstractVector{T}}, + λ̇::T, + U̇::AbstractVector{<:AbstractVector{T}}, +) where {T<:AbstractFloat} + comp = ([λ], U...) + xcomp = ([λ̇], U̇...) + copyto!(out, _segre_tangent_tensorvec(comp, xcomp)) + return out +end + +function _cp_rank1_decode_tangent_factors(::NativeCPEmbedding, dims, p, X) + λ, U = unpack_point_rank1(p, dims) + λ̇, U̇ = unpack_point_rank1(X, dims) + return λ, U, λ̇, U̇ +end + +function _cp_rank1_decode_tangent_factors(::SquaredNonnegativeCPEmbedding, dims, p, X) + λ̃, Ũ = unpack_point_rank1(p, dims) + λ̇̃, U̇̃ = unpack_point_rank1(X, dims) + λ = λ̃^2 + U = [Ũ[m] .^ 2 for m in eachindex(Ũ)] + λ̇ = 2 * λ̃ * λ̇̃ + U̇ = [2 .* Ũ[m] .* U̇̃[m] for m in eachindex(Ũ)] + return λ, U, λ̇, U̇ +end + +function _cp_rank1_decode_tangent_factors(::SoftplusNonnegativeCPEmbedding, dims, p, X) + λ̃, Ũ = unpack_point_rank1(p, dims) + λ̇̃, U̇̃ = unpack_point_rank1(X, dims) + λ = _softplus_value(λ̃) + U = [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] + λ̇ = _softplus_derivative(λ̃) * λ̇̃ + U̇ = [_softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(Ũ)] + return λ, U, λ̇, U̇ +end + +function _cp_rankr_decode_factors(::NativeCPEmbedding, dims, r, p) + return unpack_rankr_native(p, dims, r) +end + +function _cp_rankr_decode_factors(::CanonicalCPEmbedding, dims, r, p) + return unpack_rankr_canonical(p, dims, r) +end + +function _cp_rankr_decode_factors(::SquaredNonnegativeCPEmbedding, dims, r, p) + λ̃, Ũ = unpack_point_rankr(p, dims, r) + return λ̃ .^ 2, [Ũ[m] .^ 2 for m in eachindex(Ũ)] +end + +function _cp_rankr_decode_factors(::SoftplusNonnegativeCPEmbedding, dims, r, p) + λ̃, Ũ = unpack_point_rankr(p, dims, r) + return _softplus_value.(λ̃), [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] +end + +function _cp_rankr_encode_point(::NativeCPEmbedding, λ, U, r) + return pack_rankr_native(λ, U, r) +end + +function _cp_rankr_encode_point(::CanonicalCPEmbedding, λ, U, r) + return pack_rankr_canonical(λ, U, r) +end + +function _cp_rankr_encode_point(::SquaredNonnegativeCPEmbedding, λ, U, r) + T = eltype(λ) + return pack_point_rankr( + sqrt.(max.(λ, zero(T))), + [sqrt.(max.(F, zero(eltype(F)))) for F in U], + r, + ) +end + +function _cp_rankr_encode_point(::SoftplusNonnegativeCPEmbedding, λ, U, r) + T = eltype(λ) + return pack_point_rankr( + _cp_softplus_encode_value.(max.(λ, zero(T))), + [_cp_softplus_encode_value.(max.(F, zero(eltype(F)))) for F in U], + r, + ) +end + +function _cp_rankr_seed_point(::NativeCPEmbedding, λ, U, r) + return pack_rankr_native(λ, U, r) +end + +function _cp_rankr_seed_point(::CanonicalCPEmbedding, λ, U, r) + return pack_rankr_canonical(λ, U, r) +end + +function _cp_rankr_seed_point(::SquaredNonnegativeCPEmbedding, λ, U, r) + T = eltype(λ) + return pack_point_rankr( + sqrt.(max.(abs.(λ), eps(T))), + [sqrt.(max.(abs.(F), eps(eltype(F)))) for F in U], + r, + ) +end + +function _cp_rankr_seed_point(::SoftplusNonnegativeCPEmbedding, λ, U, r) + T = eltype(λ) + return pack_point_rankr( + _invsoftplus.(max.(abs.(λ), eps(T))), + [_invsoftplus.(max.(abs.(F), eps(eltype(F)))) for F in U], + r, + ) +end + +function _cp_rankr_embed_tensor(embedding::AbstractCPParameterization, dims, r, p) + λ, U = _cp_rankr_decode_factors(embedding, dims, r, p) + return reconstruct_cpd_rankr(λ, U) +end + +function _cp_rankr_tangent_tensorvec!( + out::AbstractVector{T}, + λ::AbstractVector{T}, + U::Vector{<:AbstractMatrix{T}}, + λ̇::AbstractVector{T}, + U̇::Vector{<:AbstractMatrix{T}}, +) where {T<:AbstractFloat} + fill!(out, zero(T)) + r = length(λ) + d = length(U) + Ucols = Vector{AbstractVector{T}}(undef, d) + U̇cols = Vector{AbstractVector{T}}(undef, d) + @inbounds for k = 1:r + for m = 1:d + Ucols[m] = @view U[m][:, k] + U̇cols[m] = @view U̇[m][:, k] + end + out .+= _segre_tangent_tensorvec(([λ[k]], Ucols...), ([λ̇[k]], U̇cols...)) + end + return out +end + +function _cp_rankr_decode_tangent_factors(::NativeCPEmbedding, dims, r, p, X) + λ, U = unpack_rankr_native(p, dims, r) + xparts = parts_tuple(X) + length(xparts) == r || throw( + DimensionMismatch("expected $r native tangent components, got $(length(xparts))"), + ) + T = eltype(λ) + d = length(dims) + λ̇ = similar(λ) + U̇ = [zeros(T, dims[m], r) for m = 1:d] + @inbounds for k = 1:r + xk_λ, xk_U = unpack_point_rank1(xparts[k], dims) + λ̇[k] = xk_λ + for m = 1:d + U̇[m][:, k] .= xk_U[m] + end + end + return λ, U, λ̇, U̇ +end + +function _cp_rankr_decode_tangent_factors(::CanonicalCPEmbedding, dims, r, p, X) + λ, U = unpack_rankr_canonical(p, dims, r) + λ̇, U̇ = unpack_rankr_canonical(X, dims, r) + return λ, U, λ̇, U̇ +end + +function _cp_rankr_decode_tangent_factors(::SquaredNonnegativeCPEmbedding, dims, r, p, X) + λ̃, Ũ = unpack_point_rankr(p, dims, r) + λ̇̃, U̇̃ = unpack_point_rankr(X, dims, r) + λ = λ̃ .^ 2 + U = [Ũ[m] .^ 2 for m in eachindex(Ũ)] + λ̇ = 2 .* λ̃ .* λ̇̃ + U̇ = [2 .* Ũ[m] .* U̇̃[m] for m in eachindex(Ũ)] + return λ, U, λ̇, U̇ +end + +function _cp_rankr_decode_tangent_factors(::SoftplusNonnegativeCPEmbedding, dims, r, p, X) + λ̃, Ũ = unpack_point_rankr(p, dims, r) + λ̇̃, U̇̃ = unpack_point_rankr(X, dims, r) + λ = _softplus_value.(λ̃) + U = [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] + λ̇ = _softplus_derivative.(λ̃) .* λ̇̃ + U̇ = [_softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(Ũ)] + return λ, U, λ̇, U̇ +end diff --git a/src/cpd/model/rank1.jl b/src/cpd/model/rank1.jl index 8903818..b27e55c 100644 --- a/src/cpd/model/rank1.jl +++ b/src/cpd/model/rank1.jl @@ -58,22 +58,15 @@ end tensor(model::Rank1CPDModel) = model.A manifold(model::Rank1CPDModel) = model.M +_cp_parameterization(model::Rank1CPDModel) = + model.nonnegative ? + ( + _rank1_uses_softplus_metric(model.M) ? SoftplusNonnegativeCPEmbedding() : + SquaredNonnegativeCPEmbedding() + ) : NativeCPEmbedding() + function embed_point(model::Rank1CPDModel{T,N}, p) where {T,N} - if model.nonnegative - if _rank1_uses_softplus_metric(model.M) - λ̃, Ũ = unpack_point_rank1(p, model.dims) - λ = _softplus_value(λ̃) - U = [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] - return reconstruct_cp_rank1(λ, U) - end - return embed_point_rank1_nn(p, model.dims) - end - M = manifold(model) - M isa Manifolds.Segre || return embed_point_rank1(p, model.dims) - return reshape( - ManifoldsBase.embed!(M, Vector{eltype(p[1])}(undef, prod(model.dims)), p), - model.dims, - ) + return _cp_rank1_embed_tensor(_cp_parameterization(model), model.dims, p) end function cost(model::Rank1CPDModel{T,N}, p) where {T,N} @@ -166,16 +159,7 @@ function initial_point( init == :alswarm && return initial_point(model, ALSWarmStartInit(); verbose) init_sym = _builtin_initializer_symbol(init) U0 = init_cp_rank1(model.A; init = init_sym) - if model.nonnegative - λ̃0 = _rank1_uses_softplus_metric(model.M) ? _invsoftplus(one(T)) : one(T) - Ũ = if _rank1_uses_softplus_metric(model.M) - [_invsoftplus.(max.(abs.(u), eps(T))) for u in U0] - else - [sqrt.(max.(abs.(u), eps(T))) for u in U0] - end - return pack_point_rank1(λ̃0, Ũ) - end - return pack_point_rank1_segre(one(T), U0) + return _cp_rank1_seed_point(_cp_parameterization(model), one(T), U0) end initial_point(model::Rank1CPDModel, init::PointInit; kwargs...) = init.point @@ -196,17 +180,7 @@ nonnegative squared parameterization so backend postprocessing can work with a uniform `lambda + factors` representation. """ function cpd_point(model::Rank1CPDModel{T,N}, p) where {T<:AbstractFloat,N} - if model.nonnegative - λ̃, Ũ = unpack_point_rank1(p, model.dims) - if _rank1_uses_softplus_metric(model.M) - return CPDPoint( - T[_softplus_value(λ̃)], - [reshape(_softplus_value.(Ũ[m]), :, 1) for m in eachindex(Ũ)], - ) - end - return CPDPoint(T[λ̃^2], [reshape(Ũ[m] .^ 2, :, 1) for m in eachindex(Ũ)]) - end - λ, U = unpack_point_rank1(p, model.dims) + λ, U = _cp_rank1_decode_factors(_cp_parameterization(model), model.dims, p) return CPDPoint(T[λ], [reshape(U[m], :, 1) for m in eachindex(U)]) end @@ -221,16 +195,7 @@ function pack_cpd_point( ) λ = lambda(point)[1] U = [Vector(@view F[:, 1]) for F in factors(point)] - if model.nonnegative - if _rank1_uses_softplus_metric(model.M) - return pack_point_rank1( - _invsoftplus(max(λ, zero(T))), - [_invsoftplus.(max.(u, zero(T))) for u in U], - ) - end - return pack_point_rank1(sqrt(max(λ, zero(T))), [sqrt.(max.(u, zero(T))) for u in U]) - end - return pack_point_rank1_segre(λ, U) + return _cp_rank1_encode_point(_cp_parameterization(model), λ, U) end function post_step!( diff --git a/src/cpd/model/rankr.jl b/src/cpd/model/rankr.jl index 8df2904..b9eded1 100644 --- a/src/cpd/model/rankr.jl +++ b/src/cpd/model/rankr.jl @@ -100,20 +100,15 @@ end tensor(model::RankRCPDModel) = model.A manifold(model::RankRCPDModel) = model.M +_cp_parameterization(model::RankRCPDModel) = + model.nonnegative ? + ( + model.geometry == :softplus_metric ? SoftplusNonnegativeCPEmbedding() : + SquaredNonnegativeCPEmbedding() + ) : (model.geometry == :native ? NativeCPEmbedding() : CanonicalCPEmbedding()) + function embed_point(model::RankRCPDModel{T,N}, p) where {T,N} - if model.nonnegative - λ̃, Ũ = unpack_point_rankr(p, model.dims, model.r) - if model.geometry == :softplus_metric - λ = _softplus_value.(λ̃) - U = [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] - return reconstruct_cpd_rankr(λ, U) - end - return embed_point_rankr_nn(p, model.dims, model.r) - end - λ, U = - model.geometry == :native ? unpack_rankr_native(p, model.dims, model.r) : - unpack_rankr_canonical(p, model.dims, model.r) - return reconstruct_cpd_rankr(λ, U) + return _cp_rankr_embed_tensor(_cp_parameterization(model), model.dims, model.r, p) end function cost(model::RankRCPDModel{T,N}, p) where {T,N} @@ -422,18 +417,7 @@ function initial_point( init == :alswarm && return initial_point(model, ALSWarmStartInit(); verbose) init_sym = _builtin_initializer_symbol(init) λ0, U0 = init_cpd_factors(model.A, model.r; init = init_sym) # initialize the factors - if model.nonnegative - if model.geometry == :softplus_metric - λ̃0 = _invsoftplus.(max.(abs.(λ0), eps(T))) - Ũ0 = [_invsoftplus.(max.(abs.(U0[m]), eps(T))) for m in eachindex(U0)] - else - λ̃0 = sqrt.(max.(abs.(λ0), eps(T))) - Ũ0 = [sqrt.(max.(abs.(U0[m]), eps(T))) for m in eachindex(U0)] - end - return pack_point_rankr(λ̃0, Ũ0, model.r) # structured join layout - end - return model.geometry == :native ? pack_rankr_native(λ0, U0, model.r) : - pack_rankr_canonical(λ0, U0, model.r) # native: ArrayPartition, canonical: tuple + return _cp_rankr_seed_point(_cp_parameterization(model), λ0, U0, model.r) end initial_point(model::RankRCPDModel, init::PointInit; kwargs...) = init.point @@ -445,19 +429,7 @@ supports_normalization_policy(model::RankRCPDModel, policy::AbstractNormalizatio policy isa Union{NoNormalization,SeparateLambdaNormalization} function cpd_point(model::RankRCPDModel{T,N}, p) where {T<:AbstractFloat,N} - if model.nonnegative - λ̃, Ũ = unpack_point_rankr(p, model.dims, model.r) - if model.geometry == :softplus_metric - return CPDPoint( - _softplus_value.(λ̃), - [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)], - ) - end - return CPDPoint(λ̃ .^ 2, [Ũ[m] .^ 2 for m in eachindex(Ũ)]) - end - λ, U = - model.geometry == :native ? unpack_rankr_native(p, model.dims, model.r) : - unpack_rankr_canonical(p, model.dims, model.r) + λ, U = _cp_rankr_decode_factors(_cp_parameterization(model), model.dims, model.r, p) return CPDPoint(λ, U) end @@ -470,19 +442,12 @@ function pack_cpd_point( "CPDPoint has $(length(lambda(point))) weights, expected rank $(model.r)", ), ) - if model.nonnegative - if model.geometry == :softplus_metric - λ̃ = _invsoftplus.(max.(lambda(point), zero(T))) - Ũ = [_invsoftplus.(max.(F, zero(T))) for F in factors(point)] - else - λ̃ = sqrt.(max.(lambda(point), zero(T))) - Ũ = [sqrt.(max.(F, zero(T))) for F in factors(point)] - end - return pack_point_rankr(λ̃, Ũ, model.r) - end - return model.geometry == :native ? - pack_rankr_native(lambda(point), factors(point), model.r) : - pack_rankr_canonical(lambda(point), factors(point), model.r) + return _cp_rankr_encode_point( + _cp_parameterization(model), + lambda(point), + factors(point), + model.r, + ) end function post_step!( diff --git a/src/decompositions.jl b/src/decompositions.jl index 7e8c143..1fd2713 100644 --- a/src/decompositions.jl +++ b/src/decompositions.jl @@ -8,6 +8,7 @@ include("cpd/core/mttkrp.jl") include("cpd/core/cpd_init.jl") include("cpd/core/cp_cost.jl") include("cpd/core/reconstruct.jl") +include("cpd/model/parameterizations.jl") include("cpd/model/rank1.jl") include("cpd/model/rankr.jl") diff --git a/src/join/cpd_backend.jl b/src/join/cpd_backend.jl index 123a6b5..b062160 100644 --- a/src/join/cpd_backend.jl +++ b/src/join/cpd_backend.jl @@ -96,60 +96,34 @@ function _public_cpd_factors(m, λ, U) return λ, U end -@inline function _decode_nonnegative_cpd(m, λ̃, Ũ) - use_softplus = - hasproperty(m, :geometry) ? (m.geometry == :softplus_metric) : - (hasproperty(m, :M) && _rank1_uses_softplus_metric(m.M)) - if use_softplus - return _softplus_value.(λ̃), [_softplus_value.(Ũ[j]) for j in eachindex(Ũ)] +@inline _cpd_solver_point(m, p, solver_sym::Symbol) = cpd_point(m, p) + +function _cpd_solver_point( + m::RankRCPDModel{T}, + p, + solver_sym::Symbol, +) where {T<:AbstractFloat} + if m.nonnegative && (solver_sym in _CP_ALS_FAMILY_SOLVERS) + λ, U = unpack_rankr_canonical(p, m.dims, m.r) + return CPDPoint(λ, U) end - return λ̃ .^ 2, [Ũ[j] .^ 2 for j in eachindex(Ũ)] + return cpd_point(m, p) end function _cpd_result(model::JoinModel{<:AbstractFloat,<:CPDBackend}, result, dims, r) m = cpd_model(model) solver_sym = _result_solver_symbol(solver(result)) si = _result_solver_info(result) - als_family = solver_sym in _CP_ALS_FAMILY_SOLVERS - - if r == 1 - λ̃, U_vec = unpack_point_rank1(point(result), dims) - if m.nonnegative && !als_family - λ_vec, U_sq = _decode_nonnegative_cpd(m, [λ̃], U_vec) - λ = λ_vec[1] - else - λ = λ̃ - U_sq = U_vec - end - λ_pub, U_pub = _public_cpd_factors(m, [λ], [reshape(u, :, 1) for u in U_sq]) - return CPDResult( - λ_pub, - U_pub, - cost(result), - rel_error(result), - grad_norm(result), - iterations(result), - converged(result), - solver_sym, - si, - ) - end - - comps = if m.nonnegative && !als_family - λ̃, Ũ = unpack_point_rankr(point(result), dims, r) - λ, U = _decode_nonnegative_cpd(m, λ̃, Ũ) - components_from_factors(λ, U) - else - unpack_point_rankr_components(point(result), dims, r) - end + q = _cpd_solver_point(m, point(result), solver_sym) + λ_raw = lambda(q) + U_raw = factors(q) rel_err = rel_error(result) if !isfinite(rel_err) - Xhat = reconstruct_cpd_rankr([c.λ for c in comps], factors_from_components(comps)) + Xhat = reconstruct_cpd_rankr(λ_raw, U_raw) rel_err = rel_error(m.A, Xhat) end - λ_pub, U_pub = - _public_cpd_factors(m, [c.λ for c in comps], factors_from_components(comps)) + λ_pub, U_pub = _public_cpd_factors(m, λ_raw, U_raw) return CPDResult( λ_pub, U_pub, @@ -167,31 +141,12 @@ function extract_components(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) return _extract_cpd_components(cpd_model(model), p) end -function _extract_cpd_components(m::RankRCPDModel, p) - comps = - m.nonnegative ? begin - λ̃, Ũ = unpack_point_rankr(p, m.dims, m.r) - λ, U = _decode_nonnegative_cpd(m, λ̃, Ũ) - components_from_factors(λ, U) - end : unpack_point_rankr_components(p, m.dims, m.r) +function _extract_cpd_components(m::Union{Rank1CPDModel,RankRCPDModel}, p) + q = cpd_point(m, p) + comps = components_from_factors(lambda(q), factors(q)) return [CPDComponent(pack_point_rank1(c.λ, c.vectors), c) for c in comps] end -function _extract_cpd_components(m::Rank1CPDModel, p) - λ, U = unpack_point_rank1(p, m.dims) - if m.nonnegative - if _rank1_uses_softplus_metric(m.M) - λ = _softplus_value(λ) - U = [_softplus_value.(u) for u in U] - else - λ = λ^2 - U = [u .^ 2 for u in U] - end - end - c = RankOneTensor(λ, U) - return [CPDComponent(pack_point_rank1(λ, U), c)] -end - function _extract_cpd_components(m, p) throw( ArgumentError( diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 5e828b4..1cb093d 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -11,7 +11,9 @@ _wrap_join_component(manifold::AbstractManifold) = JoinComponent(manifold) _component_manifold(component::JoinComponent) = component.manifold _component_manifold(manifold::AbstractManifold) = manifold _backend_components(backend::JoinBackend) = backend.components -_backend_components(backend::BTDBackend) = backend.manifolds +_backend_components(backend::BTDBackend) = backend.components +_backend_component(backend, k::Int) = _backend_components(backend)[k] +_backend_manifold(backend, k::Int) = _component_manifold(_backend_component(backend, k)) function _as_join_manifold_tuple(manifolds::Tuple) all(_is_manifold_like, manifolds) || throw( @@ -251,10 +253,11 @@ function _sum_backend_parts( target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - r = length(components) + components_tuple = _as_join_component_tuple(components) + r = length(components_tuple) # Keep the original target representation instead of eagerly materializing Array. tgt = target - _validate_join_ambient_compatibility(components, tgt) + _validate_join_ambient_compatibility(components_tuple, tgt) tflat = vec(tgt) tgt_len = length(tgt) @@ -262,10 +265,10 @@ function _sum_backend_parts( component_bufs = [_join_vector_workspace_like(tgt, tgt_len) for _ = 1:r] work_rec = _join_vector_workspace_like(tgt, tgt_len) work_residual = _join_vector_workspace_like(tgt, tgt_len) - manifolds = ntuple(k -> _component_manifold(components[k]), r) + manifolds = ntuple(k -> _component_manifold(components_tuple[k]), r) return (; - components, + components = components_tuple, manifolds, r, target = tgt, @@ -310,6 +313,7 @@ function _sum_backend_instance( ) where {T<:AbstractFloat,N} parts = _sum_backend_parts(components, target; init_point) return BTDBackend( + parts.components, parts.manifolds, parts.r, parts.target, diff --git a/src/solvers/btd_als.jl b/src/solvers/btd_als.jl index 2852167..1c68a28 100644 --- a/src/solvers/btd_als.jl +++ b/src/solvers/btd_als.jl @@ -2,7 +2,7 @@ export fit_btd_als @inline function _btd_block_ranks(backend::BTDBackend, b::Int) - M = backend.manifolds[b] + M = _backend_manifold(backend, b) M isa Manifolds.Tucker || throw( ArgumentError( "BTD ALSSolver expects Tucker manifolds, got $(typeof(M)) at block $b.", diff --git a/src/solvers/btd_tsd.jl b/src/solvers/btd_tsd.jl index 3f12788..7fef7d2 100644 --- a/src/solvers/btd_tsd.jl +++ b/src/solvers/btd_tsd.jl @@ -116,7 +116,7 @@ end _check_parts_len(parts, backend.r, "BTD block descent direction") pk = parts[b] eg_b = _btd_block_egrad(backend, parts, b) - rg_b = egrad_to_rgrad(backend.manifolds[b], pk, eg_b) + rg_b = egrad_to_rgrad(_backend_manifold(backend, b), pk, eg_b) decrease = _btd_tangent_dot(eg_b, rg_b) return pk, rg_b, decrease end @@ -135,7 +135,7 @@ function _btd_tsd_block_step( return p, c0, zero(T), 0, false end - Mk = backend.manifolds[b] + Mk = _backend_manifold(backend, b) α = T(solver.stepsize) α_min = T(solver.armijo_alpha_min) contraction = T(solver.armijo_contraction) diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index d0f6d34..6f23b20 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -89,56 +89,6 @@ function _lm_raw_jacobian_matrix( return J end -@inline function _cp_scaled_tangent_factors(λ, U, λ̇, U̇, ::Val{:identity}) - return λ, U, λ̇, U̇ -end - -@inline function _cp_scaled_tangent_factors(λ̃, Ũ, λ̇̃, U̇̃, ::Val{:square}) - λ = λ̃ .^ 2 - U = [Ũ[m] .^ 2 for m in eachindex(Ũ)] - λ̇ = 2 .* λ̃ .* λ̇̃ - U̇ = [2 .* Ũ[m] .* U̇̃[m] for m in eachindex(Ũ)] - return λ, U, λ̇, U̇ -end - -@inline function _cp_scaled_tangent_factors(λ̃, Ũ, λ̇̃, U̇̃, ::Val{:softplus}) - λ = _softplus_value.(λ̃) - U = [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] - λ̇ = _softplus_derivative.(λ̃) .* λ̇̃ - U̇ = [_softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(Ũ)] - return λ, U, λ̇, U̇ -end - -function _cp_rankr_tangent_tensorvec!( - out::AbstractVector{T}, - λ::AbstractVector{T}, - U::Vector{<:AbstractMatrix{T}}, - λ̇::AbstractVector{T}, - U̇::Vector{<:AbstractMatrix{T}}, -) where {T<:AbstractFloat} - fill!(out, zero(T)) - r = length(λ) - @inbounds for k = 1:r - comp = ([λ[k]], [Vector(@view U[m][:, k]) for m in eachindex(U)]...) - xcomp = ([λ̇[k]], [Vector(@view U̇[m][:, k]) for m in eachindex(U̇)]...) - out .+= _segre_tangent_tensorvec(comp, xcomp) - end - return out -end - -function _cp_rank1_tangent_tensorvec!( - out::AbstractVector{T}, - λ::T, - U::Vector{<:AbstractVector{T}}, - λ̇::T, - U̇::Vector{<:AbstractVector{T}}, -) where {T<:AbstractFloat} - comp = ([λ], U...) - xcomp = ([λ̇], U̇...) - copyto!(out, _segre_tangent_tensorvec(comp, xcomp)) - return out -end - function _lm_raw_residual_vector(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) return _lm_raw_residual_vector(cpd_model(model), p) end @@ -166,28 +116,14 @@ function _lm_raw_jacobian_matrix( d = manifold_dimension(M) J = Matrix{T}(undef, ambient_dim, d) coeff = zeros(T, d) - λp, Up = unpack_point_rank1(p, model.dims) + embedding = _cp_parameterization(model) column = Vector{T}(undef, ambient_dim) @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) Xj = ManifoldsBase.get_vector(M, p, coeff, basis) - if model.nonnegative - λ̇p, U̇p = unpack_point_rank1(Xj, model.dims) - kind = _rank1_uses_softplus_metric(model.M) ? Val(:softplus) : Val(:square) - λ, U, λ̇, U̇ = _cp_scaled_tangent_factors( - [λp], - [reshape(u, :, 1) for u in Up], - [λ̇p], - [reshape(u, :, 1) for u in U̇p], - kind, - ) - U_vec = [Vector(@view U[m][:, 1]) for m in eachindex(U)] - U̇_vec = [Vector(@view U̇[m][:, 1]) for m in eachindex(U̇)] - _cp_rank1_tangent_tensorvec!(column, λ[1], U_vec, λ̇[1], U̇_vec) - else - copyto!(column, vec(ManifoldsBase.embed(M, p, Xj))) - end + λ, U, λ̇, U̇ = _cp_rank1_decode_tangent_factors(embedding, model.dims, p, Xj) + _cp_rank1_tangent_tensorvec!(column, λ, U, λ̇, U̇) J[:, j] .= column end return J @@ -208,37 +144,13 @@ function _lm_raw_jacobian_matrix( J = Matrix{T}(undef, ambient_dim, d) coeff = zeros(T, d) column = Vector{T}(undef, ambient_dim) - if model.geometry == :native && !model.nonnegative - pparts = point_parts(p) - @inbounds for j = 1:d - fill!(coeff, zero(T)) - coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p, coeff, basis) - xparts = point_parts(Xj) - fill!(column, zero(T)) - for k = 1:model.r - column .+= _segre_tangent_tensorvec(pparts[k], xparts[k]) - end - J[:, j] .= column - end - return J - end - - λp, Up = - model.nonnegative ? unpack_point_rankr(p, model.dims, model.r) : - unpack_rankr_canonical(p, model.dims, model.r) - kind = - model.nonnegative ? - (model.geometry == :softplus_metric ? Val(:softplus) : Val(:square)) : - Val(:identity) + embedding = _cp_parameterization(model) @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) Xj = ManifoldsBase.get_vector(M, p, coeff, basis) - λ̇p, U̇p = - model.nonnegative ? unpack_point_rankr(Xj, model.dims, model.r) : - unpack_rankr_canonical(Xj, model.dims, model.r) - λ, U, λ̇, U̇ = _cp_scaled_tangent_factors(λp, Up, λ̇p, U̇p, kind) + λ, U, λ̇, U̇ = + _cp_rankr_decode_tangent_factors(embedding, model.dims, model.r, p, Xj) _cp_rankr_tangent_tensorvec!(column, λ, U, λ̇, U̇) J[:, j] .= column end diff --git a/test/basic_tests.jl b/test/basic_tests.jl index ab7913b..37c407c 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -285,6 +285,110 @@ end end end +@testset "LM CPD rank-r Jacobian matches finite differences across geometries" begin + A = randn(5, 4, 3) + r = 2 + cases = ( + (JoinModel(A, r; geometry = :canonical), 1e-7), + (JoinModel(A, r; geometry = :native), 1e-7), + (JoinModel(abs.(A), r; nonnegative = true, geometry = :squaring_metric), 5e-6), + (JoinModel(abs.(A), r; nonnegative = true, geometry = :softplus_metric), 5e-6), + ) + + for (model, tol_fd) in cases + M = TensorKitchen.manifold(model) + p = TensorKitchen._solver_point( + M, + TensorKitchen.initial_point(model, :random; verbose = false), + ) + basis = ManifoldsBase.DefaultOrthonormalBasis() + J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) + @test all(isfinite, J) + + retraction_method = TensorKitchen._solver_retraction_method(M, p) + ϵ = 1e-6 + d = manifold_dimension(M) + for j = 1:min(d, 3) + coeff = zeros(Float64, d) + coeff[j] = 1.0 + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) + p_plus = ManifoldsBase.retract(M, p, ϵ * Xj, retraction_method) + p_minus = ManifoldsBase.retract(M, p, -ϵ * Xj, retraction_method) + r_plus = TensorKitchen._lm_raw_residual_vector(model, p_plus) + r_minus = TensorKitchen._lm_raw_residual_vector(model, p_minus) + fd = (r_plus .- r_minus) ./ (2 * ϵ) + @test maximum(abs.(fd .- J[:, j])) ≤ tol_fd + end + end +end + +@testset "LM CPD Jacobian finite differences across direct parameterizations" begin + dims = (5, 4, 3) + A = randn(dims...) + r = 2 + cases = ( + ("rank1 native", TensorKitchen.Rank1CPDModel(A), 5e-7), + ("rank1 squared", TensorKitchen.Rank1CPDModel(abs.(A); nonnegative = true), 5e-6), + ( + "rank1 softplus", + TensorKitchen.Rank1CPDModel( + abs.(A); + nonnegative = true, + use_softplus_metric = true, + ), + 5e-6, + ), + ("rankr native", TensorKitchen.RankRCPDModel(A, r; geometry = :native), 5e-7), + ("rankr canonical", TensorKitchen.RankRCPDModel(A, r; geometry = :canonical), 5e-7), + ( + "rankr squared", + TensorKitchen.RankRCPDModel( + abs.(A), + r; + nonnegative = true, + geometry = :squaring_metric, + ), + 5e-6, + ), + ( + "rankr softplus", + TensorKitchen.RankRCPDModel( + abs.(A), + r; + nonnegative = true, + geometry = :softplus_metric, + ), + 5e-6, + ), + ) + + for (label, model, tol_fd) in cases + M = TensorKitchen.manifold(model) + p = TensorKitchen._solver_point( + M, + TensorKitchen.initial_point(model, :random; verbose = false), + ) + basis = ManifoldsBase.DefaultOrthonormalBasis() + J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) + @testset "$label" begin + @test all(isfinite, J) + retraction_method = TensorKitchen._solver_retraction_method(M, p) + ϵ = 1e-6 + d = manifold_dimension(M) + r0 = TensorKitchen._lm_raw_residual_vector(model, p) + for j = 1:min(d, 3) + coeff = zeros(Float64, d) + coeff[j] = 1.0 + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) + p_plus = ManifoldsBase.retract(M, p, ϵ * Xj, retraction_method) + r_plus = TensorKitchen._lm_raw_residual_vector(model, p_plus) + fd = (r_plus .- r0) ./ ϵ + @test maximum(abs.(fd .- J[:, j])) ≤ tol_fd + end + end + end +end + @testset "LM normalized and unnormalized objectives take the same step" begin A = randn(6, 5, 4) model = JoinModel(A, 2; geometry = :canonical) @@ -319,6 +423,57 @@ end @test isapprox(res_rel.rel_error, res_abs.rel_error; rtol = 1e-10, atol = 1e-10) end +@testset "CP parameterization tangent decode helpers" begin + dims = (3, 2, 2) + r = 2 + λ̃ = [1.5, -0.4] + Ũ = [randn(dims[m], r) for m = 1:length(dims)] + λ̇̃ = randn(r) + U̇̃ = [randn(dims[m], r) for m = 1:length(dims)] + p = TensorKitchen.pack_point_rankr(λ̃, Ũ, r) + X = TensorKitchen.pack_point_rankr(λ̇̃, U̇̃, r) + + λ_sq, U_sq, λ̇_sq, U̇_sq = TensorKitchen._cp_rankr_decode_tangent_factors( + TensorKitchen.SquaredNonnegativeCPEmbedding(), + dims, + r, + p, + X, + ) + @test λ_sq ≈ λ̃ .^ 2 + @test all(U_sq[m] ≈ Ũ[m] .^ 2 for m in eachindex(U_sq)) + @test λ̇_sq ≈ 2 .* λ̃ .* λ̇̃ + @test all(U̇_sq[m] ≈ 2 .* Ũ[m] .* U̇̃[m] for m in eachindex(U̇_sq)) + + λ_sp, U_sp, λ̇_sp, U̇_sp = TensorKitchen._cp_rankr_decode_tangent_factors( + TensorKitchen.SoftplusNonnegativeCPEmbedding(), + dims, + r, + p, + X, + ) + @test λ_sp ≈ TensorKitchen._softplus_value.(λ̃) + @test all(U_sp[m] ≈ TensorKitchen._softplus_value.(Ũ[m]) for m in eachindex(U_sp)) + @test λ̇_sp ≈ TensorKitchen._softplus_derivative.(λ̃) .* λ̇̃ + @test all( + U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for + m in eachindex(U̇_sp) + ) + + model_sp = TensorKitchen.RankRCPDModel( + randn(dims...), + r; + nonnegative = true, + geometry = :softplus_metric, + ) + q_zero = + CPDPoint(zeros(Float64, r), [zeros(Float64, dims[m], r) for m = 1:length(dims)]) + p_zero = TensorKitchen.pack_cpd_point(model_sp, q_zero) + λ_lat, U_lat = TensorKitchen.unpack_point_rankr(p_zero, dims, r) + @test all(isfinite, λ_lat) + @test all(F -> all(isfinite, F), U_lat) +end + @testset "cpd/approx accept LMSolver" begin A = randn(5, 4, 3) res_cpd_symbol = cpd(A, 2; solver = :lm, maxiter = 2, tol = 1e-6, verbose = false) @@ -2215,6 +2370,9 @@ end manifolds = TensorKitchen._as_join_manifold_tuple(TuckerJoin(size(A), (2, 2, 2), 2)) backend = TensorKitchen._sum_backend_instance(TensorKitchen.BTDBackend, manifolds, A) + @test length(backend.components) == 2 + @test all(c -> c isa TensorKitchen.JoinComponent, backend.components) + @test map(TensorKitchen._component_manifold, backend.components) == manifolds model_btd = JoinModel{Float64,typeof(backend)}(backend) M_btd = TensorKitchen.manifold(model_btd) p_btd = @@ -2230,8 +2388,11 @@ end end for b = 1:backend.r fast_eg = TensorKitchen._btd_block_egrad(backend, parts_btd, b) - residual_eg = - TensorKitchen._tucker_egrad(backend.manifolds[b], parts_btd[b], residual_btd) + residual_eg = TensorKitchen._tucker_egrad( + TensorKitchen._backend_manifold(backend, b), + parts_btd[b], + residual_btd, + ) @test norm(getproperty(fast_eg, :Ċ) - getproperty(residual_eg, :Ċ)) < 1e-10 @test all( norm(F - R) < 1e-10 for @@ -2243,7 +2404,11 @@ end q_btd = TensorKitchen._replace_block_part( p_btd, b, - retract(backend.manifolds[b], parts_btd[b], (-h) * block_grad), + retract( + TensorKitchen._backend_manifold(backend, b), + parts_btd[b], + (-h) * block_grad, + ), ) fd = (TensorKitchen.cost(model_btd, q_btd) - TensorKitchen.cost(model_btd, p_btd)) / From 4e1513a14adb26304e712e28f25763632524ac5a Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 14:53:03 +0200 Subject: [PATCH 28/37] relocate manopt helpers to manopt_helpers.jl --- src/solvers/abstract.jl | 44 --------------------- src/solvers/manopt_helpers.jl | 74 +++++++++++++++++++++++++++++++++++ src/solvers/rcg.jl | 38 +----------------- src/solvers/rgd.jl | 4 +- 4 files changed, 77 insertions(+), 83 deletions(-) diff --git a/src/solvers/abstract.jl b/src/solvers/abstract.jl index bab4b9e..e4c50b2 100644 --- a/src/solvers/abstract.jl +++ b/src/solvers/abstract.jl @@ -31,21 +31,6 @@ end StandardSolverTrace{T}() where {T<:AbstractFloat} = StandardSolverTrace{T}(IterationRecord{T}[], 0, 0, 0.0) -""" - _dual_stop_grad_tol(T, tol; grad_tol=nothing) - -Gradient tolerance paired with `StopWhenCostRelChangeAndGradientLess`. -Defaults to `sqrt(tol)`; callers may pass an explicit `grad_tol` (for example -`grad_tol = tol` on the nonnegative CPD manifold route). -""" -@inline function _dual_stop_grad_tol( - ::Type{T}, - tol::Real, - grad_tol = nothing, -) where {T<:Real} - return isnothing(grad_tol) ? sqrt(T(tol)) : T(grad_tol) -end - function record!( trace::StandardSolverTrace{T}, iter::Int, @@ -411,35 +396,6 @@ function _model_gradient_closure( return model_exact_join_basis_function(model) end -@inline _unwrap_solver_manifold(M) = hasproperty(M, :M) ? getproperty(M, :M) : M - -# The actual methods depend on the registered defaults, e.g. custom manifolds such -# as Segre or SoftplusEuclidean may choose ExponentialRetraction, while sphere-like -# factors may choose their ManifoldsBase default. -@inline function _default_component_retraction_method(Mi, pi) - return ManifoldsBase.default_retraction_method(Mi, typeof(pi)) -end - -function _solver_retraction_method(M, p) - return _solver_retraction_method_unwrapped(_unwrap_solver_manifold(M), p) -end - -function _solver_retraction_method_unwrapped(M::ProductManifold, p) - pparts0 = point_parts(p) - pparts = pparts0 isa Tuple ? pparts0 : Tuple(pparts0) - n = length(M.manifolds) - length(pparts) == n || throw( - ArgumentError( - "Cannot derive solver retraction method: ProductManifold has $n factors but point has $(length(pparts)) parts.", - ), - ) - methods = - ntuple(i -> _default_component_retraction_method(M.manifolds[i], pparts[i]), n) - return ManifoldsBase.ProductRetraction(methods) -end - -_solver_retraction_method_unwrapped(M, p) = _default_component_retraction_method(M, p) - """ _prepare_solver_problem(model; init, gradient_mode) diff --git a/src/solvers/manopt_helpers.jl b/src/solvers/manopt_helpers.jl index e6bc167..944e709 100644 --- a/src/solvers/manopt_helpers.jl +++ b/src/solvers/manopt_helpers.jl @@ -13,6 +13,22 @@ Base.unsafe_write(::_SolverDebugSink, ::Ptr{UInt8}, n::UInt) = Int(n) # Shared no-op IO used by Manopt debug groups. const _SOLVER_DEBUG_SINK = _SolverDebugSink() +# Gradient tolerance paired with StopWhenCostRelChangeAndGradientLess. +""" + _dual_stop_grad_tol(T, tol; grad_tol=nothing) + +Gradient tolerance paired with `StopWhenCostRelChangeAndGradientLess`. +Defaults to `sqrt(tol)`; callers may pass an explicit `grad_tol` (for example +`grad_tol = tol` on the nonnegative CPD manifold route). +""" +@inline function _dual_stop_grad_tol( + ::Type{T}, + tol::Real, + grad_tol = nothing, +) where {T<:Real} + return isnothing(grad_tol) ? sqrt(T(tol)) : T(grad_tol) +end + # Stop when both the relative cost change and Riemannian gradient norm are small. mutable struct StopWhenCostRelChangeAndGradientLess{T<:Real} <: Manopt.StoppingCriterion tol_cost::T @@ -136,6 +152,64 @@ function _solver_point(M, p0) end +# Unwrap solver manifold wrappers down to the underlying manifold object. +@inline _unwrap_solver_manifold(M) = hasproperty(M, :M) ? getproperty(M, :M) : M + + +# The actual methods depend on the registered defaults, e.g. custom manifolds such +# as Segre or SoftplusEuclidean may choose ExponentialRetraction, while sphere-like +# factors may choose their ManifoldsBase default. +@inline function _default_component_retraction_method(Mi, pi) + return ManifoldsBase.default_retraction_method(Mi, typeof(pi)) +end + + +# Choose a solver retraction method, including per-factor product retractions. +function _solver_retraction_method(M, p) + return _solver_retraction_method_unwrapped(_unwrap_solver_manifold(M), p) +end + +function _solver_retraction_method_unwrapped(M::ProductManifold, p) + pparts0 = point_parts(p) + pparts = pparts0 isa Tuple ? pparts0 : Tuple(pparts0) + n = length(M.manifolds) + length(pparts) == n || throw( + ArgumentError( + "Cannot derive solver retraction method: ProductManifold has $n factors but point has $(length(pparts)) parts.", + ), + ) + methods = + ntuple(i -> _default_component_retraction_method(M.manifolds[i], pparts[i]), n) + return ManifoldsBase.ProductRetraction(methods) +end + +_solver_retraction_method_unwrapped(M, p) = _default_component_retraction_method(M, p) + + +# Conservative compatibility probe for vector transports used by Manopt solvers. +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 + + +# Choose a vector transport method that works with the current manifold/point layout. +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 ManifoldsBase.default_vector_transport_method(M, typeof(p)) +end + + # Detect pullback nonnegative geometries that need conservative line search. function _contains_sqeuclidean_manifold(M) M2 = _unwrap_solver_manifold(M) diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index f58f039..3f12af8 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -1,42 +1,6 @@ # solvers/rcg.jl — Riemannian Conjugate Gradient export RCGSolver -# Vector transport selection -""" - _supports_vector_transport_to(M, p, vt, retraction_method) - -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 - -""" - _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 ManifoldsBase.default_vector_transport_method(M, typeof(p)) -end - # RCG coefficient and restart rule selection function _rcg_coefficient_rule( M, @@ -174,7 +138,7 @@ function solve_rcg( ) return _manopt_finish_result( - get_solver_result(state), + _tk_get_solver_result(state), state, callbacks.progress, diagnostics_recorder, diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index e7517b5..f87d199 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -107,7 +107,7 @@ function solve_rgd( ) return _manopt_finish_result( - get_solver_result(state), + _tk_get_solver_result(state), state, callbacks.progress, diagnostics_recorder, @@ -188,7 +188,7 @@ function solve_rgd_fixed( ) return _manopt_finish_result( - get_solver_result(state), + _tk_get_solver_result(state), state, callbacks.progress, diagnostics_recorder, From c2efa31bd3aa7470823ef6f347d4b1d00e9d2481 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 17:27:15 +0200 Subject: [PATCH 29/37] refine component abstraction with LM Jacobian matrix ordering test --- src/join/join_backend.jl | 60 ++++++++++++++++++++++---- src/join/join_model.jl | 26 +++++++++++- src/solvers/lm.jl | 76 +++++++++++++++++++-------------- test/basic_tests.jl | 92 ++++++++++++++++++++++++++++++++++++++++ 4 files changed, 211 insertions(+), 43 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 1cb093d..863159e 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -8,7 +8,9 @@ _is_join_component_like(x) = _is_manifold_like(x) _wrap_join_component(component::JoinComponent) = component _wrap_join_component(manifold::AbstractManifold) = JoinComponent(manifold) -_component_manifold(component::JoinComponent) = component.manifold +_component_embedding(component::JoinComponent) = component_embedding(component) +_component_embedding(::AbstractManifold) = DefaultJoinEmbedding() +_component_manifold(component::JoinComponent) = component_manifold(component) _component_manifold(manifold::AbstractManifold) = manifold _backend_components(backend::JoinBackend) = backend.components _backend_components(backend::BTDBackend) = backend.components @@ -120,8 +122,12 @@ function _component_egrad(::DefaultJoinEmbedding, M::Manifolds.Segre, p, residua return pack_tangent_rank1_segre(grad_λ, grad_U) end -_component_egrad(component::JoinComponent, p, residual) = - _component_egrad(component.embedding, component.manifold, p, residual) +_component_egrad(component::JoinComponent, p, residual) = _component_egrad( + _component_embedding(component), + _component_manifold(component), + p, + residual, +) _component_egrad(M, p, residual) = _component_egrad(DefaultJoinEmbedding(), M, p, residual) _manifold_egrad(M, p, residual) = _component_egrad(M, p, residual) @@ -191,7 +197,32 @@ end ambient_length(M::Manifolds.Segre) = prod(factor_dims(M)) ambient_length(M::Manifolds.Tucker) = prod(factor_dims(M)) -ambient_length(component::JoinComponent) = ambient_length(component.manifold) +ambient_length(component::JoinComponent) = ambient_length(component_manifold(component)) + +function component_basis_vector( + component, + p, + coeffs; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) + return ManifoldsBase.get_vector(component_manifold(component), p, coeffs, basis) +end + +function component_basis_vector( + component, + p, + j::Integer; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) + T = _scalar_eltype(p) + d = component_tangent_dimension(component, p) + 1 <= j <= d || throw( + BoundsError("Component basis index $j is out of bounds for tangent dimension $d."), + ) + coeffs = zeros(T, d) + coeffs[j] = one(T) + return component_basis_vector(component, p, coeffs; basis) +end """ _join_vector_workspace_like(target, n) returns AbstractVector @@ -613,15 +644,25 @@ function _component_ambient_embedding!( end function _component_ambient_embedding!(out::AbstractVector, component::JoinComponent, p) - return _component_ambient_embedding!(out, component.embedding, component.manifold, p) + return _component_ambient_embedding!( + out, + _component_embedding(component), + _component_manifold(component), + p, + ) end +component_ambient_embedding!(out::AbstractVector, component, p) = + _component_ambient_embedding!(out, component, p) + """ _component_ambient_pushforward!(out, component, p, X) Write the ambient pushforward `DΦ(p)[X]` of one join component into `out`. This is the component-level differential used by LM Jacobian assembly on -generic `JoinModel`s. +generic `JoinModel`s. Implementations must overwrite `out` completely rather +than accumulating into it, since LM Jacobian assembly reuses the same work +buffer across columns. """ function _component_ambient_pushforward!(out::AbstractVector, M, p, X) return _component_ambient_pushforward!(out, DefaultJoinEmbedding(), M, p, X) @@ -667,13 +708,16 @@ function _component_ambient_pushforward!( ) return _component_ambient_pushforward!( out, - component.embedding, - component.manifold, + _component_embedding(component), + _component_manifold(component), p, X, ) end +component_ambient_pushforward!(out::AbstractVector, component, p, X) = + _component_ambient_pushforward!(out, component, p, X) + function _subtract_ambient_tensor!( residual::AbstractArray{T,N}, component, diff --git a/src/join/join_model.jl b/src/join/join_model.jl index 1b92e2c..1a42c87 100644 --- a/src/join/join_model.jl +++ b/src/join/join_model.jl @@ -1,6 +1,17 @@ # join/join_model.jl — Join-front-end model and backend type definitions -export AbstractJoinBackend, JoinComponent, JoinModel, CPDBackend, JoinBackend, BTDBackend +export AbstractJoinBackend, + JoinComponent, + JoinModel, + CPDBackend, + JoinBackend, + BTDBackend, + component_manifold, + component_embedding, + component_tangent_dimension, + component_basis_vector, + component_ambient_embedding!, + component_ambient_pushforward! # BTDBackend is defined in `btd/model.jl` (includes contraction workspace). abstract type AbstractJoinBackend end @@ -15,7 +26,18 @@ struct DefaultJoinEmbedding end JoinComponent(manifold::M) where {M} = JoinComponent{M,DefaultJoinEmbedding}(manifold, DefaultJoinEmbedding()) -manifold(component::JoinComponent) = component.manifold +component_manifold(component::JoinComponent) = component.manifold +component_embedding(component::JoinComponent) = component.embedding +component_tangent_dimension(component::JoinComponent) = + manifold_dimension(component_manifold(component)) +component_tangent_dimension(component::JoinComponent, p) = + component_tangent_dimension(component) +component_manifold(M::AbstractManifold) = M +component_embedding(::AbstractManifold) = DefaultJoinEmbedding() +component_tangent_dimension(M::AbstractManifold) = manifold_dimension(M) +component_tangent_dimension(M::AbstractManifold, p) = component_tangent_dimension(M) + +manifold(component::JoinComponent) = component_manifold(component) struct JoinModel{T<:AbstractFloat,B<:AbstractJoinBackend} <: AbstractDecompositionModel{T} backend::B diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 6f23b20..560870e 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -36,30 +36,6 @@ solver_symbol(::LMSolver) = :lm one(T) / sqrt(T(normA2)) : one(T) end -function _join_tangent_ambient_vector!( - out::AbstractVector, - backend::Union{JoinBackend,BTDBackend}, - p, - X, -) - parts = point_parts(p) - xparts = point_parts(X) - _check_parts_len(parts, backend.r, "_join_tangent_ambient_vector!") - _check_parts_len(xparts, backend.r, "_join_tangent_ambient_vector!") - fill!(out, zero(eltype(out))) - components = _backend_components(backend) - @inbounds for k = 1:backend.r - _component_ambient_pushforward!( - backend.component_bufs[k], - components[k], - parts[k], - xparts[k], - ) - out .+= backend.component_bufs[k] - end - return out -end - function _lm_raw_residual_vector( model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, p, @@ -67,6 +43,24 @@ function _lm_raw_residual_vector( return copy(_join_residual!(model.backend, p)) end +function _join_component_jacobian_block!( + J::AbstractMatrix{T}, + next_col::Int, + component, + p_component, + work_vec::AbstractVector{T}; + basis = ManifoldsBase.DefaultOrthonormalBasis(), +) where {T<:AbstractFloat} + d_component = component_tangent_dimension(component, p_component) + @inbounds for j = 1:d_component + Xj = component_basis_vector(component, p_component, j; basis) + component_ambient_pushforward!(work_vec, component, p_component, Xj) + J[:, next_col] .= work_vec + next_col += 1 + end + return next_col +end + function _lm_raw_jacobian_matrix( model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, M, @@ -75,17 +69,33 @@ function _lm_raw_jacobian_matrix( ) T = _scalar_eltype(p) ambient_dim = length(tensor(model)) - d = manifold_dimension(M) + backend = model.backend + components = _backend_components(backend) + parts = point_parts(p) + _check_parts_len(parts, backend.r, "_lm_raw_jacobian_matrix") + d = sum(component_tangent_dimension(components[k], parts[k]) for k = 1:backend.r) + d == manifold_dimension(M) || throw( + DimensionMismatch( + "JoinModel Jacobian assembly expected tangent dimension $d from components but manifold reports $(manifold_dimension(M)).", + ), + ) J = Matrix{T}(undef, ambient_dim, d) - coeff = zeros(T, d) - column = similar(model.backend.work_rec, T, ambient_dim) - @inbounds for j = 1:d - fill!(coeff, zero(T)) - coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p, coeff, basis) - _join_tangent_ambient_vector!(column, model.backend, p, Xj) - J[:, j] .= column + next_col = 1 + @inbounds for k = 1:backend.r + next_col = _join_component_jacobian_block!( + J, + next_col, + components[k], + parts[k], + backend.component_bufs[k]; + basis, + ) end + next_col == d + 1 || throw( + DimensionMismatch( + "JoinModel Jacobian assembly filled $(next_col - 1) columns but expected $d.", + ), + ) return J end diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 37c407c..5646036 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -202,6 +202,10 @@ end JoinModel((TensorKitchen.JoinComponent(segres[1]), segres[2], segres[3]), A) @test model_component.backend.components[1] isa TensorKitchen.JoinComponent @test map(TensorKitchen.manifold, model_component.backend.components) == segres + @test TensorKitchen.component_manifold(model_component.backend.components[1]) == + segres[1] + @test TensorKitchen.component_embedding(model_component.backend.components[1]) isa + TensorKitchen.DefaultJoinEmbedding p = TensorKitchen.initial_point(model, :random; verbose = false) @test length(TensorKitchen.point_parts(p)) == 3 @@ -255,6 +259,36 @@ end end end +function _reference_join_jacobian_from_product_basis(model, M, p; basis) + backend = model.backend + parts = TensorKitchen.point_parts(p) + T = TensorKitchen._scalar_eltype(p) + ambient_dim = length(TensorKitchen.tensor(model)) + d = manifold_dimension(M) + J = Matrix{T}(undef, ambient_dim, d) + coeff = zeros(T, d) + col = Vector{T}(undef, ambient_dim) + buf = Vector{T}(undef, ambient_dim) + for j = 1:d + fill!(coeff, zero(T)) + coeff[j] = one(T) + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) + xparts = TensorKitchen.point_parts(Xj) + fill!(col, zero(T)) + for k = 1:backend.r + TensorKitchen.component_ambient_pushforward!( + buf, + TensorKitchen._backend_component(backend, k), + parts[k], + xparts[k], + ) + col .+= buf + end + J[:, j] .= col + end + return J +end + @testset "LM generic join Jacobian uses component pushforwards" begin A = randn(5, 4, 3) model = JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), A) @@ -266,6 +300,33 @@ end basis = ManifoldsBase.DefaultOrthonormalBasis() J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) @test all(isfinite, J) + @test sum( + TensorKitchen.component_tangent_dimension( + TensorKitchen._backend_component(model.backend, k), + TensorKitchen.point_parts(p)[k], + ) for k = 1:model.backend.r + ) == manifold_dimension(M) + J_ref = _reference_join_jacobian_from_product_basis(model, M, p; basis) + @test maximum(abs.(J .- J_ref)) ≤ 1e-12 + + parts = TensorKitchen.point_parts(p) + c1 = TensorKitchen._backend_component(model.backend, 1) + ξ1 = TensorKitchen.component_basis_vector(c1, parts[1], 1; basis) + buf = similar(model.backend.work_rec) + TensorKitchen.component_ambient_pushforward!(buf, c1, parts[1], ξ1) + @test maximum(abs.(buf .- J[:, 1])) ≤ 1e-10 + + c2 = TensorKitchen._backend_component(model.backend, 2) + ξ2 = TensorKitchen.component_basis_vector(c2, parts[2], 1; basis) + TensorKitchen.component_ambient_pushforward!(buf, c2, parts[2], ξ2) + offset2 = TensorKitchen.component_tangent_dimension(c1, parts[1]) + 1 + @test maximum(abs.(buf .- J[:, offset2])) ≤ 1e-10 + + fill!(buf, NaN) + TensorKitchen.component_ambient_pushforward!(buf, c1, parts[1], ξ1) + fresh = similar(buf) + TensorKitchen.component_ambient_pushforward!(fresh, c1, parts[1], ξ1) + @test buf == fresh retraction_method = TensorKitchen._solver_retraction_method(M, p) ϵ = 1e-6 @@ -2377,6 +2438,37 @@ end M_btd = TensorKitchen.manifold(model_btd) p_btd = TensorKitchen._solver_point(M_btd, TensorKitchen.initial_point(model_btd, :random)) + basis_btd = ManifoldsBase.DefaultOrthonormalBasis() + J_btd = + TensorKitchen._lm_raw_jacobian_matrix(model_btd, M_btd, p_btd; basis = basis_btd) + @test all(isfinite, J_btd) + @test sum( + TensorKitchen.component_tangent_dimension( + TensorKitchen._backend_component(backend, k), + TensorKitchen.point_parts(p_btd)[k], + ) for k = 1:backend.r + ) == manifold_dimension(M_btd) + J_btd_ref = _reference_join_jacobian_from_product_basis( + model_btd, + M_btd, + p_btd; + basis = basis_btd, + ) + @test maximum(abs.(J_btd .- J_btd_ref)) ≤ 1e-12 + + retraction_method_btd = TensorKitchen._solver_retraction_method(M_btd, p_btd) + residual0_btd = TensorKitchen._lm_raw_residual_vector(model_btd, p_btd) + ϵ_btd = 1e-6 + for j = 1:min(manifold_dimension(M_btd), 2) + coeff = zeros(Float64, manifold_dimension(M_btd)) + coeff[j] = 1.0 + Xj = ManifoldsBase.get_vector(M_btd, p_btd, coeff, basis_btd) + p_plus = ManifoldsBase.retract(M_btd, p_btd, ϵ_btd * Xj, retraction_method_btd) + r_plus = TensorKitchen._lm_raw_residual_vector(model_btd, p_plus) + fd = (r_plus .- residual0_btd) ./ ϵ_btd + @test maximum(abs.(fd .- J_btd[:, j])) ≤ 5e-6 + end + parts_btd = TensorKitchen.point_parts(p_btd) residual_btd = TensorKitchen._join_residual!(backend, p_btd) tangent_dot_btd(a, b) = begin From 3d50ab1865d7203e93078b64b629e5c0efd64552 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Sat, 20 Jun 2026 18:49:16 +0200 Subject: [PATCH 30/37] rename CPEmbedding to CPparam in AbstractCPParametrization --- src/cpd/model/parameterizations.jl | 64 +++++++++++++++--------------- src/cpd/model/rank1.jl | 6 +-- src/cpd/model/rankr.jl | 6 +-- test/basic_tests.jl | 4 +- 4 files changed, 38 insertions(+), 42 deletions(-) diff --git a/src/cpd/model/parameterizations.jl b/src/cpd/model/parameterizations.jl index 6d2f35f..63e2ee5 100644 --- a/src/cpd/model/parameterizations.jl +++ b/src/cpd/model/parameterizations.jl @@ -2,40 +2,40 @@ abstract type AbstractCPParameterization end -struct NativeCPEmbedding <: AbstractCPParameterization end -struct CanonicalCPEmbedding <: AbstractCPParameterization end -struct SquaredNonnegativeCPEmbedding <: AbstractCPParameterization end -struct SoftplusNonnegativeCPEmbedding <: AbstractCPParameterization end +struct NativeCPParam <: AbstractCPParameterization end +struct CanonicalCPParam <: AbstractCPParameterization end +struct SquaredNNCPParam <: AbstractCPParameterization end +struct SoftplusNNCPParam <: AbstractCPParameterization end @inline _cp_softplus_encode_value(x::T) where {T<:AbstractFloat} = _invsoftplus(max(x, eps(T))) -function _cp_rank1_decode_factors(::NativeCPEmbedding, dims, p) +function _cp_rank1_decode_factors(::NativeCPParam, dims, p) return unpack_point_rank1(p, dims) end -function _cp_rank1_decode_factors(::SquaredNonnegativeCPEmbedding, dims, p) +function _cp_rank1_decode_factors(::SquaredNNCPParam, dims, p) λ̃, Ũ = unpack_point_rank1(p, dims) return λ̃^2, [Ũ[m] .^ 2 for m in eachindex(Ũ)] end -function _cp_rank1_decode_factors(::SoftplusNonnegativeCPEmbedding, dims, p) +function _cp_rank1_decode_factors(::SoftplusNNCPParam, dims, p) λ̃, Ũ = unpack_point_rank1(p, dims) return _softplus_value(λ̃), [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] end -function _cp_rank1_encode_point(::NativeCPEmbedding, λ, U) +function _cp_rank1_encode_point(::NativeCPParam, λ, U) return pack_point_rank1_segre(λ, U) end -function _cp_rank1_encode_point(::SquaredNonnegativeCPEmbedding, λ, U) +function _cp_rank1_encode_point(::SquaredNNCPParam, λ, U) return pack_point_rank1( sqrt(max(λ, zero(λ))), [sqrt.(max.(u, zero(eltype(u)))) for u in U], ) end -function _cp_rank1_encode_point(::SoftplusNonnegativeCPEmbedding, λ, U) +function _cp_rank1_encode_point(::SoftplusNNCPParam, λ, U) T = typeof(λ) return pack_point_rank1( _cp_softplus_encode_value(max(λ, zero(T))), @@ -43,11 +43,11 @@ function _cp_rank1_encode_point(::SoftplusNonnegativeCPEmbedding, λ, U) ) end -function _cp_rank1_seed_point(::NativeCPEmbedding, λ, U) +function _cp_rank1_seed_point(::NativeCPParam, λ, U) return pack_point_rank1_segre(λ, U) end -function _cp_rank1_seed_point(::SquaredNonnegativeCPEmbedding, λ, U) +function _cp_rank1_seed_point(::SquaredNNCPParam, λ, U) T = typeof(λ) return pack_point_rank1( sqrt(max(abs(λ), eps(T))), @@ -55,7 +55,7 @@ function _cp_rank1_seed_point(::SquaredNonnegativeCPEmbedding, λ, U) ) end -function _cp_rank1_seed_point(::SoftplusNonnegativeCPEmbedding, λ, U) +function _cp_rank1_seed_point(::SoftplusNNCPParam, λ, U) T = typeof(λ) return pack_point_rank1( _invsoftplus(max(abs(λ), eps(T))), @@ -81,13 +81,13 @@ function _cp_rank1_tangent_tensorvec!( return out end -function _cp_rank1_decode_tangent_factors(::NativeCPEmbedding, dims, p, X) +function _cp_rank1_decode_tangent_factors(::NativeCPParam, dims, p, X) λ, U = unpack_point_rank1(p, dims) λ̇, U̇ = unpack_point_rank1(X, dims) return λ, U, λ̇, U̇ end -function _cp_rank1_decode_tangent_factors(::SquaredNonnegativeCPEmbedding, dims, p, X) +function _cp_rank1_decode_tangent_factors(::SquaredNNCPParam, dims, p, X) λ̃, Ũ = unpack_point_rank1(p, dims) λ̇̃, U̇̃ = unpack_point_rank1(X, dims) λ = λ̃^2 @@ -97,7 +97,7 @@ function _cp_rank1_decode_tangent_factors(::SquaredNonnegativeCPEmbedding, dims, return λ, U, λ̇, U̇ end -function _cp_rank1_decode_tangent_factors(::SoftplusNonnegativeCPEmbedding, dims, p, X) +function _cp_rank1_decode_tangent_factors(::SoftplusNNCPParam, dims, p, X) λ̃, Ũ = unpack_point_rank1(p, dims) λ̇̃, U̇̃ = unpack_point_rank1(X, dims) λ = _softplus_value(λ̃) @@ -107,33 +107,33 @@ function _cp_rank1_decode_tangent_factors(::SoftplusNonnegativeCPEmbedding, dims return λ, U, λ̇, U̇ end -function _cp_rankr_decode_factors(::NativeCPEmbedding, dims, r, p) +function _cp_rankr_decode_factors(::NativeCPParam, dims, r, p) return unpack_rankr_native(p, dims, r) end -function _cp_rankr_decode_factors(::CanonicalCPEmbedding, dims, r, p) +function _cp_rankr_decode_factors(::CanonicalCPParam, dims, r, p) return unpack_rankr_canonical(p, dims, r) end -function _cp_rankr_decode_factors(::SquaredNonnegativeCPEmbedding, dims, r, p) +function _cp_rankr_decode_factors(::SquaredNNCPParam, dims, r, p) λ̃, Ũ = unpack_point_rankr(p, dims, r) return λ̃ .^ 2, [Ũ[m] .^ 2 for m in eachindex(Ũ)] end -function _cp_rankr_decode_factors(::SoftplusNonnegativeCPEmbedding, dims, r, p) +function _cp_rankr_decode_factors(::SoftplusNNCPParam, dims, r, p) λ̃, Ũ = unpack_point_rankr(p, dims, r) return _softplus_value.(λ̃), [_softplus_value.(Ũ[m]) for m in eachindex(Ũ)] end -function _cp_rankr_encode_point(::NativeCPEmbedding, λ, U, r) +function _cp_rankr_encode_point(::NativeCPParam, λ, U, r) return pack_rankr_native(λ, U, r) end -function _cp_rankr_encode_point(::CanonicalCPEmbedding, λ, U, r) +function _cp_rankr_encode_point(::CanonicalCPParam, λ, U, r) return pack_rankr_canonical(λ, U, r) end -function _cp_rankr_encode_point(::SquaredNonnegativeCPEmbedding, λ, U, r) +function _cp_rankr_encode_point(::SquaredNNCPParam, λ, U, r) T = eltype(λ) return pack_point_rankr( sqrt.(max.(λ, zero(T))), @@ -142,7 +142,7 @@ function _cp_rankr_encode_point(::SquaredNonnegativeCPEmbedding, λ, U, r) ) end -function _cp_rankr_encode_point(::SoftplusNonnegativeCPEmbedding, λ, U, r) +function _cp_rankr_encode_point(::SoftplusNNCPParam, λ, U, r) T = eltype(λ) return pack_point_rankr( _cp_softplus_encode_value.(max.(λ, zero(T))), @@ -151,15 +151,15 @@ function _cp_rankr_encode_point(::SoftplusNonnegativeCPEmbedding, λ, U, r) ) end -function _cp_rankr_seed_point(::NativeCPEmbedding, λ, U, r) +function _cp_rankr_seed_point(::NativeCPParam, λ, U, r) return pack_rankr_native(λ, U, r) end -function _cp_rankr_seed_point(::CanonicalCPEmbedding, λ, U, r) +function _cp_rankr_seed_point(::CanonicalCPParam, λ, U, r) return pack_rankr_canonical(λ, U, r) end -function _cp_rankr_seed_point(::SquaredNonnegativeCPEmbedding, λ, U, r) +function _cp_rankr_seed_point(::SquaredNNCPParam, λ, U, r) T = eltype(λ) return pack_point_rankr( sqrt.(max.(abs.(λ), eps(T))), @@ -168,7 +168,7 @@ function _cp_rankr_seed_point(::SquaredNonnegativeCPEmbedding, λ, U, r) ) end -function _cp_rankr_seed_point(::SoftplusNonnegativeCPEmbedding, λ, U, r) +function _cp_rankr_seed_point(::SoftplusNNCPParam, λ, U, r) T = eltype(λ) return pack_point_rankr( _invsoftplus.(max.(abs.(λ), eps(T))), @@ -204,7 +204,7 @@ function _cp_rankr_tangent_tensorvec!( return out end -function _cp_rankr_decode_tangent_factors(::NativeCPEmbedding, dims, r, p, X) +function _cp_rankr_decode_tangent_factors(::NativeCPParam, dims, r, p, X) λ, U = unpack_rankr_native(p, dims, r) xparts = parts_tuple(X) length(xparts) == r || throw( @@ -224,13 +224,13 @@ function _cp_rankr_decode_tangent_factors(::NativeCPEmbedding, dims, r, p, X) return λ, U, λ̇, U̇ end -function _cp_rankr_decode_tangent_factors(::CanonicalCPEmbedding, dims, r, p, X) +function _cp_rankr_decode_tangent_factors(::CanonicalCPParam, dims, r, p, X) λ, U = unpack_rankr_canonical(p, dims, r) λ̇, U̇ = unpack_rankr_canonical(X, dims, r) return λ, U, λ̇, U̇ end -function _cp_rankr_decode_tangent_factors(::SquaredNonnegativeCPEmbedding, dims, r, p, X) +function _cp_rankr_decode_tangent_factors(::SquaredNNCPParam, dims, r, p, X) λ̃, Ũ = unpack_point_rankr(p, dims, r) λ̇̃, U̇̃ = unpack_point_rankr(X, dims, r) λ = λ̃ .^ 2 @@ -240,7 +240,7 @@ function _cp_rankr_decode_tangent_factors(::SquaredNonnegativeCPEmbedding, dims, return λ, U, λ̇, U̇ end -function _cp_rankr_decode_tangent_factors(::SoftplusNonnegativeCPEmbedding, dims, r, p, X) +function _cp_rankr_decode_tangent_factors(::SoftplusNNCPParam, dims, r, p, X) λ̃, Ũ = unpack_point_rankr(p, dims, r) λ̇̃, U̇̃ = unpack_point_rankr(X, dims, r) λ = _softplus_value.(λ̃) diff --git a/src/cpd/model/rank1.jl b/src/cpd/model/rank1.jl index b27e55c..303a918 100644 --- a/src/cpd/model/rank1.jl +++ b/src/cpd/model/rank1.jl @@ -60,10 +60,8 @@ manifold(model::Rank1CPDModel) = model.M _cp_parameterization(model::Rank1CPDModel) = model.nonnegative ? - ( - _rank1_uses_softplus_metric(model.M) ? SoftplusNonnegativeCPEmbedding() : - SquaredNonnegativeCPEmbedding() - ) : NativeCPEmbedding() + (_rank1_uses_softplus_metric(model.M) ? SoftplusNNCPParam() : SquaredNNCPParam()) : + NativeCPParam() function embed_point(model::Rank1CPDModel{T,N}, p) where {T,N} return _cp_rank1_embed_tensor(_cp_parameterization(model), model.dims, p) diff --git a/src/cpd/model/rankr.jl b/src/cpd/model/rankr.jl index b9eded1..cba003e 100644 --- a/src/cpd/model/rankr.jl +++ b/src/cpd/model/rankr.jl @@ -102,10 +102,8 @@ manifold(model::RankRCPDModel) = model.M _cp_parameterization(model::RankRCPDModel) = model.nonnegative ? - ( - model.geometry == :softplus_metric ? SoftplusNonnegativeCPEmbedding() : - SquaredNonnegativeCPEmbedding() - ) : (model.geometry == :native ? NativeCPEmbedding() : CanonicalCPEmbedding()) + (model.geometry == :softplus_metric ? SoftplusNNCPParam() : SquaredNNCPParam()) : + (model.geometry == :native ? NativeCPParam() : CanonicalCPParam()) function embed_point(model::RankRCPDModel{T,N}, p) where {T,N} return _cp_rankr_embed_tensor(_cp_parameterization(model), model.dims, model.r, p) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 5646036..9729d9e 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -495,7 +495,7 @@ end X = TensorKitchen.pack_point_rankr(λ̇̃, U̇̃, r) λ_sq, U_sq, λ̇_sq, U̇_sq = TensorKitchen._cp_rankr_decode_tangent_factors( - TensorKitchen.SquaredNonnegativeCPEmbedding(), + TensorKitchen.SquaredNNCPParam(), dims, r, p, @@ -507,7 +507,7 @@ end @test all(U̇_sq[m] ≈ 2 .* Ũ[m] .* U̇̃[m] for m in eachindex(U̇_sq)) λ_sp, U_sp, λ̇_sp, U̇_sp = TensorKitchen._cp_rankr_decode_tangent_factors( - TensorKitchen.SoftplusNonnegativeCPEmbedding(), + TensorKitchen.SoftplusNNCPParam(), dims, r, p, From 99181679089509acd6992877cd318f2b958bd1fd Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 2 Jul 2026 15:03:30 +0200 Subject: [PATCH 31/37] update Manopt version and fix the compatibility --- Project.toml | 2 +- src/solvers/lm.jl | 24 ++++++++++++++++++------ 2 files changed, 19 insertions(+), 7 deletions(-) diff --git a/Project.toml b/Project.toml index a3bb5e9..3f25b5f 100644 --- a/Project.toml +++ b/Project.toml @@ -17,7 +17,7 @@ TensorOperations = "6aa20fa7-93e2-5fca-9bc0-fbd0db3c71a2" JuliaFormatter = "2.8.5" Manifolds = "0.11.28" ManifoldsBase = "2.3.5" -Manopt = "0.5.37" +Manopt = "0.6" ProgressMeter = "1.11.0" RecursiveArrayTools = "4.3" TensorOperations = "5.6" diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 560870e..5f00204 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -238,6 +238,14 @@ function solve_lm( jacobian = _lm_jacobian_function(model, T, normA2, setup.uses_relative_objective; basis) initial_residual_values = residual(M, p0_local) initial_jacobian_f = jacobian(M, p0_local) + tangent_space = TangentSpace(M, p0_local) + lm_subsolver_state = Manopt.CoordinatesNormalSystemState( + tangent_space, + zero_vector(M, p0_local); + evaluation = Manopt.InplaceEvaluation(), + linsolve = linear_subsolver, + basis = basis, + ) retraction_method = _solver_retraction_method(M, p0_local) stopping = StopWhenAny( StopAfterIteration(maxiter), @@ -269,16 +277,20 @@ function solve_lm( p0_local; evaluation = Manopt.AllocatingEvaluation(), function_type = Manopt.FunctionVectorialType(), - jacobian_type = Manopt.CoordinateVectorialType(basis), + jacobian_type = Manopt.CoefficientVectorialType(basis), retraction_method = retraction_method, stopping_criterion = stopping, initial_residual_values = initial_residual_values, - initial_jacobian_f = initial_jacobian_f, - η = η, + initial_jacobian_matrices = [initial_jacobian_f], + candidate_acceptance_threshold = η, + damping_increase_factor = β, + damping_increase_threshold = η, + damping_reduction_threshold = expect_zero_residual ? η : Inf, + damping_reduction_factor = inv(T(β)), damping_term_min = damping_term_min, - β = β, - expect_zero_residual = expect_zero_residual, - linear_subsolver! = linear_subsolver, + initial_damping_term = damping_term_min, + use_unified_basis = true, + sub_state = lm_subsolver_state, debug = callbacks.debug_actions, return_state = true, ) From 410acb85361f6693563172ae3f8306ffaf6360f1 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Thu, 2 Jul 2026 15:36:23 +0200 Subject: [PATCH 32/37] Implement LM updates before resolving conflicts --- src/join/cpd_backend.jl | 10 ++ src/join/join_backend.jl | 48 +++++++++ src/solvers/lm.jl | 213 ++++++++++++++------------------------- test/basic_tests.jl | 32 ++++++ 4 files changed, 166 insertions(+), 137 deletions(-) diff --git a/src/join/cpd_backend.jl b/src/join/cpd_backend.jl index b062160..7d87b0e 100644 --- a/src/join/cpd_backend.jl +++ b/src/join/cpd_backend.jl @@ -59,6 +59,16 @@ initial_point( ) = initial_point(cpd_model(model), init; kwargs...) cost(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = cost(cpd_model(model), p) egrad(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = egrad(cpd_model(model), p) +residual(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = residual(cpd_model(model), p) +differential_action!(out::AbstractVector, model::JoinModel{<:AbstractFloat,<:CPDBackend}, p, X) = + differential_action!(out, cpd_model(model), p, X) +adjoint_action( + model::JoinModel{<:AbstractFloat,<:CPDBackend}, + p, + a::AbstractVector; + kwargs..., +) = + adjoint_action(cpd_model(model), p, a; kwargs...) supports_rgrad(model::JoinModel{<:AbstractFloat,<:CPDBackend}) = supports_rgrad(cpd_model(model)) rgrad(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = rgrad(cpd_model(model), p) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 863159e..f49b323 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -792,3 +792,51 @@ function _join_residual!(backend::Union{JoinBackend,BTDBackend}, p) backend.work_residual .-= backend.target_flat return backend.work_residual end + +function residual(model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, p) + return copy(_join_residual!(model.backend, p)) +end + +function differential_action!( + out::AbstractVector{T}, + model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, + p, + X, +) where {T<:AbstractFloat} + backend = model.backend + parts = point_parts(p) + xparts = point_parts(X) + _check_parts_len(parts, backend.r, "differential_action!") + _check_parts_len(xparts, backend.r, "differential_action!") + length(out) == length(backend.target_flat) || throw( + DimensionMismatch( + "differential_action! output length $(length(out)) != ambient length $(length(backend.target_flat)).", + ), + ) + fill!(out, zero(T)) + @inbounds for k = 1:backend.r + component_ambient_pushforward!( + backend.component_bufs[k], + _backend_component(backend, k), + parts[k], + xparts[k], + ) + out .+= backend.component_bufs[k] + end + return out +end + +function adjoint_action( + model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, + p, + a::AbstractVector; + kwargs..., +) + backend = model.backend + length(a) == length(backend.target_flat) || throw( + DimensionMismatch( + "adjoint_action expected ambient vector of length $(length(backend.target_flat)), got $(length(a)).", + ), + ) + return _join_basis_project(_backend_components(backend), p, a) +end diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 5f00204..01db871 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -36,164 +36,91 @@ solver_symbol(::LMSolver) = :lm one(T) / sqrt(T(normA2)) : one(T) end -function _lm_raw_residual_vector( - model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, - p, -) - return copy(_join_residual!(model.backend, p)) -end - -function _join_component_jacobian_block!( - J::AbstractMatrix{T}, - next_col::Int, - component, - p_component, - work_vec::AbstractVector{T}; - basis = ManifoldsBase.DefaultOrthonormalBasis(), -) where {T<:AbstractFloat} - d_component = component_tangent_dimension(component, p_component) - @inbounds for j = 1:d_component - Xj = component_basis_vector(component, p_component, j; basis) - component_ambient_pushforward!(work_vec, component, p_component, Xj) - J[:, next_col] .= work_vec - next_col += 1 - end - return next_col -end +_lm_raw_residual_vector(model::AbstractDecompositionModel, p) = residual(model, p) function _lm_raw_jacobian_matrix( - model::JoinModel{<:AbstractFloat,<:Union{JoinBackend,BTDBackend}}, + model::AbstractDecompositionModel, M, p; basis = ManifoldsBase.DefaultOrthonormalBasis(), ) T = _scalar_eltype(p) ambient_dim = length(tensor(model)) - backend = model.backend - components = _backend_components(backend) - parts = point_parts(p) - _check_parts_len(parts, backend.r, "_lm_raw_jacobian_matrix") - d = sum(component_tangent_dimension(components[k], parts[k]) for k = 1:backend.r) - d == manifold_dimension(M) || throw( - DimensionMismatch( - "JoinModel Jacobian assembly expected tangent dimension $d from components but manifold reports $(manifold_dimension(M)).", - ), - ) - J = Matrix{T}(undef, ambient_dim, d) - next_col = 1 - @inbounds for k = 1:backend.r - next_col = _join_component_jacobian_block!( - J, - next_col, - components[k], - parts[k], - backend.component_bufs[k]; - basis, - ) - end - next_col == d + 1 || throw( - DimensionMismatch( - "JoinModel Jacobian assembly filled $(next_col - 1) columns but expected $d.", - ), - ) - return J -end - -function _lm_raw_residual_vector(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) - return _lm_raw_residual_vector(cpd_model(model), p) -end - -function _lm_raw_jacobian_matrix( - model::JoinModel{<:AbstractFloat,<:CPDBackend}, - M, - p; - basis = ManifoldsBase.DefaultOrthonormalBasis(), -) - return _lm_raw_jacobian_matrix(cpd_model(model), M, p; basis) -end - -function _lm_raw_residual_vector(model::Rank1CPDModel{T}, p) where {T<:AbstractFloat} - return vec(embed_point(model, p)) .- vec(model.A) -end - -function _lm_raw_jacobian_matrix( - model::Rank1CPDModel{T}, - M, - p; - basis = ManifoldsBase.DefaultOrthonormalBasis(), -) where {T<:AbstractFloat} - ambient_dim = length(model.A) d = manifold_dimension(M) J = Matrix{T}(undef, ambient_dim, d) coeff = zeros(T, d) - embedding = _cp_parameterization(model) column = Vector{T}(undef, ambient_dim) @inbounds for j = 1:d fill!(coeff, zero(T)) coeff[j] = one(T) Xj = ManifoldsBase.get_vector(M, p, coeff, basis) - λ, U, λ̇, U̇ = _cp_rank1_decode_tangent_factors(embedding, model.dims, p, Xj) - _cp_rank1_tangent_tensorvec!(column, λ, U, λ̇, U̇) + differential_action!(column, model, p, Xj) J[:, j] .= column end return J end -function _lm_raw_residual_vector(model::RankRCPDModel{T}, p) where {T<:AbstractFloat} - return vec(embed_point(model, p)) .- vec(model.A) +function _lm_residual_function( + model::AbstractDecompositionModel, + ::Type{T}, + normA2, + normalized_objective::Bool, +) where {T<:AbstractFloat} + scale = _lm_scaling_factor(T, normA2, normalized_objective) + return (M, p) -> scale .* _lm_raw_residual_vector(model, p) end -function _lm_raw_jacobian_matrix( - model::RankRCPDModel{T}, - M, - p; +function _lm_jacobian_function( + model::AbstractDecompositionModel, + ::Type{T}, + normA2, + normalized_objective::Bool; basis = ManifoldsBase.DefaultOrthonormalBasis(), ) where {T<:AbstractFloat} - ambient_dim = length(model.A) - d = manifold_dimension(M) - J = Matrix{T}(undef, ambient_dim, d) - coeff = zeros(T, d) - column = Vector{T}(undef, ambient_dim) - embedding = _cp_parameterization(model) - @inbounds for j = 1:d - fill!(coeff, zero(T)) - coeff[j] = one(T) - Xj = ManifoldsBase.get_vector(M, p, coeff, basis) - λ, U, λ̇, U̇ = - _cp_rankr_decode_tangent_factors(embedding, model.dims, model.r, p, Xj) - _cp_rankr_tangent_tensorvec!(column, λ, U, λ̇, U̇) - J[:, j] .= column - end - return J -end - -function _lm_raw_residual_vector(model::AbstractDecompositionModel, p) - throw(ArgumentError("LMSolver residual is not implemented for model $(typeof(model)).")) + scale = _lm_scaling_factor(T, normA2, normalized_objective) + return (M, p) -> scale .* _lm_raw_jacobian_matrix(model, M, p; basis) end -function _lm_raw_jacobian_matrix(model::AbstractDecompositionModel, M, p; basis) - throw(ArgumentError("LMSolver Jacobian is not implemented for model $(typeof(model)).")) +function _lm_differential_action_function( + model::AbstractDecompositionModel, + ::Type{T}, + normA2, + normalized_objective::Bool, +) where {T<:AbstractFloat} + scale = _lm_scaling_factor(T, normA2, normalized_objective) + return (M, p, X) -> scale .* differential_action(model, p, X) end -function _lm_residual_function( +function _lm_adjoint_action_function( model::AbstractDecompositionModel, ::Type{T}, normA2, normalized_objective::Bool, ) where {T<:AbstractFloat} scale = _lm_scaling_factor(T, normA2, normalized_objective) - return (M, p) -> scale .* _lm_raw_residual_vector(model, p) + return (M, p, a) -> adjoint_action(model, p, scale .* a) end -function _lm_jacobian_function( +function _lm_vector_differential_function( model::AbstractDecompositionModel, ::Type{T}, normA2, - normalized_objective::Bool; - basis = ManifoldsBase.DefaultOrthonormalBasis(), + normalized_objective::Bool, ) where {T<:AbstractFloat} - scale = _lm_scaling_factor(T, normA2, normalized_objective) - return (M, p) -> scale .* _lm_raw_jacobian_matrix(model, M, p; basis) + ambient_dim = length(tensor(model)) + residual_f = _lm_residual_function(model, T, normA2, normalized_objective) + differential_f = _lm_differential_action_function(model, T, normA2, normalized_objective) + adjoint_f = _lm_adjoint_action_function(model, T, normA2, normalized_objective) + return Manopt.VectorDifferentialFunction( + residual_f, + differential_f, + adjoint_f, + ambient_dim; + evaluation = Manopt.AllocatingEvaluation(), + function_type = Manopt.FunctionVectorialType(), + jacobian_type = Manopt.FunctionVectorialType(), + adjoint_jacobian_type = Manopt.FunctionVectorialType(), + ) end function solve_lm( @@ -233,18 +160,31 @@ function solve_lm( ) p0_local = setup.p0 T = setup.T - basis = ManifoldsBase.DefaultOrthonormalBasis() - residual = _lm_residual_function(model, T, normA2, setup.uses_relative_objective) - jacobian = _lm_jacobian_function(model, T, normA2, setup.uses_relative_objective; basis) - initial_residual_values = residual(M, p0_local) - initial_jacobian_f = jacobian(M, p0_local) - tangent_space = TangentSpace(M, p0_local) - lm_subsolver_state = Manopt.CoordinatesNormalSystemState( - tangent_space, - zero_vector(M, p0_local); - evaluation = Manopt.InplaceEvaluation(), - linsolve = linear_subsolver, - basis = basis, + vdf = _lm_vector_differential_function(model, T, normA2, setup.uses_relative_objective) + initial_residual_values = residual(model, p0_local) + scale = _lm_scaling_factor(T, normA2, setup.uses_relative_objective) + if scale != one(T) + initial_residual_values .*= scale + end + nlso = Manopt.ManifoldNonlinearLeastSquaresObjective( + vdf, + Manopt.ComponentwiseRobustifierFunction(Manopt.IdentityRobustifier()), + ) + initial_jacobian_matrices = fill(nothing, 1) + sub_objective = Manopt.construct_lm_subobjective( + false, + nlso, + damping_term_min, + 1.0e-6, + :Strict, + initial_residual_values, + initial_jacobian_matrices, + ) + sub_state = Manopt.ConjugateResidualState( + TangentSpace(M, p0_local), + sub_objective; + stopping_criterion = StopAfterIteration(max(4 * manifold_dimension(M), 50)) | + StopWhenGradientNormLess(T(1e-14)), ) retraction_method = _solver_retraction_method(M, p0_local) stopping = StopWhenAny( @@ -272,16 +212,11 @@ function solve_lm( ) state = Manopt.LevenbergMarquardt( M, - residual, - jacobian, + nlso, p0_local; - evaluation = Manopt.AllocatingEvaluation(), - function_type = Manopt.FunctionVectorialType(), - jacobian_type = Manopt.CoefficientVectorialType(basis), retraction_method = retraction_method, stopping_criterion = stopping, initial_residual_values = initial_residual_values, - initial_jacobian_matrices = [initial_jacobian_f], candidate_acceptance_threshold = η, damping_increase_factor = β, damping_increase_threshold = η, @@ -289,8 +224,9 @@ function solve_lm( damping_reduction_factor = inv(T(β)), damping_term_min = damping_term_min, initial_damping_term = damping_term_min, - use_unified_basis = true, - sub_state = lm_subsolver_state, + use_unified_basis = false, + sub_objective = sub_objective, + sub_state = sub_state, debug = callbacks.debug_actions, return_state = true, ) @@ -316,6 +252,9 @@ function solve_lm( damping_term_min = Float64(damping_term_min), β = Float64(β), expect_zero_residual = expect_zero_residual, + uses_operator_jacobian = true, + uses_direct_adjoint_action = true, + uses_coordinate_linear_solver = false, uses_vector_transport = !isnothing(vector_transport_method), ), ) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 9729d9e..3ec4836 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -259,6 +259,38 @@ end end end +@testset "Operator interface matches Jacobian and gradient" begin + A = randn(5, 4, 3) + cases = ( + ("generic_join", JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), A)), + ("cp_canonical", JoinModel(A, 2; geometry = :canonical)), + ("cp_softplus", JoinModel(abs.(A), 2; geometry = :softplus_metric, nonnegative = true)), + ) + for (label, model) in cases + M = TensorKitchen.manifold(model) + p = TensorKitchen._solver_point( + M, + TensorKitchen.initial_point(model, :random; verbose = false), + ) + basis = ManifoldsBase.DefaultOrthonormalBasis() + r = TensorKitchen.residual(model, p) + J = TensorKitchen._lm_raw_jacobian_matrix(model, M, p; basis) + @testset "$label" begin + @test r ≈ TensorKitchen._lm_raw_residual_vector(model, p) + d = manifold_dimension(M) + for j = 1:min(d, 3) + coeff = zeros(Float64, d) + coeff[j] = 1.0 + Xj = ManifoldsBase.get_vector(M, p, coeff, basis) + @test TensorKitchen.differential_action(model, p, Xj) ≈ J[:, j] + end + g_adj = TensorKitchen.adjoint_action(model, p, r; basis) + g_model = TensorKitchen.rgrad(model, p) + @test norm(M, p, g_adj - g_model) ≤ 1e-7 * max(1.0, norm(M, p, g_model)) + end + end +end + function _reference_join_jacobian_from_product_basis(model, M, p; basis) backend = model.backend parts = TensorKitchen.point_parts(p) From b532e42ae427f6e3ce80666e59eb27f976de9aa1 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 3 Jul 2026 12:24:50 +0200 Subject: [PATCH 33/37] juliaFormatter --- src/cpd/model/rank1.jl | 10 ++++++++-- src/cpd/model/rankr.jl | 16 +++++++++++++--- src/solvers/lm.jl | 3 ++- test/basic_tests.jl | 10 ++++++++-- 4 files changed, 31 insertions(+), 8 deletions(-) diff --git a/src/cpd/model/rank1.jl b/src/cpd/model/rank1.jl index 27d2ca2..ff642d5 100644 --- a/src/cpd/model/rank1.jl +++ b/src/cpd/model/rank1.jl @@ -82,11 +82,17 @@ function differential_action!( "differential_action! output length $(length(out)) != ambient length $(length(model.A)).", ), ) - λ, U, λ̇, U̇ = _cp_rank1_decode_tangent_factors(_cp_parameterization(model), model.dims, p, X) + λ, U, λ̇, U̇ = + _cp_rank1_decode_tangent_factors(_cp_parameterization(model), model.dims, p, X) return _cp_rank1_tangent_tensorvec!(out, λ, U, λ̇, U̇) end -function adjoint_action(model::Rank1CPDModel{T,N}, p, a::AbstractVector; kwargs...) where {T<:AbstractFloat,N} +function adjoint_action( + model::Rank1CPDModel{T,N}, + p, + a::AbstractVector; + kwargs..., +) where {T<:AbstractFloat,N} length(a) == length(model.A) || throw( DimensionMismatch( "adjoint_action expected ambient vector of length $(length(model.A)) for $(typeof(model)), got $(length(a)).", diff --git a/src/cpd/model/rankr.jl b/src/cpd/model/rankr.jl index f98d40f..363f1e0 100644 --- a/src/cpd/model/rankr.jl +++ b/src/cpd/model/rankr.jl @@ -124,12 +124,22 @@ function differential_action!( "differential_action! output length $(length(out)) != ambient length $(length(model.A)).", ), ) - λ, U, λ̇, U̇ = - _cp_rankr_decode_tangent_factors(_cp_parameterization(model), model.dims, model.r, p, X) + λ, U, λ̇, U̇ = _cp_rankr_decode_tangent_factors( + _cp_parameterization(model), + model.dims, + model.r, + p, + X, + ) return _cp_rankr_tangent_tensorvec!(out, λ, U, λ̇, U̇) end -function adjoint_action(model::RankRCPDModel{T,N}, p, a::AbstractVector; kwargs...) where {T<:AbstractFloat,N} +function adjoint_action( + model::RankRCPDModel{T,N}, + p, + a::AbstractVector; + kwargs..., +) where {T<:AbstractFloat,N} length(a) == length(model.A) || throw( DimensionMismatch( "adjoint_action expected ambient vector of length $(length(model.A)) for $(typeof(model)), got $(length(a)).", diff --git a/src/solvers/lm.jl b/src/solvers/lm.jl index 8432042..6f6df41 100644 --- a/src/solvers/lm.jl +++ b/src/solvers/lm.jl @@ -98,7 +98,8 @@ function _lm_vector_differential_function( ) where {T<:AbstractFloat} ambient_dim = length(tensor(model)) residual_f = _lm_residual_function(model, T, normA2, normalized_objective) - differential_f = _lm_differential_action_function(model, T, normA2, normalized_objective) + differential_f = + _lm_differential_action_function(model, T, normA2, normalized_objective) adjoint_f = _lm_adjoint_action_function(model, T, normA2, normalized_objective) return Manopt.VectorDifferentialFunction( residual_f, diff --git a/test/basic_tests.jl b/test/basic_tests.jl index e2c2ab2..d83ec16 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -262,9 +262,15 @@ end @testset "Operator interface matches Jacobian and gradient" begin A = randn(5, 4, 3) cases = ( - ("generic_join", JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), A)), + ( + "generic_join", + JoinModel((Manifolds.Segre((5, 4, 3)), Manifolds.Segre((5, 4, 3))), A), + ), ("cp_canonical", JoinModel(A, 2; geometry = :canonical)), - ("cp_softplus", JoinModel(abs.(A), 2; geometry = :softplus_metric, nonnegative = true)), + ( + "cp_softplus", + JoinModel(abs.(A), 2; geometry = :softplus_metric, nonnegative = true), + ), ) for (label, model) in cases M = TensorKitchen.manifold(model) From cc9f9b5fb874df50bf21af66c4fa1457c4d2776d Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 3 Jul 2026 13:52:37 +0200 Subject: [PATCH 34/37] fixed the version of JuliaFormatter for format_check --- .github/workflows/format_check.yml | 3 +-- Project.toml | 2 +- src/join/cpd_backend.jl | 11 +++++++---- test/basic_tests.jl | 3 +-- 4 files changed, 10 insertions(+), 9 deletions(-) diff --git a/.github/workflows/format_check.yml b/.github/workflows/format_check.yml index 2b61af9..aaefa39 100644 --- a/.github/workflows/format_check.yml +++ b/.github/workflows/format_check.yml @@ -18,8 +18,7 @@ jobs: - uses: actions/checkout@v6 - name: Install JuliaFormatter and format run: | - julia --project=. -e 'using Pkg; Pkg.instantiate()' - julia --project=. -e 'using JuliaFormatter; format(["./src", "./test"], verbose=true)' + julia -e 'using Pkg; Pkg.activate(temp=true); Pkg.add(PackageSpec(name="JuliaFormatter", version="2.10.1")); using JuliaFormatter; format(["./src", "./test"], verbose=true)' - name: Format check run: | julia -e ' diff --git a/Project.toml b/Project.toml index 3f25b5f..6e1c292 100644 --- a/Project.toml +++ b/Project.toml @@ -14,7 +14,7 @@ RecursiveArrayTools = "731186ca-8d62-57ce-b412-fbd966d074cd" TensorOperations = "6aa20fa7-93e2-5fca-9bc0-fbd0db3c71a2" [compat] -JuliaFormatter = "2.8.5" +JuliaFormatter = "=2.10.1" Manifolds = "0.11.28" ManifoldsBase = "2.3.5" Manopt = "0.6" diff --git a/src/join/cpd_backend.jl b/src/join/cpd_backend.jl index 7d87b0e..ab3d7e3 100644 --- a/src/join/cpd_backend.jl +++ b/src/join/cpd_backend.jl @@ -60,15 +60,18 @@ initial_point( cost(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = cost(cpd_model(model), p) egrad(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = egrad(cpd_model(model), p) residual(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = residual(cpd_model(model), p) -differential_action!(out::AbstractVector, model::JoinModel{<:AbstractFloat,<:CPDBackend}, p, X) = - differential_action!(out, cpd_model(model), p, X) +differential_action!( + out::AbstractVector, + model::JoinModel{<:AbstractFloat,<:CPDBackend}, + p, + X, +) = differential_action!(out, cpd_model(model), p, X) adjoint_action( model::JoinModel{<:AbstractFloat,<:CPDBackend}, p, a::AbstractVector; kwargs..., -) = - adjoint_action(cpd_model(model), p, a; kwargs...) +) = adjoint_action(cpd_model(model), p, a; kwargs...) supports_rgrad(model::JoinModel{<:AbstractFloat,<:CPDBackend}) = supports_rgrad(cpd_model(model)) rgrad(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) = rgrad(cpd_model(model), p) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index d83ec16..a537745 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -553,8 +553,7 @@ end @test all(U_sp[m] ≈ TensorKitchen._softplus_value.(Ũ[m]) for m in eachindex(U_sp)) @test λ̇_sp ≈ TensorKitchen._softplus_derivative.(λ̃) .* λ̇̃ @test all( - U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for - m in eachindex(U̇_sp) + U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(U̇_sp) ) model_sp = TensorKitchen.RankRCPDModel( From 95ea4ecb7266170f126c06c49f3348f5f3920631 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 3 Jul 2026 15:28:56 +0200 Subject: [PATCH 35/37] rewrite the test for BTD LM to bypass --- test/basic_tests.jl | 75 ++++++++++++++++----------------------------- 1 file changed, 26 insertions(+), 49 deletions(-) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index a537745..afac95c 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -604,65 +604,42 @@ end @test res_approx.solver == :lm end -@testset "BTD accepts LMSolver on nested Tucker layouts" begin +@testset "BTD exposes LM residual/Jacobian hooks on nested Tucker layouts" begin A = randn(7, 6, 5) ranks = (2, 2, 2) manifolds = TensorKitchen._as_join_manifold_tuple(TuckerJoin(size(A), ranks, 2)) backend = TensorKitchen._sum_backend_instance(TensorKitchen.BTDBackend, manifolds, A) model = TensorKitchen.JoinModel{Float64,typeof(backend)}(backend) p0 = TensorKitchen.initial_point(model, :random; verbose = false) + M = TensorKitchen.manifold(model) + p0_solver = TensorKitchen._solver_point(M, p0) + basis = ManifoldsBase.DefaultOrthonormalBasis() @test p0 isa ArrayPartition @test TensorKitchen.point_parts(p0)[1] isa Manifolds.TuckerPoint - - low = solve( - LMSolver(), - model; - p0, - maxiter = 2, - tol = 1e-6, - verbose = false, - return_stats = true, - ) - low_parts = TensorKitchen.point_parts(low.point) - @test low.solver == :lm - @test low.point isa ArrayPartition - @test length(low_parts) == 2 - @test low_parts[1] isa Manifolds.TuckerPoint - - res_btd = btd( - A, - 2, - ranks; - solver = :lm, - warm_rel_error_gate = nothing, - maxiter = 2, - tol = 1e-6, - verbose = false, - ) - @test res_btd isa BTDResult - @test res_btd.solver == :lm - @test length(res_btd.components) == 2 - @test !get(res_btd.solver_info, :btd_skipped_manifold_polish, false) - - res_btd_alswarm_lm = btd( - A, - 2, - ranks; - solver = :lm, - init = :alswarm, - warm_init = BTDHOSVDMultistartInit(2; screening_steps = 0, block_maxiter = 1), - warm_steps = 1, - warm_block_maxiter = 1, - warm_rel_error_gate = nothing, - maxiter = 2, - tol = 1e-6, - verbose = false, + @test p0_solver isa ArrayPartition + @test TensorKitchen.point_parts(p0_solver)[1] isa Manifolds.TuckerPoint + + residual0 = TensorKitchen._lm_raw_residual_vector(model, p0_solver) + J0 = TensorKitchen._lm_raw_jacobian_matrix(model, M, p0_solver; basis = basis) + @test length(residual0) == length(A) + @test size(J0) == (length(A), manifold_dimension(M)) + @test all(isfinite, residual0) + @test all(isfinite, J0) + + coeff = zeros(Float64, manifold_dimension(M)) + coeff[1] = 1.0 + X = ManifoldsBase.get_vector(M, p0_solver, coeff, basis) + JX = TensorKitchen.differential_action(model, p0_solver, X) + ambient = randn(size(A)) + lhs = dot(JX, vec(ambient)) + rhs = ManifoldsBase.inner( + M, + p0_solver, + X, + TensorKitchen.adjoint_action(model, p0_solver, vec(ambient)), ) - @test res_btd_alswarm_lm isa BTDResult - @test res_btd_alswarm_lm.solver == :lm - @test hasproperty(res_btd_alswarm_lm.solver_info, :btd_als_warm_start_iters) - @test res_btd_alswarm_lm.solver_info.btd_als_warm_start_requested_solver == :lm + @test isapprox(lhs, rhs; atol = 1e-8, rtol = 1e-8) end # ========================================================================= From 0e26ac84fb5361b29c81f0231f1f68b24d603e27 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 3 Jul 2026 15:52:43 +0200 Subject: [PATCH 36/37] Temporarily disable BTD LM frontend --- src/api/btd.jl | 15 ++++++++++++++- test/basic_tests.jl | 28 +++++++++++++++++++++++++++- 2 files changed, 41 insertions(+), 2 deletions(-) diff --git a/src/api/btd.jl b/src/api/btd.jl index 1127881..f2307af 100644 --- a/src/api/btd.jl +++ b/src/api/btd.jl @@ -63,6 +63,15 @@ _btd_uses_warm_start(::AbstractSolver, ::BTDALSWarmStartInit) = true _btd_should_polish(::ALSSolver, ::Integer) = false _btd_should_polish(::AbstractSolver, polish_n::Integer) = polish_n > 0 +function _reject_unsupported_btd_solver(solver_obj) + solver_obj isa LMSolver || return nothing + throw( + ArgumentError( + "BTD currently does not support LM refinement because the required Manopt operator path is not yet available for nested Tucker layouts. Use :rgd, :rcg, :lbfgs, :als, or :btd_tsd instead.", + ), + ) +end + function _btd_warm_start_result( model::JoinModel{T,<:BTDBackend}, backend::BTDBackend, @@ -159,7 +168,10 @@ refines it. Returns a [`BTDResult`](@ref). - `:als`: Alternating least squares. - `:rcg`: Riemannian conjugate gradient. - `:lbfgs`: Limited-memory quasi-Newton refinement. - - `:lm`: Levenberg-Marquardt refinement. + - `:btd_tsd`: Blockwise tangent-subspace descent for BTD. + +`solver = :lm` is currently not supported for BTD because the required Manopt +LM operator path is not yet available for nested Tucker layouts. ## Extended Options @@ -234,6 +246,7 @@ function btd( kwargs..., ) where {T<:AbstractFloat,N} solver_obj = _solver_object(solver, stepsize; kwargs...) + _reject_unsupported_btd_solver(solver_obj) solver_sym = _btd_solver_symbol(solver_obj) init_resolved = _resolve_btd_init(init, solver_obj) init_eff = diff --git a/test/basic_tests.jl b/test/basic_tests.jl index afac95c..1cef4a5 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -553,7 +553,8 @@ end @test all(U_sp[m] ≈ TensorKitchen._softplus_value.(Ũ[m]) for m in eachindex(U_sp)) @test λ̇_sp ≈ TensorKitchen._softplus_derivative.(λ̃) .* λ̇̃ @test all( - U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(U̇_sp) + U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for + m in eachindex(U̇_sp) ) model_sp = TensorKitchen.RankRCPDModel( @@ -642,6 +643,31 @@ end @test isapprox(lhs, rhs; atol = 1e-8, rtol = 1e-8) end +@testset "BTD rejects LMSolver until nested Tucker LM support lands" begin + A = randn(7, 6, 5) + ranks = (2, 2, 2) + + @test_throws ArgumentError btd( + A, + 2, + ranks; + solver = :lm, + maxiter = 2, + tol = 1e-6, + verbose = false, + ) + + @test_throws ArgumentError btd( + A, + 2, + ranks; + solver = LMSolver(), + maxiter = 2, + tol = 1e-6, + verbose = false, + ) +end + # ========================================================================= # cpd/cp_rank.jl (cost/egrad functions) # ========================================================================= From a7286b52595060f4f71e86627bb036222ec98429 Mon Sep 17 00:00:00 2001 From: Se Eun Choi Date: Fri, 3 Jul 2026 16:03:50 +0200 Subject: [PATCH 37/37] Apply JuliaFormatter output --- test/basic_tests.jl | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 1cef4a5..00056e6 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -553,8 +553,7 @@ end @test all(U_sp[m] ≈ TensorKitchen._softplus_value.(Ũ[m]) for m in eachindex(U_sp)) @test λ̇_sp ≈ TensorKitchen._softplus_derivative.(λ̃) .* λ̇̃ @test all( - U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for - m in eachindex(U̇_sp) + U̇_sp[m] ≈ TensorKitchen._softplus_derivative.(Ũ[m]) .* U̇̃[m] for m in eachindex(U̇_sp) ) model_sp = TensorKitchen.RankRCPDModel(