Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"
Expand Down
3 changes: 1 addition & 2 deletions src/api/approx.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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."))
Expand Down
4 changes: 2 additions & 2 deletions src/api/btd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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,
(
Expand Down Expand Up @@ -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,
Expand Down
12 changes: 6 additions & 6 deletions src/api/cpd.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down Expand Up @@ -579,13 +579,14 @@ end
function _validate_cpd_solver_supported(solver::AbstractSolver)
throw(
ArgumentError(
"Unsupported CPD solver $(typeof(solver)). Use :als, :rgd, :rgd_fixed, or :rcg.",
"Unsupported CPD solver $(typeof(solver)). Use :als, :rgd, :rgd_fixed, :rcg, or :lbfgs.",
),
)
end

_validate_cpd_solver_supported(::Union{ALSSolver,RGDSolver,RGDFixedSolver,RCGSolver}) =
nothing
_validate_cpd_solver_supported(
::Union{ALSSolver,RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver},
) = nothing

function _validate_cpd_solver_options(
solver::AbstractSolver,
Expand Down Expand Up @@ -752,8 +753,6 @@ function _cpd_manifold_grad_tol(
solver::Union{RGDSolver,RGDFixedSolver,RCGSolver,LBFGSSolver},
tol::Real,
)
inner = cpd_model(model)
inner.nonnegative || return nothing
return tol
end

Expand Down Expand Up @@ -992,6 +991,7 @@ If `r` is omitted, uses the smallest tensor mode as a heuristic rank.
- `rgd` (default): Riemannian gradient descent
- `rgd_fixed`: Riemannian gradient descent with fixed step size
- `rcg`: Riemannian conjugate gradient
- `lbfgs`: Limited-memory Riemannian quasi-Newton
- `als`: Alternating Least Squares

## Extended Options
Expand Down
1 change: 1 addition & 0 deletions src/backend.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
67 changes: 56 additions & 11 deletions src/core/progress.jl
Original file line number Diff line number Diff line change
Expand Up @@ -176,6 +176,40 @@ 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 &&
_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

Expand All @@ -192,14 +226,15 @@ 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)
_sync_rendered!(progress)
end
return nothing
end
Expand All @@ -215,24 +250,22 @@ 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))]
else
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
Expand All @@ -247,16 +280,28 @@ 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)
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
Expand Down
4 changes: 1 addition & 3 deletions src/core/types.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)

"""
Expand Down
2 changes: 1 addition & 1 deletion src/cpd/core/cpd_init.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion src/join/cpd_backend.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
32 changes: 11 additions & 21 deletions src/results/conversion.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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))."))
Expand Down
17 changes: 4 additions & 13 deletions src/results/rel_error.jl
Original file line number Diff line number Diff line change
Expand Up @@ -18,19 +18,10 @@ 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

Expand Down
Loading
Loading