@@ -22,6 +22,104 @@ function _merge_res_solver_info(res, patch::NamedTuple)
2222 )
2323end
2424
25+ mutable struct _CPDComponentTraceRecorder{M}
26+ model:: M
27+ previous:: Any
28+ previous_cost:: Float64
29+ iterations:: Vector{Int}
30+ cost_history:: Vector{Float64}
31+ cost_rel_change_history:: Vector{Float64}
32+ max_component_delta_history:: Vector{Float64}
33+ component_delta_history:: Vector{Vector{Float64}}
34+ end
35+
36+ function _CPDComponentTraceRecorder (model)
37+ return _CPDComponentTraceRecorder (
38+ model,
39+ nothing ,
40+ NaN ,
41+ Int[],
42+ Float64[],
43+ Float64[],
44+ Float64[],
45+ Vector{Float64}[],
46+ )
47+ end
48+
49+ function _rankone_norm2 (λ, U, k:: Int )
50+ val = abs2 (λ[k])
51+ @inbounds for m = 1 : length (U)
52+ val *= sum (abs2, @view U[m][:, k])
53+ end
54+ return Float64 (val)
55+ end
56+
57+ function _rankone_inner (λa, Ua, λb, Ub, k:: Int )
58+ val = λa[k] * λb[k]
59+ @inbounds for m = 1 : length (Ua)
60+ val *= dot (@view (Ua[m][:, k]), @view (Ub[m][:, k]))
61+ end
62+ return Float64 (val)
63+ end
64+
65+ function _cpd_component_deltas (prev:: CPDPoint , curr:: CPDPoint )
66+ λ_prev = lambda (prev)
67+ U_prev = factors (prev)
68+ λ_curr = lambda (curr)
69+ U_curr = factors (curr)
70+ r = length (λ_curr)
71+ deltas = Vector {Float64} (undef, r)
72+ @inbounds for k = 1 : r
73+ n_prev = _rankone_norm2 (λ_prev, U_prev, k)
74+ n_curr = _rankone_norm2 (λ_curr, U_curr, k)
75+ cross = _rankone_inner (λ_prev, U_prev, λ_curr, U_curr, k)
76+ delta = sqrt (max (n_prev + n_curr - 2 * cross, 0.0 ))
77+ deltas[k] = delta / max (sqrt (max (n_prev, 0.0 )), 1.0 )
78+ end
79+ return deltas
80+ end
81+
82+ function _record_cpd_component_trace! (rec:: _CPDComponentTraceRecorder , p, iter:: Int )
83+ q = cpd_point (rec. model, p)
84+ cost_val = Float64 (cost (rec. model, p))
85+ if rec. previous != = nothing
86+ deltas = _cpd_component_deltas (rec. previous, q)
87+ rel_change = abs (rec. previous_cost - cost_val) / max (abs (rec. previous_cost), 1.0 )
88+ push! (rec. iterations, iter)
89+ push! (rec. cost_history, cost_val)
90+ push! (rec. cost_rel_change_history, rel_change)
91+ push! (rec. max_component_delta_history, maximum (deltas))
92+ push! (rec. component_delta_history, deltas)
93+ end
94+ rec. previous = q
95+ rec. previous_cost = cost_val
96+ return nothing
97+ end
98+
99+ function _cpd_component_trace_callback (rec:: _CPDComponentTraceRecorder )
100+ return function (problem, state, k)
101+ p = try
102+ Manopt. get_iterate (state)
103+ catch
104+ return nothing
105+ end
106+ _record_cpd_component_trace! (rec, p, Int (k))
107+ return nothing
108+ end
109+ end
110+
111+ function _cpd_component_trace_info (rec:: _CPDComponentTraceRecorder )
112+ return (
113+ component_trace_iterations = rec. iterations,
114+ component_trace_cost_history = rec. cost_history,
115+ component_trace_cost_rel_change_history = rec. cost_rel_change_history,
116+ component_trace_max_delta_history = rec. max_component_delta_history,
117+ component_trace_delta_history = rec. component_delta_history,
118+ component_trace_final_max_delta = isempty (rec. max_component_delta_history) ? NaN :
119+ rec. max_component_delta_history[end ],
120+ )
121+ end
122+
25123function _pack_cpd_explicit_p0 (model, p0)
26124 p0 isa CPDPoint && return pack_cpd_point (model, p0)
27125 p0 isa CPDResult && return pack_cpd_point (model, cpd_point (p0))
@@ -159,8 +257,12 @@ function _run_cpd_solver(
159257 verbose:: Bool ,
160258 vector_transport_method,
161259 pullback_eps,
260+ component_trace,
162261 kwargs... ,
163262)
263+ trace_recorder = component_trace ? _CPDComponentTraceRecorder (model) : nothing
264+ iteration_callbacks =
265+ isnothing (trace_recorder) ? () : (_cpd_component_trace_callback (trace_recorder),)
164266 p_solve = if init_eff isa ALSWarmStartInit && isnothing (p0) && ! (solver isa ALSSolver)
165267 _cpd_als_warm_then_pack (
166268 model,
@@ -175,7 +277,7 @@ function _run_cpd_solver(
175277 _pack_cpd_explicit_p0 (model, p0)
176278 end
177279
178- return _solve_model (
280+ result = _solve_model (
179281 model;
180282 init = init_eff,
181283 p0 = p_solve,
@@ -188,8 +290,11 @@ function _run_cpd_solver(
188290 verbose,
189291 refinement_verbose = verbose,
190292 vector_transport_method,
293+ iteration_callbacks,
191294 kwargs... ,
192295 )
296+ return isnothing (trace_recorder) ? result :
297+ _merge_res_solver_info (result, _cpd_component_trace_info (trace_recorder))
193298end
194299
195300function _cpd_impl (
@@ -212,6 +317,7 @@ function _cpd_impl(
212317 verbose,
213318 vector_transport_method,
214319 pullback_eps = 1e-8 ,
320+ component_trace:: Bool = false ,
215321 kwargs... ,
216322) where {T<: AbstractFloat ,N}
217323 haskey (kwargs, :softplus_beta ) && throw (
@@ -254,6 +360,9 @@ function _cpd_impl(
254360 throw (ArgumentError (" geometry=$geometry_eff requires nonnegative=true." ))
255361 end
256362 if solver_obj isa ALSSolver
363+ component_trace && throw (
364+ ArgumentError (" component_trace=true is only supported for manifold solvers." ),
365+ )
257366 geometry_eff == :canonical || throw (
258367 ArgumentError (
259368 " solver=:als does not use manifold geometry. Use geometry=:canonical." ,
@@ -292,6 +401,7 @@ function _cpd_impl(
292401 verbose,
293402 vector_transport_method,
294403 pullback_eps = pullback_eps_eff,
404+ component_trace,
295405 nonnegative,
296406 kwargs... ,
297407 )
@@ -361,6 +471,9 @@ If `r` is omitted, uses the smallest tensor mode as a heuristic rank.
361471* `verbose = true`: Enables progress output.
362472* `nonnegative::Bool = false`: Nonnegative CPD option to be selected by the user. (same as `nncpd`)
363473* `pullback_eps = 1e-8`: Regularization parameter for pullback-style nonnegative geometries.
474+ * `component_trace = false`: For manifold solvers, records per-iteration movement of
475+ each CP rank-one term in `solver_info`. Use this to diagnose whether a flat cost
476+ means the rank-one terms are also stuck.
364477
365478## Notes
366479* `solver = :als` does not use manifold geometry. In that case:
@@ -403,6 +516,7 @@ function cpd(
403516 verbose = true ,
404517 vector_transport_method = nothing ,
405518 pullback_eps = 1e-8 ,
519+ component_trace:: Bool = false ,
406520 kwargs... ,
407521) where {T<: AbstractFloat ,N}
408522 if nonnegative
@@ -436,6 +550,7 @@ function cpd(
436550 scale_by_lambda = scale_by_lambda,
437551 lambda_eps = lambda_eps,
438552 pullback_eps = pullback_eps,
553+ component_trace = component_trace,
439554 verbose = verbose,
440555 vector_transport_method = vector_transport_method,
441556 kwargs... ,
@@ -459,6 +574,7 @@ function cpd(
459574 lambda_eps = lambda_eps,
460575 nonnegative = false ,
461576 pullback_eps = pullback_eps,
577+ component_trace = component_trace,
462578 verbose = verbose,
463579 vector_transport_method = vector_transport_method,
464580 kwargs... ,
0 commit comments