From 050353ec26b8a7bd7d580fcdca88855ef465d979 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 06:13:10 +0200 Subject: [PATCH 1/7] Reduce join backend allocations and add reconstruct_tucker!. Split sum-backend construction by type, reconstruct Tucker components in-place, document the join residual model, and test reconstruct_tucker!. --- src/join/join_backend.jl | 155 ++++++++++++++++++++++++++------------- src/tucker/hosvd.jl | 50 ++++++++++++- test/basic_tests.jl | 3 + 3 files changed, 156 insertions(+), 52 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 824987d..941983b 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -117,10 +117,6 @@ function _ambient_tensor(M, p, target_shape::Tuple) return reshape(_ambient_vector(M, p, prod(target_shape)), target_shape) end -supports_ambient_length(::AbstractManifold) = false -supports_ambient_length(M::Manifolds.Segre) = true -supports_ambient_length(M::Manifolds.Tucker) = true - """ ambient_length(M) -> Int @@ -191,6 +187,18 @@ function _sum_backend_instance( manifolds, target::AbstractArray{T,N}; init_point = nothing, +) where {T<:AbstractFloat,N} + throw( + ArgumentError( + "Unsupported join backend type $B. Use JoinBackend or BTDBackend.", + ), + ) +end + +function _sum_backend_parts( + manifolds, + target::AbstractArray{T,N}; + init_point = nothing, ) where {T<:AbstractFloat,N} r = length(manifolds) # Keep the original target representation instead of eagerly materializing Array. @@ -204,36 +212,64 @@ function _sum_backend_instance( work_rec = _join_vector_workspace_like(tgt, tgt_len) work_residual = _join_vector_workspace_like(tgt, tgt_len) - if B === BTDBackend - return B( - manifolds, - r, - tgt, - size(tgt), - tflat, - ProductManifold(manifolds...), - init_point, - work_rec, - work_residual, - BTDContractionWorkspace{T,N}(), - sum(abs2, tgt), - component_bufs, - ) - else - return B( - manifolds, - r, - tgt, - size(tgt), - tflat, - ProductManifold(manifolds...), - init_point, - work_rec, - work_residual, - _JoinResidualWORO(_join_vector_workspace_like(tgt, tgt_len)), - component_bufs, - ) - end + return (; + manifolds, + r, + target = tgt, + target_size = size(tgt), + target_flat = tflat, + target_len = tgt_len, + product = ProductManifold(manifolds...), + init_point, + work_rec, + work_residual, + component_bufs, + ) +end + +function _sum_backend_instance( + ::Type{JoinBackend}, + manifolds, + target::AbstractArray{T,N}; + init_point = nothing, +) where {T<:AbstractFloat,N} + parts = _sum_backend_parts(manifolds, target; init_point) + return JoinBackend( + parts.manifolds, + parts.r, + parts.target, + parts.target_size, + parts.target_flat, + parts.product, + parts.init_point, + parts.work_rec, + parts.work_residual, + _JoinResidualWORO(_join_vector_workspace_like(parts.target, parts.target_len)), + parts.component_bufs, + ) +end + +function _sum_backend_instance( + ::Type{BTDBackend}, + manifolds, + target::AbstractArray{T,N}; + init_point = nothing, +) where {T<:AbstractFloat,N} + parts = _sum_backend_parts(manifolds, target; init_point) + return BTDBackend( + parts.manifolds, + parts.r, + parts.target, + parts.target_size, + parts.target_flat, + parts.product, + parts.init_point, + parts.work_rec, + parts.work_residual, + BTDContractionWorkspace{T,N}(), + sum(abs2, parts.target), + parts.component_bufs, + ) end """ @@ -311,17 +347,6 @@ function initial_point( return ArrayPartition(parts...) end -function _join_reconstruct!(out::AbstractArray, manifolds::Tuple, r::Int, p) - parts = point_parts(p) - _check_parts_len(parts, r, "_join_reconstruct") - fill!(out, zero(eltype(out))) - out_len = length(out) - @inbounds for k = 1:r - out .+= _ambient_vector(manifolds[k], parts[k], out_len) - end - return out -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) @@ -430,8 +455,7 @@ function _ambient_vector!(out::AbstractVector, M, p) if M isa Manifolds.Tucker core = p.hosvd.core factors = p.hosvd.U - X = reconstruct_tucker(core, factors) - copyto!(out, vec(X)) + reconstruct_tucker!(reshape(out, factor_dims(M)), core, factors) return out end @@ -451,12 +475,41 @@ function _subtract_ambient_tensor!( ), ) _ambient_vector!(work_vec, M, p) - @inbounds for i in eachindex(residual, work_vec) - residual[i] -= work_vec[i] + residual_vec = vec(residual) + @inbounds for i in eachindex(residual_vec, work_vec) + residual_vec[i] -= work_vec[i] end return residual end +""" + _join_reconstruct!(out, backend, p) + +Reconstruct the ambient join approximation represented by `p` into `out`. + +For component manifolds `M_k` with embeddings +`Phi_k : M_k -> R^n`, the generic join model optimizes + +```math +f(p_1, ..., p_r) = + \\frac{1}{2}\\left\\|\\sum_{k=1}^r \\Phi_k(p_k) - A\\right\\|^2. +``` + +This backend implements the mathematical core of that model: + +- `_join_reconstruct!` computes `sum_k Phi_k(p_k)`. +- `_join_residual!` computes `sum_k Phi_k(p_k) - A`. +- `cost(model, p)` computes `1/2 * ||residual||^2`. +- `_manifold_egrad` applies the adjoint embedding derivative + `DPhi_k(p_k)'` to the residual for each component. +- `rgrad(model, p)` projects those component gradients to tangent spaces. +- `extract_components(model, p)` converts the optimized component points + `p_k` into result components. + +The method writes each component embedding into preallocated component buffers, +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 r = backend.r @@ -468,10 +521,10 @@ function _join_reconstruct!(out::AbstractArray, backend::Union{JoinBackend,BTDBa fill!(out, zero(eltype(out))) @inbounds for k = 1:r - # 1. 각 컴포넌트(블록/랭크)를 미리 할당된 개별 버퍼에 in-place 재구성 + # Reconstruct each component into its preallocated workspace. _ambient_vector!(bufs[k], manifolds[k], parts[k]) - # 2. 전체 결과를 담는 out 텐서에 누적 + # Accumulate into the output tensor without allocating a Khatri-Rao-sized object. out .+= bufs[k] end diff --git a/src/tucker/hosvd.jl b/src/tucker/hosvd.jl index 72b4b75..afa036e 100644 --- a/src/tucker/hosvd.jl +++ b/src/tucker/hosvd.jl @@ -1,4 +1,4 @@ -export tucker_hosvd, reconstruct_tucker +export tucker_hosvd, reconstruct_tucker, reconstruct_tucker! """ tucker_hosvd(A, ranks) computes (core, factors) @@ -46,4 +46,52 @@ function reconstruct_tucker(core::AbstractArray{T}, factors) where {T<:AbstractF return A end +""" + reconstruct_tucker!(out, core, factors) + +Reconstruct a Tucker tensor directly into `out`. + +This low-memory kernel avoids allocating the full reconstructed tensor before +copying it into `out`. It is useful in workspace-based paths where avoiding +large temporary tensors is more important than using BLAS-heavy mode products. +""" +function reconstruct_tucker!( + out::AbstractArray{T,N}, + core::AbstractArray{T,N}, + factors, +) where {T<:AbstractFloat,N} + length(factors) == N || throw( + DimensionMismatch( + "reconstruct_tucker!: got $(length(factors)) factors for an order-$N core.", + ), + ) + @inbounds for mode = 1:N + size(out, mode) == size(factors[mode], 1) || throw( + DimensionMismatch( + "reconstruct_tucker!: output mode $mode has size $(size(out, mode)), " * + "expected $(size(factors[mode], 1)).", + ), + ) + size(core, mode) == size(factors[mode], 2) || throw( + DimensionMismatch( + "reconstruct_tucker!: core mode $mode has size $(size(core, mode)), " * + "expected $(size(factors[mode], 2)).", + ), + ) + end + + @inbounds for I in CartesianIndices(out) + acc = zero(T) + for J in CartesianIndices(core) + val = core[J] + for mode = 1:N + val *= factors[mode][I[mode], J[mode]] + end + acc += val + end + out[I] = acc + end + return out +end + # reconstruction_error(A, core, factors) — defined in results/rel_error.jl (alias of rel_error). diff --git a/test/basic_tests.jl b/test/basic_tests.jl index 3ba6fe3..a2bc8f2 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -15,6 +15,9 @@ end Ahat = reconstruct_tucker(core, factors) @test size(Ahat) == size(A) + Ahat_inplace = similar(Ahat) + @test reconstruct_tucker!(Ahat_inplace, core, factors) === Ahat_inplace + @test Ahat_inplace ≈ Ahat @test reconstruction_error(A, core, factors) >= 0 @test reconstruction_error(A, core, factors) <= 1 + 1e-10 end From bda91543b755041330663d9e78d9f0cf6434d348 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 07:03:31 +0200 Subject: [PATCH 2/7] Format src/join/join_backend.jlto fix formatting job --- src/join/join_backend.jl | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 941983b..ee26b6e 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -188,11 +188,7 @@ function _sum_backend_instance( target::AbstractArray{T,N}; init_point = nothing, ) where {T<:AbstractFloat,N} - throw( - ArgumentError( - "Unsupported join backend type $B. Use JoinBackend or BTDBackend.", - ), - ) + throw(ArgumentError("Unsupported join backend type $B. Use JoinBackend or BTDBackend.")) end function _sum_backend_parts( From 353992e9962e5e41c659a558c57ba0d3f9b2fde9 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 07:08:43 +0200 Subject: [PATCH 3/7] docstring update --- src/join/join_backend.jl | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index ee26b6e..16a8843 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -74,7 +74,7 @@ function _manifold_egrad(M, p, residual) end """ - _ambient_vector(M, p, target_len) -> AbstractVector + _ambient_vector(M, p, target_len) returns AbstractVector Embed a component point into the flattened ambient tensor space, checking that its length matches the target. @@ -108,7 +108,7 @@ function _ambient_vector(M, p, target_len::Int) end """ - _ambient_tensor(M, p, target_shape) -> AbstractArray + _ambient_tensor(M, p, target_shape) returns AbstractArray Embed a component point into ambient space and reshape it to the target tensor shape. @@ -118,7 +118,7 @@ function _ambient_tensor(M, p, target_shape::Tuple) end """ - ambient_length(M) -> Int + ambient_length(M) returns Int Return the flattened ambient tensor length represented by a component manifold. """ @@ -137,7 +137,7 @@ ambient_length(M::Manifolds.Segre) = prod(factor_dims(M)) ambient_length(M::Manifolds.Tucker) = prod(factor_dims(M)) """ - _join_vector_workspace_like(target, n) + _join_vector_workspace_like(target, n) returns AbstractVector Allocate a flattened work vector with the same storage style and scalar type as the target tensor. @@ -151,7 +151,7 @@ the target tensor. end """ - _validate_join_ambient_compatibility(manifolds, target) + _validate_join_ambient_compatibility(manifolds, target) validates the join ambient compatibility Ensure every join component embeds into the same flattened ambient space as the target tensor. @@ -177,7 +177,7 @@ function _validate_join_ambient_compatibility(manifolds::Tuple, target::Abstract end """ - _sum_backend_instance(B, manifolds, target; init_point=nothing) + _sum_backend_instance(B, manifolds, target; init_point=nothing) returns AbstractJoinBackend Construct either a generic `JoinBackend` or `BTDBackend` with shared target, product manifold, and reusable reconstruction/residual buffers. From 5a605f2fad4880248e47ce3ecb3b592d411716c2 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 07:39:56 +0200 Subject: [PATCH 4/7] refactor join_backend.jl join.jl to julia dispatch structure --- src/join/cpd_backend.jl | 50 ++++++++++-------- src/join/join_backend.jl | 111 +++++++++++++++++++++------------------ src/manifolds/join.jl | 37 +++++++------ 3 files changed, 109 insertions(+), 89 deletions(-) diff --git a/src/join/cpd_backend.jl b/src/join/cpd_backend.jl index 5b9f4e2..172bcfe 100644 --- a/src/join/cpd_backend.jl +++ b/src/join/cpd_backend.jl @@ -108,7 +108,7 @@ end function _cpd_result(model::JoinModel{<:AbstractFloat,<:CPDBackend}, result, dims, r) m = cpd_model(model) - solver_sym = solver(result) isa Symbol ? solver(result) : :unknown + solver_sym = _result_solver_symbol(solver(result)) si = hasproperty(result, :solver_info) ? solver_info(result) : (;) als_family = solver_sym in _CP_ALS_FAMILY_SOLVERS @@ -164,29 +164,35 @@ function _cpd_result(model::JoinModel{<:AbstractFloat,<:CPDBackend}, result, dim end function extract_components(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p) - m = cpd_model(model) - if m isa RankRCPDModel - 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) - return [CPDComponent(pack_point_rank1(c.λ, c.vectors), c) for c in comps] - elseif m isa Rank1CPDModel - λ, 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 + 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) + 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 - c = RankOneTensor(λ, U) - return [CPDComponent(pack_point_rank1(λ, U), c)] end + c = RankOneTensor(λ, U) + return [CPDComponent(pack_point_rank1(λ, U), c)] +end + +function _extract_cpd_components(m, p) throw( ArgumentError( "Unsupported CPD backend model $(typeof(m)) for component extraction.", diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 16a8843..4ab1ffc 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -1,6 +1,7 @@ # join/join_backend.jl — Generic join and BTD backend implementation -@inline _is_manifold_like(x) = x isa AbstractManifold +_is_manifold_like(::AbstractManifold) = true +_is_manifold_like(_) = false function _as_join_manifold_tuple(manifolds::Tuple) all(_is_manifold_like, manifolds) || throw( @@ -30,15 +31,14 @@ _as_join_manifold_tuple(M::ProductManifold) = Tuple(M.manifolds) ) end -function _manifold_init(M, target, init) - init_sym = _builtin_initializer_symbol(init) - if M isa Manifolds.Sphere - return _sphere_init(M, target, init_sym) - elseif M isa Manifolds.Segre - return _segre_init(M, target, init_sym) - elseif M isa Manifolds.Tucker - return _tucker_init(M, target, init_sym) - end +_manifold_init(M, target, init) = + _manifold_init(M, target, _builtin_initializer_symbol(init)) + +_manifold_init(M::Manifolds.Sphere, target, init::Symbol) = _sphere_init(M, target, init) +_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) + +function _manifold_init(M, target, init_sym::Symbol) init_sym == :random && return rand(M) throw( ArgumentError( @@ -53,24 +53,22 @@ end Compute the component Euclidean gradient induced by a join residual. Tucker components use their native tensor gradient; other components copy the residual. """ -function _manifold_egrad(M, p, residual) - if M isa Manifolds.Tucker - return _tucker_egrad(M, p, residual) - elseif M isa Manifolds.Segre - dims = factor_dims(M) - R = reshape(residual, dims) - parts = point_parts(p) - λ = parts[1][1] - grad_λ = rank1_inner_parts(R, parts) - grad_U = Vector{Vector{eltype(R)}}(undef, length(dims)) - @inbounds for m = 1:length(dims) - g = rank1_mode_contract_parts(R, parts, m) - rmul!(g, λ) - grad_U[m] = g - end - return pack_tangent_rank1_segre(grad_λ, grad_U) +_manifold_egrad(M, p, residual) = copy(residual) +_manifold_egrad(M::Manifolds.Tucker, p, residual) = _tucker_egrad(M, p, residual) + +function _manifold_egrad(M::Manifolds.Segre, p, residual) + dims = factor_dims(M) + R = reshape(residual, dims) + parts = point_parts(p) + λ = parts[1][1] + grad_λ = rank1_inner_parts(R, parts) + grad_U = Vector{Vector{eltype(R)}}(undef, length(dims)) + @inbounds for m = 1:length(dims) + g = rank1_mode_contract_parts(R, parts, m) + rmul!(g, λ) + grad_U[m] = g end - return copy(residual) + return pack_tangent_rank1_segre(grad_λ, grad_U) end """ @@ -80,23 +78,6 @@ Embed a component point into the flattened ambient tensor space, checking that its length matches the target. """ function _ambient_vector(M, p, target_len::Int) - if M isa Manifolds.Tucker - p isa Manifolds.TuckerPoint || throw( - ArgumentError( - "Expected native TuckerPoint for Manifolds.Tucker, got $(typeof(p)).", - ), - ) - core = p.hosvd.core - factors = p.hosvd.U - X = reconstruct_tucker(core, factors) - length(X) == target_len || throw( - DimensionMismatch( - "Tucker reconstructed tensor length $(length(X)) != target length $target_len.", - ), - ) - return vec(X) - end - emb = ManifoldsBase.embed(M, p) length(emb) == target_len || throw( DimensionMismatch( @@ -107,6 +88,26 @@ function _ambient_vector(M, p, target_len::Int) return vec(emb) end +function _ambient_vector(M::Manifolds.Tucker, p::Manifolds.TuckerPoint, target_len::Int) + core = p.hosvd.core + factors = p.hosvd.U + X = reconstruct_tucker(core, factors) + length(X) == target_len || throw( + DimensionMismatch( + "Tucker reconstructed tensor length $(length(X)) != target length $target_len.", + ), + ) + return vec(X) +end + +function _ambient_vector(M::Manifolds.Tucker, p, target_len::Int) + throw( + ArgumentError( + "Expected native TuckerPoint for Manifolds.Tucker, got $(typeof(p)).", + ), + ) +end + """ _ambient_tensor(M, p, target_shape) returns AbstractArray @@ -448,17 +449,25 @@ function rgrad(model::JoinModel{<:AbstractFloat,<:JoinBackend}, p) end function _ambient_vector!(out::AbstractVector, M, p) - if M isa Manifolds.Tucker - core = p.hosvd.core - factors = p.hosvd.U - reconstruct_tucker!(reshape(out, factor_dims(M)), core, factors) - return out - end - ManifoldsBase.embed!(M, out, p) return out end +function _ambient_vector!(out::AbstractVector, M::Manifolds.Tucker, p::Manifolds.TuckerPoint) + core = p.hosvd.core + factors = p.hosvd.U + reconstruct_tucker!(reshape(out, factor_dims(M)), core, factors) + return out +end + +function _ambient_vector!(out::AbstractVector, M::Manifolds.Tucker, p) + throw( + ArgumentError( + "Expected native TuckerPoint for Manifolds.Tucker, got $(typeof(p)).", + ), + ) +end + function _subtract_ambient_tensor!( residual::AbstractArray{T,N}, M, diff --git a/src/manifolds/join.jl b/src/manifolds/join.jl index 6531ccb..f67f03d 100644 --- a/src/manifolds/join.jl +++ b/src/manifolds/join.jl @@ -3,7 +3,8 @@ # Why join_product instead of "just ProductManifold"? ProductManifold is always the # result type; join_product(base, r) decides how to expand base into r components: # - Manifolds.Segre → flattened (Euclidean(1), Sphere, ...) × r (one λ + spheres per rank-1). -# - Manifolds.Tucker / ProductManifold → repeat each factor r times. +# - ProductManifold → repeat each existing product factor r times. +# - Manifolds.Tucker → repeat the Tucker manifold r times. # - Generic manifold → ProductManifold(base, base, ..., base). # So we use join_product (via CPJoin/TuckerJoin) wherever we want this expansion; # the pipeline uses it for BTD (btd → TuckerJoin) and for the CPJoin/TuckerJoin APIs. @@ -29,23 +30,27 @@ end Decides how to expand base into r components: * Manifolds.Segre → flattened (Euclidean(1), Sphere, ...) × r (one λ + spheres per rank-1). -* Manifolds.Tucker / ProductManifold → repeat each factor r times. +* ProductManifold → repeat each existing product factor r times. +* Manifolds.Tucker → repeat the Tucker manifold r times. * Generic manifold → ProductManifold(base, base, ..., base). """ -function join_product(base::T, r::Int) where {T<:AbstractManifold} - r >= 2 || throw(ArgumentError("Join rank must be at least 2, got r=$r.")) - - M = if base isa Manifolds.Segre - factors = _segre_flat_factors(base) - ProductManifold([deepcopy(f) for _ = 1:r for f in factors]...) - elseif base isa ProductManifold - factors = base.manifolds - ProductManifold([deepcopy(f) for _ = 1:r for f in factors]...) - else - ProductManifold(ntuple(_ -> deepcopy(base), r)...) - end - - return M +function join_product(base::AbstractManifold, r::Int) + r >= 1 || throw(ArgumentError("Join rank must be at least 1, got r=$r.")) + return _join_product(base, r) +end + +function _join_product(base::Manifolds.Segre, r::Int) + factors = _segre_flat_factors(base) + return ProductManifold([deepcopy(f) for _ = 1:r for f in factors]...) +end + +function _join_product(base::ProductManifold, r::Int) + factors = base.manifolds + return ProductManifold([deepcopy(f) for _ = 1:r for f in factors]...) +end + +function _join_product(base::AbstractManifold, r::Int) + return ProductManifold(ntuple(_ -> deepcopy(base), r)...) end product(M::ProductManifold) = M From adb75ba4e4bcd2f81803c8383a3e0f866b7c6757 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 07:44:50 +0200 Subject: [PATCH 5/7] Format src/join/join_backend.jl cpd_backend.jl src/manifolds/join.jl --- src/join/join_backend.jl | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/src/join/join_backend.jl b/src/join/join_backend.jl index 4ab1ffc..5ae1730 100644 --- a/src/join/join_backend.jl +++ b/src/join/join_backend.jl @@ -453,7 +453,11 @@ function _ambient_vector!(out::AbstractVector, M, p) return out end -function _ambient_vector!(out::AbstractVector, M::Manifolds.Tucker, p::Manifolds.TuckerPoint) +function _ambient_vector!( + out::AbstractVector, + M::Manifolds.Tucker, + p::Manifolds.TuckerPoint, +) core = p.hosvd.core factors = p.hosvd.U reconstruct_tucker!(reshape(out, factor_dims(M)), core, factors) From f0854acc7afd82aec2cf24d8a8a7b6856be6fcd4 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 09:35:59 +0200 Subject: [PATCH 6/7] refactor to dispatch structure --- src/api/btd.jl | 4 +- src/api/cpd.jl | 118 ++++++++++++++++++++++++++++++- src/cpd/core/cp_normalization.jl | 116 +++++++++++++++++++----------- src/dispatch/approx_routing.jl | 10 ++- src/results/conversion.jl | 7 +- src/solvers/abstract.jl | 35 +++++---- src/solvers/lbfgs.jl | 4 ++ src/solvers/rcg.jl | 21 +++--- src/solvers/rgd.jl | 8 +++ src/solvers/solve_dispatch.jl | 2 + src/tucker/hooi.jl | 70 ++++++++++-------- test/basic_tests.jl | 38 ++++++++++ 12 files changed, 332 insertions(+), 101 deletions(-) diff --git a/src/api/btd.jl b/src/api/btd.jl index 3846f4f..bf664cd 100644 --- a/src/api/btd.jl +++ b/src/api/btd.jl @@ -50,8 +50,8 @@ end return solver == :als ? BTDHOSVDMultistartInit() : :alswarm end -@inline _btd_solver_symbol(solver::AbstractSolver) = - solver isa ALSSolver ? :als : solver_symbol(solver) +@inline _btd_solver_symbol(::ALSSolver) = :als +@inline _btd_solver_symbol(solver::AbstractSolver) = solver_symbol(solver) function _btd_effective_init( solver::Symbol, diff --git a/src/api/cpd.jl b/src/api/cpd.jl index 43a555c..f3719de 100644 --- a/src/api/cpd.jl +++ b/src/api/cpd.jl @@ -22,6 +22,104 @@ function _merge_res_solver_info(res, patch::NamedTuple) ) end +mutable struct _CPDComponentTraceRecorder{M} + model::M + previous::Any + previous_cost::Float64 + iterations::Vector{Int} + cost_history::Vector{Float64} + cost_rel_change_history::Vector{Float64} + max_component_delta_history::Vector{Float64} + component_delta_history::Vector{Vector{Float64}} +end + +function _CPDComponentTraceRecorder(model) + return _CPDComponentTraceRecorder( + model, + nothing, + NaN, + Int[], + Float64[], + Float64[], + Float64[], + Vector{Float64}[], + ) +end + +function _rankone_norm2(λ, U, k::Int) + val = abs2(λ[k]) + @inbounds for m = 1:length(U) + val *= sum(abs2, @view U[m][:, k]) + end + return Float64(val) +end + +function _rankone_inner(λa, Ua, λb, Ub, k::Int) + val = λa[k] * λb[k] + @inbounds for m = 1:length(Ua) + val *= dot(@view(Ua[m][:, k]), @view(Ub[m][:, k])) + end + return Float64(val) +end + +function _cpd_component_deltas(prev::CPDPoint, curr::CPDPoint) + λ_prev = lambda(prev) + U_prev = factors(prev) + λ_curr = lambda(curr) + U_curr = factors(curr) + r = length(λ_curr) + deltas = Vector{Float64}(undef, r) + @inbounds for k = 1:r + n_prev = _rankone_norm2(λ_prev, U_prev, k) + n_curr = _rankone_norm2(λ_curr, U_curr, k) + cross = _rankone_inner(λ_prev, U_prev, λ_curr, U_curr, k) + delta = sqrt(max(n_prev + n_curr - 2 * cross, 0.0)) + deltas[k] = delta / max(sqrt(max(n_prev, 0.0)), 1.0) + end + return deltas +end + +function _record_cpd_component_trace!(rec::_CPDComponentTraceRecorder, p, iter::Int) + q = cpd_point(rec.model, p) + cost_val = Float64(cost(rec.model, p)) + if rec.previous !== nothing + deltas = _cpd_component_deltas(rec.previous, q) + rel_change = abs(rec.previous_cost - cost_val) / max(abs(rec.previous_cost), 1.0) + push!(rec.iterations, iter) + push!(rec.cost_history, cost_val) + push!(rec.cost_rel_change_history, rel_change) + push!(rec.max_component_delta_history, maximum(deltas)) + push!(rec.component_delta_history, deltas) + end + rec.previous = q + rec.previous_cost = cost_val + return nothing +end + +function _cpd_component_trace_callback(rec::_CPDComponentTraceRecorder) + return function (problem, state, k) + p = try + Manopt.get_iterate(state) + catch + return nothing + end + _record_cpd_component_trace!(rec, p, Int(k)) + return nothing + end +end + +function _cpd_component_trace_info(rec::_CPDComponentTraceRecorder) + return ( + component_trace_iterations = rec.iterations, + component_trace_cost_history = rec.cost_history, + component_trace_cost_rel_change_history = rec.cost_rel_change_history, + component_trace_max_delta_history = rec.max_component_delta_history, + component_trace_delta_history = rec.component_delta_history, + component_trace_final_max_delta = isempty(rec.max_component_delta_history) ? + NaN : rec.max_component_delta_history[end], + ) +end + function _pack_cpd_explicit_p0(model, p0) p0 isa CPDPoint && return pack_cpd_point(model, p0) p0 isa CPDResult && return pack_cpd_point(model, cpd_point(p0)) @@ -159,8 +257,12 @@ function _run_cpd_solver( verbose::Bool, vector_transport_method, pullback_eps, + component_trace, kwargs..., ) + trace_recorder = component_trace ? _CPDComponentTraceRecorder(model) : nothing + iteration_callbacks = + isnothing(trace_recorder) ? () : (_cpd_component_trace_callback(trace_recorder),) p_solve = if init_eff isa ALSWarmStartInit && isnothing(p0) && !(solver isa ALSSolver) _cpd_als_warm_then_pack( model, @@ -175,7 +277,7 @@ function _run_cpd_solver( _pack_cpd_explicit_p0(model, p0) end - return _solve_model( + result = _solve_model( model; init = init_eff, p0 = p_solve, @@ -188,8 +290,11 @@ function _run_cpd_solver( verbose, refinement_verbose = verbose, vector_transport_method, + iteration_callbacks, kwargs..., ) + return isnothing(trace_recorder) ? result : + _merge_res_solver_info(result, _cpd_component_trace_info(trace_recorder)) end function _cpd_impl( @@ -212,6 +317,7 @@ function _cpd_impl( verbose, vector_transport_method, pullback_eps = 1e-8, + component_trace::Bool = false, kwargs..., ) where {T<:AbstractFloat,N} haskey(kwargs, :softplus_beta) && throw( @@ -254,6 +360,9 @@ function _cpd_impl( throw(ArgumentError("geometry=$geometry_eff requires nonnegative=true.")) end if solver_obj isa ALSSolver + component_trace && throw( + ArgumentError("component_trace=true is only supported for manifold solvers."), + ) geometry_eff == :canonical || throw( ArgumentError( "solver=:als does not use manifold geometry. Use geometry=:canonical.", @@ -292,6 +401,7 @@ function _cpd_impl( verbose, vector_transport_method, pullback_eps = pullback_eps_eff, + component_trace, nonnegative, kwargs..., ) @@ -361,6 +471,9 @@ If `r` is omitted, uses the smallest tensor mode as a heuristic rank. * `verbose = true`: Enables progress output. * `nonnegative::Bool = false`: Nonnegative CPD option to be selected by the user. (same as `nncpd`) * `pullback_eps = 1e-8`: Regularization parameter for pullback-style nonnegative geometries. +* `component_trace = false`: For manifold solvers, records per-iteration movement of + each CP rank-one term in `solver_info`. Use this to diagnose whether a flat cost + means the rank-one terms are also stuck. ## Notes * `solver = :als` does not use manifold geometry. In that case: @@ -403,6 +516,7 @@ function cpd( verbose = true, vector_transport_method = nothing, pullback_eps = 1e-8, + component_trace::Bool = false, kwargs..., ) where {T<:AbstractFloat,N} if nonnegative @@ -436,6 +550,7 @@ function cpd( scale_by_lambda = scale_by_lambda, lambda_eps = lambda_eps, pullback_eps = pullback_eps, + component_trace = component_trace, verbose = verbose, vector_transport_method = vector_transport_method, kwargs..., @@ -459,6 +574,7 @@ function cpd( lambda_eps = lambda_eps, nonnegative = false, pullback_eps = pullback_eps, + component_trace = component_trace, verbose = verbose, vector_transport_method = vector_transport_method, kwargs..., diff --git a/src/cpd/core/cp_normalization.jl b/src/cpd/core/cp_normalization.jl index d9390c3..afa0362 100644 --- a/src/cpd/core/cp_normalization.jl +++ b/src/cpd/core/cp_normalization.jl @@ -73,52 +73,86 @@ function normalize_components!( ) end - if policy isa NoNormalization - return factors - elseif policy isa SeparateLambdaNormalization - @inbounds for k = 1:r - scale = lambda[k] - for m = 1:d - col = @view factors[m][:, k] - scale = _normalize_column_into_lambda!(col, scale) - end - lambda[k] = scale + return _normalize_components_policy!(factors, lambda, policy) +end + +_normalize_components_policy!( + factors::Vector{Matrix{T}}, + lambda::Vector{T}, + ::NoNormalization, +) where {T<:AbstractFloat} = factors + +function _normalize_components_policy!( + factors::Vector{Matrix{T}}, + lambda::Vector{T}, + ::SeparateLambdaNormalization, +) where {T<:AbstractFloat} + r = length(lambda) + d = length(factors) + @inbounds for k = 1:r + scale = lambda[k] + for m = 1:d + col = @view factors[m][:, k] + scale = _normalize_column_into_lambda!(col, scale) end - return factors - elseif policy isa LastModeNormalization - last_mode = d - @inbounds for k = 1:r - total_scale = lambda[k] - for m = 1:d - col = @view factors[m][:, k] - nu = _safe_column_norm!(col) - col ./= nu - total_scale *= nu - end - mag = abs(total_scale) - factors[last_mode][:, k] .*= mag - lambda[k] = _sign_or_zero(total_scale) + lambda[k] = scale + end + return factors +end + +function _normalize_components_policy!( + factors::Vector{Matrix{T}}, + lambda::Vector{T}, + ::LastModeNormalization, +) where {T<:AbstractFloat} + r = length(lambda) + d = length(factors) + last_mode = d + @inbounds for k = 1:r + total_scale = lambda[k] + for m = 1:d + col = @view factors[m][:, k] + nu = _safe_column_norm!(col) + col ./= nu + total_scale *= nu end - return factors - elseif policy isa EvenDistributionNormalization - @inbounds for k = 1:r - total_scale = lambda[k] - for m = 1:d - col = @view factors[m][:, k] - nu = _safe_column_norm!(col) - col ./= nu - total_scale *= nu - end - mag = abs(total_scale) - scale = mag <= eps(T) ? zero(T) : mag^(inv(T(d))) - for m = 1:d - factors[m][:, k] .*= scale - end - lambda[k] = _sign_or_zero(total_scale) + mag = abs(total_scale) + factors[last_mode][:, k] .*= mag + lambda[k] = _sign_or_zero(total_scale) + end + return factors +end + +function _normalize_components_policy!( + factors::Vector{Matrix{T}}, + lambda::Vector{T}, + ::EvenDistributionNormalization, +) where {T<:AbstractFloat} + r = length(lambda) + d = length(factors) + @inbounds for k = 1:r + total_scale = lambda[k] + for m = 1:d + col = @view factors[m][:, k] + nu = _safe_column_norm!(col) + col ./= nu + total_scale *= nu end - return factors + mag = abs(total_scale) + scale = mag <= eps(T) ? zero(T) : mag^(inv(T(d))) + for m = 1:d + factors[m][:, k] .*= scale + end + lambda[k] = _sign_or_zero(total_scale) end + return factors +end +function _normalize_components_policy!( + factors::Vector{Matrix{T}}, + lambda::Vector{T}, + policy::AbstractNormalizationPolicy, +) where {T<:AbstractFloat} throw(ArgumentError("Unsupported normalization policy $(typeof(policy)).")) end diff --git a/src/dispatch/approx_routing.jl b/src/dispatch/approx_routing.jl index d888f91..776d3a2 100644 --- a/src/dispatch/approx_routing.jl +++ b/src/dispatch/approx_routing.jl @@ -3,16 +3,22 @@ # If every summand is a Manifolds.Segre with the same factor_dims, the call is a # plain rank-r CPD. Route those directly into the cpd() tree so they get the # CPDBackend, CPDResult, and all CPD-specific kwargs (geometry, nonnegative, ...). +_is_segre_manifold(::Manifolds.Segre) = true +_is_segre_manifold(::AbstractManifold) = false + function _all_segre_uniform(manifolds) isempty(manifolds) && return false - all(m -> m isa Manifolds.Segre, manifolds) || return false + all(_is_segre_manifold, manifolds) || return false d0 = factor_dims(first(manifolds)) return all(m -> factor_dims(m) == d0, manifolds) end +_is_tucker_manifold(::Manifolds.Tucker) = true +_is_tucker_manifold(::AbstractManifold) = false + function _all_tucker_uniform(manifolds, target_shape::Tuple) isempty(manifolds) && return false - all(m -> m isa Manifolds.Tucker, manifolds) || return false + all(_is_tucker_manifold, manifolds) || return false dims0 = factor_dims(first(manifolds)) ranks0 = multilinear_rank(first(manifolds)) dims0 == target_shape || return false diff --git a/src/results/conversion.jl b/src/results/conversion.jl index 53b92c4..7c5795c 100644 --- a/src/results/conversion.jl +++ b/src/results/conversion.jl @@ -1,8 +1,11 @@ # results/conversion.jl — optimization result wrappers +_result_solver_symbol(solver::Symbol) = solver +_result_solver_symbol(solver) = :unknown + function _to_approx_result(model::JoinModel{T}, result) where {T<:AbstractFloat} comps = extract_components(model, result.point) - solver_sym = result.solver isa Symbol ? result.solver : :unknown + solver_sym = _result_solver_symbol(result.solver) solver_info = hasproperty(result, :solver_info) ? result.solver_info : (;) return ApproxResult( result.point, @@ -19,7 +22,7 @@ end function _to_btd_result(model::JoinModel{T}, result) where {T<:AbstractFloat} comps = extract_components(model, result.point) - solver_sym = result.solver isa Symbol ? result.solver : :unknown + 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( diff --git a/src/solvers/abstract.jl b/src/solvers/abstract.jl index 684a73d..b0a4575 100644 --- a/src/solvers/abstract.jl +++ b/src/solvers/abstract.jl @@ -209,6 +209,7 @@ function solve( verbose::Bool = true, return_stats::Bool = false, vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, + iteration_callbacks = (), ) where {T<:AbstractFloat} setup = _prepare_solver_problem(model; init, p0, gradient_mode, verbose) normalization_policy = _normalization_policy(normalization) @@ -231,6 +232,7 @@ function solve( vector_transport_method, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) end @@ -246,6 +248,7 @@ function solve( verbose::Bool = true, return_stats::Bool = false, vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, + iteration_callbacks = (), ) where {T<:AbstractFloat} setup = _prepare_solver_problem(model; init, p0, gradient_mode) normalization_policy = _normalization_policy(normalization) @@ -268,6 +271,7 @@ function solve( vector_transport_method, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) end @@ -364,23 +368,24 @@ end end function _solver_retraction_method(M, p) - M2 = _unwrap_solver_manifold(M) - if M2 isa ProductManifold - pparts0 = point_parts(p) - pparts = pparts0 isa Tuple ? pparts0 : Tuple(pparts0) - n = length(M2.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(M2.manifolds[i], pparts[i]), n) - return ManifoldsBase.ProductRetraction(methods) - end - return _default_component_retraction_method(M2, 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/lbfgs.jl b/src/solvers/lbfgs.jl index a649607..41e58fc 100644 --- a/src/solvers/lbfgs.jl +++ b/src/solvers/lbfgs.jl @@ -67,6 +67,7 @@ function solve_lbfgs( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback = nothing, diagnostics_recorder = nothing, + iteration_callbacks = (), memory_size::Int = 1, cautious_update::Bool = true, initial_scale::Real = 1.0, @@ -129,6 +130,7 @@ function solve_lbfgs( post_step_callback, diagnostics_callback, progress_callback, + iteration_callbacks..., ), count = [:Cost, :Gradient], return_state = true, @@ -200,6 +202,7 @@ function run_second_order_solver( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) return solve_lbfgs( setup.model_cost, @@ -215,6 +218,7 @@ function run_second_order_solver( vector_transport_method = vector_transport_method, post_step_callback, diagnostics_recorder, + iteration_callbacks, memory_size = solver.memory_size, cautious_update = solver.cautious_update, initial_scale = solver.initial_scale, diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index 8ab188a..c3f17da 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -4,17 +4,16 @@ export RCGSolver struct SegreProjectionTransport <: ManifoldsBase.AbstractVectorTransportMethod end function _uses_segre_projection_transport(M) - M2 = _unwrap_solver_manifold(M) - if M2 isa Manifolds.Segre - return true - elseif M2 isa ProductManifold - return all(_uses_segre_projection_transport, M2.manifolds) - elseif hasproperty(M2, :native) - return _uses_segre_projection_transport(getproperty(M2, :native)) - end - return false + return _uses_segre_projection_transport_unwrapped(_unwrap_solver_manifold(M)) 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, @@ -105,6 +104,7 @@ function solve_rcg( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback = nothing, diagnostics_recorder = nothing, + iteration_callbacks = (), ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -152,6 +152,7 @@ function solve_rcg( post_step_callback, diagnostics_callback, progress_callback, + iteration_callbacks..., ), count = [:Cost, :Gradient], return_state = true, @@ -223,6 +224,7 @@ function run_first_order_solver( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) return solve_rcg( setup.model_cost, @@ -238,5 +240,6 @@ function run_first_order_solver( vector_transport_method = vector_transport_method, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) end diff --git a/src/solvers/rgd.jl b/src/solvers/rgd.jl index e458f63..f52ea16 100644 --- a/src/solvers/rgd.jl +++ b/src/solvers/rgd.jl @@ -547,6 +547,7 @@ function solve_rgd( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback = nothing, diagnostics_recorder = nothing, + iteration_callbacks = (), ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -626,6 +627,7 @@ function solve_rgd( post_step_callback, diagnostics_callback, progress_callback, + iteration_callbacks..., ), count = [:Cost, :Gradient], return_state = true, @@ -691,6 +693,7 @@ function solve_rgd_fixed( model_grad = nothing, post_step_callback = nothing, diagnostics_recorder = nothing, + iteration_callbacks = (), ) p0_local = _solver_point(M, p0) T = _scalar_eltype(p0_local) @@ -725,6 +728,7 @@ function solve_rgd_fixed( post_step_callback, diagnostics_callback, progress_callback, + iteration_callbacks..., ), count = [:Cost, :Gradient], return_state = true, @@ -801,6 +805,7 @@ function run_first_order_solver( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) return solve_rgd( setup.model_cost, @@ -817,6 +822,7 @@ function run_first_order_solver( vector_transport_method = vector_transport_method, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) end @@ -850,6 +856,7 @@ function run_first_order_solver( vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) return solve_rgd_fixed( setup.model_cost, @@ -865,5 +872,6 @@ function run_first_order_solver( model_grad = setup.model_grad, post_step_callback, diagnostics_recorder, + iteration_callbacks, ) end diff --git a/src/solvers/solve_dispatch.jl b/src/solvers/solve_dispatch.jl index 0ee6b71..198eba7 100644 --- a/src/solvers/solve_dispatch.jl +++ b/src/solvers/solve_dispatch.jl @@ -58,6 +58,7 @@ function _solve_with_solver( normalization = NoNormalization(), verbose::Bool, vector_transport_method::Union{ManifoldsBase.AbstractVectorTransportMethod,Nothing} = nothing, + iteration_callbacks = (), kwargs..., ) return solve( @@ -72,6 +73,7 @@ function _solve_with_solver( verbose, return_stats = true, vector_transport_method, + iteration_callbacks, ) end diff --git a/src/tucker/hooi.jl b/src/tucker/hooi.jl index 86fa6a6..35920ff 100644 --- a/src/tucker/hooi.jl +++ b/src/tucker/hooi.jl @@ -17,35 +17,7 @@ function hooi( d = N processing_order = collect(1:d) - # Initialization - if init isa TuckerResult - td0 = init::TuckerResult{T,N} - size(td0.core) == ranks || throw( - DimensionMismatch( - "hooi: TuckerResult.core has size $(size(td0.core)), expected core size $ranks", - ), - ) - for m = 1:d - size(td0.factors[m], 1) == dims[m] || throw( - DimensionMismatch( - "hooi: TuckerResult factor $m has $(size(td0.factors[m],1)) rows, expected $(dims[m])", - ), - ) - size(td0.factors[m], 2) == ranks[m] || throw( - DimensionMismatch( - "hooi: TuckerResult factor $m has $(size(td0.factors[m],2)) cols, expected $(ranks[m])", - ), - ) - end - factors0 = td0.factors - singular_vals = [copy(td0.singular_values[m]) for m = 1:d] - elseif init == :sthosvd - td0 = sthosvd(A, ranks) - factors0 = td0.factors - singular_vals = td0.singular_values - else - error("Unknown init: $init. Use :sthosvd or a TuckerResult.") - end + factors0, singular_vals = _hooi_initial_factors(A, ranks, init) factors = [copy(factors0[m]) for m = 1:d] prev_rel_error = T(Inf) @@ -118,3 +90,43 @@ function hooi(A::AbstractArray{T,N}, ranks::Vector{Int}; kwargs...) where {T,N} @assert length(ranks) == N return hooi(A, Tuple(Int.(ranks)); kwargs...) end + +function _hooi_initial_factors( + A::AbstractArray{T,N}, + ranks::NTuple{N,Int}, + init::TuckerResult{T,N}, +) where {T<:AbstractFloat,N} + dims = size(A) + size(init.core) == ranks || throw( + DimensionMismatch( + "hooi: TuckerResult.core has size $(size(init.core)), expected core size $ranks", + ), + ) + @inbounds for m = 1:N + size(init.factors[m], 1) == dims[m] || throw( + DimensionMismatch( + "hooi: TuckerResult factor $m has $(size(init.factors[m], 1)) rows, expected $(dims[m])", + ), + ) + size(init.factors[m], 2) == ranks[m] || throw( + DimensionMismatch( + "hooi: TuckerResult factor $m has $(size(init.factors[m], 2)) cols, expected $(ranks[m])", + ), + ) + end + return init.factors, [copy(init.singular_values[m]) for m = 1:N] +end + +function _hooi_initial_factors( + A::AbstractArray{T,N}, + ranks::NTuple{N,Int}, + init::Symbol, +) where {T<:AbstractFloat,N} + init == :sthosvd || error("Unknown init: $init. Use :sthosvd or a TuckerResult.") + td0 = sthosvd(A, ranks) + return td0.factors, td0.singular_values +end + +function _hooi_initial_factors(A::AbstractArray, ranks::Tuple, init) + error("Unknown init: $init. Use :sthosvd or a TuckerResult.") +end diff --git a/test/basic_tests.jl b/test/basic_tests.jl index a2bc8f2..5c9fd36 100644 --- a/test/basic_tests.jl +++ b/test/basic_tests.jl @@ -180,6 +180,13 @@ end Jt = TuckerJoin(dims, mlrank, 2) @test Jt isa ProductManifold @test length(Jt.manifolds) == 2 + @test length(join_product(Mt, 1).manifolds) == 1 + + base_product = ProductManifold(Manifolds.Sphere(1), Manifolds.Sphere(2)) + Jp = join_product(base_product, 2) + @test Jp isa ProductManifold + @test length(Jp.manifolds) == 4 + @test_throws ArgumentError join_product(Mt, 0) Mt_vec = Manifolds.Tucker(collect(dims), collect(mlrank)) @test factor_dims(Mt_vec) == dims @@ -480,6 +487,37 @@ end @test res_init_sym.solver_info.function_evaluations >= 0 @test res_init_sym.solver_info.gradient_evaluations >= 1 + res_trace = cpd( + A, + r; + solver = :rgd, + init = :tucker, + maxiter = 5, + tol = 1e-6, + verbose = false, + component_trace = true, + ) + trace_info = res_trace.solver_info + @test hasproperty(trace_info, :component_trace_iterations) + @test hasproperty(trace_info, :component_trace_max_delta_history) + @test hasproperty(trace_info, :component_trace_delta_history) + @test length(trace_info.component_trace_iterations) == + length(trace_info.component_trace_max_delta_history) + @test length(trace_info.component_trace_delta_history) == + length(trace_info.component_trace_max_delta_history) + @test all(length(deltas) == r for deltas in trace_info.component_trace_delta_history) + @test all(isfinite, trace_info.component_trace_max_delta_history) + @test_throws ArgumentError cpd( + A, + r; + solver = :als, + init = :tucker, + maxiter = 1, + tol = 1e-6, + verbose = false, + component_trace = true, + ) + res_alswarm_obj = cpd( A, r; From 08b978ee3a841353ae91ab1596af59ff6dc23221 Mon Sep 17 00:00:00 2001 From: SE Choi Date: Mon, 18 May 2026 09:38:33 +0200 Subject: [PATCH 7/7] Format code to pass formatting checks --- src/api/cpd.jl | 4 ++-- src/solvers/abstract.jl | 3 ++- src/solvers/rcg.jl | 4 ++-- 3 files changed, 6 insertions(+), 5 deletions(-) diff --git a/src/api/cpd.jl b/src/api/cpd.jl index f3719de..94bd08c 100644 --- a/src/api/cpd.jl +++ b/src/api/cpd.jl @@ -115,8 +115,8 @@ function _cpd_component_trace_info(rec::_CPDComponentTraceRecorder) component_trace_cost_rel_change_history = rec.cost_rel_change_history, component_trace_max_delta_history = rec.max_component_delta_history, component_trace_delta_history = rec.component_delta_history, - component_trace_final_max_delta = isempty(rec.max_component_delta_history) ? - NaN : rec.max_component_delta_history[end], + component_trace_final_max_delta = isempty(rec.max_component_delta_history) ? NaN : + rec.max_component_delta_history[end], ) end diff --git a/src/solvers/abstract.jl b/src/solvers/abstract.jl index b0a4575..a14e7d0 100644 --- a/src/solvers/abstract.jl +++ b/src/solvers/abstract.jl @@ -380,7 +380,8 @@ function _solver_retraction_method_unwrapped(M::ProductManifold, p) "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) + methods = + ntuple(i -> _default_component_retraction_method(M.manifolds[i], pparts[i]), n) return ManifoldsBase.ProductRetraction(methods) end diff --git a/src/solvers/rcg.jl b/src/solvers/rcg.jl index c3f17da..08e17ad 100644 --- a/src/solvers/rcg.jl +++ b/src/solvers/rcg.jl @@ -11,8 +11,8 @@ _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 + hasproperty(M, :native) ? _uses_segre_projection_transport(getproperty(M, :native)) : + false function ManifoldsBase.vector_transport_to( M::Manifolds.Segre,