Skip to content

Commit e3f1e7e

Browse files
authored
Merge pull request #17 from TensorKitchen/fix-decomp-comp-accessor
Fix DecompositionComponent accessor
2 parents 4adb3bb + d2f1605 commit e3f1e7e

4 files changed

Lines changed: 95 additions & 27 deletions

File tree

‎src/api/btd.jl‎

Lines changed: 12 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,9 @@ end
5050
return solver == :als ? BTDHOSVDMultistartInit() : :alswarm
5151
end
5252

53+
@inline _btd_solver_symbol(solver::AbstractSolver) =
54+
solver isa ALSSolver ? :als : solver_symbol(solver)
55+
5356
function _btd_effective_init(
5457
solver::Symbol,
5558
init,
@@ -246,7 +249,9 @@ function btd(
246249
restart_seed = nothing,
247250
kwargs...,
248251
) where {T<:AbstractFloat,N}
249-
init_resolved = _resolve_btd_init(init, solver)
252+
solver_obj = _solver_object(solver, stepsize; kwargs...)
253+
solver_sym = _btd_solver_symbol(solver_obj)
254+
init_resolved = _resolve_btd_init(init, solver_sym)
250255
init_eff =
251256
init_resolved == :alswarm ?
252257
BTDALSWarmStartInit(
@@ -264,12 +269,12 @@ function btd(
264269
warm_info = (;)
265270
short_circuited = Ref(false)
266271
result = with_phase_progress() do
267-
if solver ∉ (:als,) && init_eff isa BTDALSWarmStartInit
272+
if solver_sym ∉ (:als,) && init_eff isa BTDALSWarmStartInit
268273
warm = _btd_warm_start_result(
269274
model,
270275
b,
271276
init_eff,
272-
solver;
277+
solver_sym;
273278
verbose = verbose,
274279
warm_rel_error_gate,
275280
)
@@ -285,7 +290,7 @@ function btd(
285290
model;
286291
init = solve_init,
287292
p0 = solve_p0,
288-
solver = solver,
293+
solver = solver_obj,
289294
maxiter = maxiter,
290295
stepsize = stepsize,
291296
tol = tol,
@@ -311,15 +316,15 @@ function btd(
311316
isempty(propertynames(warm_info)) ? result :
312317
_merge_btd_solver_info(result, warm_info)
313318
polish_n = if btd_als_polish_maxiter === nothing
314-
solver == :als ? 0 : clamp(maxiter ÷ 2, 20, 500)
319+
solver_sym == :als ? 0 : clamp(maxiter ÷ 2, 20, 500)
315320
else
316321
btd_als_polish_maxiter
317322
end
318-
if polish_n > 0 && solver ∉ (:als,)
323+
if polish_n > 0 && solver_sym ∉ (:als,)
319324
result = _polish_btd_with_als(
320325
b,
321326
result,
322-
solver;
327+
solver_sym;
323328
block_method = block_method,
324329
block_maxiter = block_maxiter,
325330
polish_maxiter = polish_n,

‎src/core/types.jl‎

Lines changed: 33 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
# core/types.jl — RankOneTensor, CPDResult, ApproxResult, TuckerResult
22
export RankOneTensor,
33
CPDResult,
4+
CPDComponent,
45
ApproxResult,
56
BTDResult,
67
TuckerResult,
@@ -37,6 +38,20 @@ struct RankOneTensor{T<:AbstractFloat}
3738
vectors::Vector{Vector{T}}
3839
end
3940

41+
"""
42+
CPDComponent{T}
43+
44+
Lazy public component wrapper for CPD-backed join results.
45+
46+
Each component stores its component-specific rank-one point and decoded
47+
rank-one tensor data. Dense tensors are reconstructed on demand through
48+
`component.tensor` or [`tensor(component)`](@ref).
49+
"""
50+
struct CPDComponent{T<:AbstractFloat,P}
51+
point::P
52+
component::RankOneTensor{T}
53+
end
54+
4055
"""
4156
CPDResult{T}
4257
@@ -244,7 +259,8 @@ end
244259
Return the factor vectors of a rank-one tensor component.
245260
"""
246261
vectors(c::RankOneTensor) = c.vectors
247-
kind(c::DecompositionComponent) = c.kind
262+
kind(::CPDComponent) = :Segre
263+
kind(c::DecompositionComponent) = typeof(c.manifold)
248264

249265
"""
250266
point(x)
@@ -254,28 +270,42 @@ Return the optimization point stored in a result or component.
254270
Supported inputs include [`CPDResult`](@ref), [`ApproxResult`](@ref),
255271
[`BTDResult`](@ref), and [`DecompositionComponent`](@ref).
256272
"""
273+
point(c::CPDComponent) = c.point
257274
point(c::DecompositionComponent) = c.point
258275

259276
"""
260277
tensor(c::DecompositionComponent)
261278
262279
Reconstruct a component into its ambient tensor representation.
263280
"""
281+
reconstruct(c::CPDComponent) = reconstruct_cp_rank1(λ(c.component), vectors(c.component))
282+
tensor(c::CPDComponent) = reconstruct(c)
264283
tensor(c::DecompositionComponent) = c.tensor
265284

266285
"""
267286
core(x)
268287
269288
Return the Tucker core stored in a Tucker component.
270289
"""
271-
core(c::DecompositionComponent) = c.core
290+
core(c::DecompositionComponent) = _component_core(c.point)
272291

273292
"""
274293
factors(x)
275294
276295
Return factor matrices for a CP or Tucker result/component.
277296
"""
278-
factors(c::DecompositionComponent) = c.factors
297+
weights(c::CPDComponent) = [λ(c.component)]
298+
factors(c::CPDComponent) = [reshape(v, :, 1) for v in vectors(c.component)]
299+
factors(c::DecompositionComponent) = _component_factors(c.point)
300+
301+
function Base.getproperty(c::CPDComponent, name::Symbol)
302+
name === :kind && return kind(c)
303+
name === :weights && return weights(c)
304+
name === :factors && return factors(c)
305+
name === :tensor && return reconstruct(c)
306+
return getfield(c, name)
307+
end
308+
279309
point(r::CPDResult) = cpd_point(r)
280310
point(r::ApproxResult) = r.point
281311
point(r::BTDResult) = r.point

‎src/join/cpd_backend.jl‎

Lines changed: 4 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ function JoinModel(
2929
scale_by_lambda = scale_by_lambda,
3030
lambda_eps = lambda_eps,
3131
nonnegative = nonnegative,
32-
use_pullback_metric = use_pullback_metric,
32+
use_pullback_metric = (geometry_eff == :squaring_metric) || use_pullback_metric,
3333
pullback_eps = pullback_eps,
3434
)
3535
b = CPDBackend(inner)
@@ -172,15 +172,7 @@ function extract_components(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p)
172172
λ, U = _decode_nonnegative_cpd(m, λ̃, Ũ)
173173
components_from_factors(λ, U)
174174
end : unpack_point_rankr_components(p, m.dims, m.r)
175-
return [
176-
(
177-
kind = :Segre,
178-
point = p,
179-
weights = [c.λ],
180-
factors = [reshape(c.vectors[j], :, 1) for j = 1:length(c.vectors)],
181-
tensor = reconstruct_cp_rank1(c.λ, c.vectors),
182-
) for c in comps
183-
]
175+
return [CPDComponent(pack_point_rank1(c.λ, c.vectors), c) for c in comps]
184176
elseif m isa Rank1CPDModel
185177
λ, U = unpack_point_rank1(p, m.dims)
186178
if m.nonnegative
@@ -192,13 +184,8 @@ function extract_components(model::JoinModel{<:AbstractFloat,<:CPDBackend}, p)
192184
U = [u .^ 2 for u in U]
193185
end
194186
end
195-
return [(
196-
kind = :Segre,
197-
point = p,
198-
weights = [λ],
199-
factors = [reshape(U[j], :, 1) for j = 1:length(U)],
200-
tensor = reconstruct_cp_rank1(λ, U),
201-
)]
187+
c = RankOneTensor(λ, U)
188+
return [CPDComponent(pack_point_rank1(λ, U), c)]
202189
end
203190
throw(
204191
ArgumentError(

‎test/basic_tests.jl‎

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -223,6 +223,21 @@ end
223223
@test size(Ahat) == size(A)
224224
@test rel_error(A, res) == TensorKitchen.relative_frobenius_error(A, Ahat)
225225
@test rel_error(A, Ahat) == rel_error(A, res)
226+
227+
model = JoinModel(A, r; geometry = :canonical)
228+
p = TensorKitchen.initial_point(model, :random)
229+
comps = TensorKitchen.extract_components(model, p)
230+
@test length(comps) == r
231+
@test comps[1] isa TensorKitchen.CPDComponent
232+
@test !(:tensor in fieldnames(typeof(comps[1])))
233+
@test comps[1].point !== p
234+
@test comps[1].kind == :Segre
235+
@test size(comps[1].tensor) == size(A)
236+
Xparts = zero(A)
237+
for c in comps
238+
Xparts .+= c.tensor
239+
end
240+
@test TensorKitchen.cost(model, p) ≈ 0.5 * sum(abs2, A .- Xparts)
226241
end
227242

228243
@testset "frontend defaults through public APIs" begin
@@ -562,6 +577,10 @@ end
562577
@test all(
563578
isapprox(norm(q_sep.factors[m][:, k]), 1; atol = 1e-10) for m = 1:3 for k = 1:r
564579
)
580+
U_sep, λ_sep = normalize_components(U, λ, SeparateLambdaNormalization())
581+
@test U_sep isa Vector{Matrix{Float64}}
582+
@test λ_sep isa Vector{Float64}
583+
@test reconstruct_cpd_rankr(λ_sep, U_sep) ≈ A_ref
565584

566585
q_last = normalize_components(CPDPoint(λ, U), :last_mode)
567586
@test reconstruct_cpd_rankr(q_last.lambda, q_last.factors) ≈ A_ref
@@ -886,6 +905,11 @@ end
886905
)
887906
@test TensorKitchen.manifold(model_sm) isa ProductManifold
888907
@test all(m -> m isa SqEuclidean, TensorKitchen.manifold(model_sm).manifolds)
908+
join_model_sm = JoinModel(A, r; nonnegative = true, geometry = :squaring_metric)
909+
@test all(
910+
m -> m isa SqEuclidean,
911+
TensorKitchen.manifold(TensorKitchen.cpd_model(join_model_sm)).manifolds,
912+
)
889913
@test getproperty(model_sm, :scale_by_lambda) == false
890914
res_nn_sm = cpd(
891915
A,
@@ -1712,6 +1736,28 @@ end
17121736
length(res_btd_tsd.solver_info.line_search_trial_history)
17131737
@test isfinite(res_btd_tsd.rel_error)
17141738
@test res_btd_tsd.rel_error ≈ norm(A - reconstruct(res_btd_tsd)) / norm(A)
1739+
res_btd_tsd_object = btd(
1740+
A,
1741+
2,
1742+
(2, 2, 2);
1743+
solver = BTDTSDSolver(stepsize = 1.0),
1744+
init = :hosvd_multistart,
1745+
maxiter = 1,
1746+
schedule = :cyclic,
1747+
btd_als_polish_maxiter = 0,
1748+
tol = 1e-6,
1749+
verbose = false,
1750+
)
1751+
@test res_btd_tsd_object isa BTDResult
1752+
@test res_btd_tsd_object.solver == :btd_tsd
1753+
@test_throws ArgumentError btd(
1754+
A,
1755+
2,
1756+
(2, 2, 2);
1757+
solver = :tsd,
1758+
maxiter = 1,
1759+
verbose = false,
1760+
)
17151761

17161762
manifolds = TensorKitchen._as_join_manifold_tuple(TuckerJoin(size(A), (2, 2, 2), 2))
17171763
backend = TensorKitchen._sum_backend_instance(TensorKitchen.BTDBackend, manifolds, A)

0 commit comments

Comments
 (0)