Skip to content

Commit d7e71ce

Browse files
authored
Merge pull request #21 from TensorKitchen/refactor
Refactor dispatch routing and dual stopping criteria
2 parents cc531bf + 86e2ffd commit d7e71ce

14 files changed

Lines changed: 684 additions & 298 deletions

File tree

src/api/approx.jl

Lines changed: 143 additions & 72 deletions
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@ approx(M, target; verbose = false)
8585
# returns an BTDResult
8686
```
8787
88-
* `approx(base, r, target; kwargs...)` : builds a rank-r Segre join and routes to CPD when base isa Manifolds.Segre and BTD when base isa Manifolds.Tucker unless generic
88+
* `approx(base, r, target; kwargs...)` : builds a rank-r Segre join and routes by the type of `base`: `Manifolds.Segre` uses CPD and `Manifolds.Tucker` uses BTD unless generic
8989
dispatch is explicitly requested.
9090
- base means one manifold template, not yet a full join.
9191
- `approx(base, r, target; ...)` repeats that same manifold r times to build a join.
@@ -133,31 +133,81 @@ For the generic join path:
133133
* `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.
134134
* For generic mixed joins, use manifold solvers such as `:rgd`, `:rcg`, or `:lbfgs`.
135135
"""
136+
function _approx_manifold_collection(
137+
dispatch::AutoApproxDispatch,
138+
manifolds,
139+
target::AbstractArray;
140+
kwargs...,
141+
)
142+
_all_segre_uniform(manifolds) && return _approx_manifold_collection(
143+
CPDApproxDispatch(),
144+
manifolds,
145+
target;
146+
kwargs...,
147+
)
148+
_all_tucker_uniform(manifolds, size(target)) && return _approx_manifold_collection(
149+
BTDApproxDispatch(),
150+
manifolds,
151+
target;
152+
kwargs...,
153+
)
154+
return _approx_manifold_collection(
155+
GenericApproxDispatch(),
156+
manifolds,
157+
target;
158+
kwargs...,
159+
)
160+
end
161+
162+
function _approx_manifold_collection(
163+
::CPDApproxDispatch,
164+
manifolds,
165+
target::AbstractArray;
166+
kwargs...,
167+
)
168+
_all_segre_uniform(manifolds) || throw(
169+
ArgumentError(
170+
"approx(...; dispatch=:cpd) requires all manifolds to be Manifolds.Segre with identical factor_dims.",
171+
),
172+
)
173+
return cpd(target, length(manifolds); kwargs...)
174+
end
175+
176+
function _approx_manifold_collection(
177+
::BTDApproxDispatch,
178+
manifolds,
179+
target::AbstractArray;
180+
kwargs...,
181+
)
182+
_all_tucker_uniform(manifolds, size(target)) || throw(
183+
ArgumentError(
184+
"approx(...; dispatch=:btd) requires all manifolds to be Manifolds.Tucker with identical factor_dims/multilinear_rank matching the target.",
185+
),
186+
)
187+
return btd(target, length(manifolds), multilinear_rank(first(manifolds)); kwargs...)
188+
end
189+
190+
function _approx_manifold_collection(
191+
::GenericApproxDispatch,
192+
manifolds,
193+
target::AbstractArray;
194+
kwargs...,
195+
)
196+
return approx(JoinModel(manifolds, target); kwargs...)
197+
end
198+
136199
function approx(
137200
manifolds::Tuple{Vararg{AbstractManifold}},
138201
target::AbstractArray{T,N};
139202
dispatch::Symbol = :auto,
140203
kwargs...,
141204
) where {T<:AbstractFloat,N}
142-
dispatch = _normalize_approx_dispatch(dispatch)
143-
if dispatch == :cpd || (dispatch == :auto && _all_segre_uniform(manifolds))
144-
_all_segre_uniform(manifolds) || throw(
145-
ArgumentError(
146-
"approx(...; dispatch=:cpd) requires all manifolds to be Manifolds.Segre with identical factor_dims.",
147-
),
148-
)
149-
return cpd(target, length(manifolds); kwargs...)
150-
end
151-
if dispatch == :btd ||
152-
(dispatch == :auto && _all_tucker_uniform(manifolds, size(target)))
153-
_all_tucker_uniform(manifolds, size(target)) || throw(
154-
ArgumentError(
155-
"approx(...; dispatch=:btd) requires all manifolds to be Manifolds.Tucker with identical factor_dims/multilinear_rank matching the target.",
156-
),
157-
)
158-
return btd(target, length(manifolds), multilinear_rank(first(manifolds)); kwargs...)
159-
end
160-
return approx(JoinModel(manifolds, target); kwargs...)
205+
return _approx_manifold_collection(
206+
approx_dispatch(dispatch),
207+
manifolds,
208+
target;
209+
kwargs...,
210+
)
161211
end
162212

163213
function approx(
@@ -166,25 +216,12 @@ function approx(
166216
dispatch::Symbol = :auto,
167217
kwargs...,
168218
) where {T<:AbstractFloat,N}
169-
dispatch = _normalize_approx_dispatch(dispatch)
170-
if dispatch == :cpd || (dispatch == :auto && _all_segre_uniform(manifolds))
171-
_all_segre_uniform(manifolds) || throw(
172-
ArgumentError(
173-
"approx(...; dispatch=:cpd) requires all manifolds to be Manifolds.Segre with identical factor_dims.",
174-
),
175-
)
176-
return cpd(target, length(manifolds); kwargs...)
177-
end
178-
if dispatch == :btd ||
179-
(dispatch == :auto && _all_tucker_uniform(manifolds, size(target)))
180-
_all_tucker_uniform(manifolds, size(target)) || throw(
181-
ArgumentError(
182-
"approx(...; dispatch=:btd) requires all manifolds to be Manifolds.Tucker with identical factor_dims/multilinear_rank matching the target.",
183-
),
184-
)
185-
return btd(target, length(manifolds), multilinear_rank(first(manifolds)); kwargs...)
186-
end
187-
return approx(JoinModel(manifolds, target); kwargs...)
219+
return _approx_manifold_collection(
220+
approx_dispatch(dispatch),
221+
manifolds,
222+
target;
223+
kwargs...,
224+
)
188225
end
189226

190227
"""
@@ -199,27 +236,46 @@ function approx(
199236
dispatch::Symbol = :auto,
200237
kwargs...,
201238
) where {T<:AbstractFloat,N}
202-
dispatch = _normalize_approx_dispatch(dispatch)
203239
mfs = Tuple(M.manifolds)
204-
if dispatch == :cpd || (dispatch == :auto && _all_segre_uniform(mfs))
205-
_all_segre_uniform(mfs) || throw(
206-
ArgumentError(
207-
"approx(...; dispatch=:cpd) requires all manifolds to be Manifolds.Segre with identical factor_dims.",
208-
),
209-
)
210-
return cpd(target, length(mfs); kwargs...)
211-
end
212-
if dispatch == :btd || (dispatch == :auto && _all_tucker_uniform(mfs, size(target)))
213-
_all_tucker_uniform(mfs, size(target)) || throw(
214-
ArgumentError(
215-
"approx(...; dispatch=:btd) requires all manifolds to be Manifolds.Tucker with identical factor_dims/multilinear_rank matching the target.",
216-
),
217-
)
218-
return btd(target, length(mfs), multilinear_rank(first(mfs)); kwargs...)
219-
end
240+
return _approx_product_manifold(approx_dispatch(dispatch), M, mfs, target; kwargs...)
241+
end
242+
243+
function _approx_product_manifold(
244+
dispatch::Union{AutoApproxDispatch,CPDApproxDispatch,BTDApproxDispatch},
245+
M::ProductManifold,
246+
manifolds,
247+
target::AbstractArray;
248+
kwargs...,
249+
)
250+
return _approx_manifold_collection(dispatch, manifolds, target; kwargs...)
251+
end
252+
253+
function _approx_product_manifold(
254+
::GenericApproxDispatch,
255+
M::ProductManifold,
256+
manifolds,
257+
target::AbstractArray;
258+
kwargs...,
259+
)
220260
return approx(JoinModel(M, target); kwargs...)
221261
end
222262

263+
_approx_segre_rank(
264+
::Union{AutoApproxDispatch,CPDApproxDispatch},
265+
base::Manifolds.Segre,
266+
r::Int,
267+
target::AbstractArray;
268+
kwargs...,
269+
) = cpd(target, r; kwargs...)
270+
271+
_approx_segre_rank(
272+
::AbstractApproxDispatch,
273+
base::Manifolds.Segre,
274+
r::Int,
275+
target::AbstractArray;
276+
kwargs...,
277+
) = approx(JoinModel(base, r, target); kwargs...)
278+
223279
"""
224280
approx(base::Manifolds.Segre, r, target; dispatch=:auto, kwargs...)
225281
@@ -233,11 +289,25 @@ function approx(
233289
dispatch::Symbol = :auto,
234290
kwargs...,
235291
) where {T<:AbstractFloat,N}
236-
dispatch = _normalize_approx_dispatch(dispatch)
237-
return dispatch in (:auto, :cpd) ? cpd(target, r; kwargs...) :
238-
approx(JoinModel(base, r, target); kwargs...)
292+
return _approx_segre_rank(approx_dispatch(dispatch), base, r, target; kwargs...)
239293
end
240294

295+
_approx_tucker_rank(
296+
::Union{AutoApproxDispatch,BTDApproxDispatch},
297+
base::Manifolds.Tucker,
298+
r::Int,
299+
target::AbstractArray;
300+
kwargs...,
301+
) = btd(target, r, multilinear_rank(base); kwargs...)
302+
303+
_approx_tucker_rank(
304+
::AbstractApproxDispatch,
305+
base::Manifolds.Tucker,
306+
r::Int,
307+
target::AbstractArray;
308+
kwargs...,
309+
) = approx(JoinModel(base, r, target); kwargs...)
310+
241311
"""
242312
approx(base::Manifolds.Tucker, r, target; dispatch=:auto, kwargs...)
243313
@@ -251,9 +321,18 @@ function approx(
251321
dispatch::Symbol = :auto,
252322
kwargs...,
253323
) where {T<:AbstractFloat,N}
254-
dispatch = _normalize_approx_dispatch(dispatch)
255-
return dispatch in (:auto, :btd) ? btd(target, r, multilinear_rank(base); kwargs...) :
256-
approx(JoinModel(base, r, target); kwargs...)
324+
return _approx_tucker_rank(approx_dispatch(dispatch), base, r, target; kwargs...)
325+
end
326+
327+
_reject_generic_rank_dispatch(::AutoApproxDispatch) = nothing
328+
_reject_generic_rank_dispatch(::GenericApproxDispatch) = nothing
329+
330+
function _reject_generic_rank_dispatch(::CPDApproxDispatch)
331+
throw(ArgumentError("approx(...; dispatch=:cpd) requires Manifolds.Segre inputs."))
332+
end
333+
334+
function _reject_generic_rank_dispatch(::BTDApproxDispatch)
335+
throw(ArgumentError("approx(...; dispatch=:btd) requires Manifolds.Tucker inputs."))
257336
end
258337

259338
"""
@@ -269,11 +348,7 @@ function approx(
269348
dispatch::Symbol = :auto,
270349
kwargs...,
271350
) where {T<:AbstractFloat,N}
272-
dispatch = _normalize_approx_dispatch(dispatch)
273-
dispatch == :cpd &&
274-
throw(ArgumentError("approx(...; dispatch=:cpd) requires Manifolds.Segre inputs."))
275-
dispatch == :btd &&
276-
throw(ArgumentError("approx(...; dispatch=:btd) requires Manifolds.Tucker inputs."))
351+
_reject_generic_rank_dispatch(approx_dispatch(dispatch))
277352
return approx(JoinModel(base, r, target); kwargs...)
278353
end
279354

@@ -289,11 +364,7 @@ function approx(
289364
dispatch::Symbol = :auto,
290365
kwargs...,
291366
) where {T<:AbstractFloat,N}
292-
dispatch = _normalize_approx_dispatch(dispatch)
293-
dispatch == :cpd &&
294-
throw(ArgumentError("approx(...; dispatch=:cpd) requires Manifolds.Segre inputs."))
295-
dispatch == :btd &&
296-
throw(ArgumentError("approx(...; dispatch=:btd) requires Manifolds.Tucker inputs."))
367+
_reject_generic_rank_dispatch(approx_dispatch(dispatch))
297368
return approx(JoinModel(base, target); kwargs...)
298369
end
299370

src/api/btd.jl

Lines changed: 13 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -45,40 +45,23 @@ function _polish_btd_with_als(
4545
)
4646
end
4747

48-
@inline function _resolve_btd_init(init, solver::Symbol)
48+
_default_btd_init(::ALSSolver) = BTDHOSVDMultistartInit()
49+
_default_btd_init(::AbstractSolver) = :alswarm
50+
51+
@inline function _resolve_btd_init(init, solver::AbstractSolver)
4952
init == :auto || return init
50-
return solver == :als ? BTDHOSVDMultistartInit() : :alswarm
53+
return _default_btd_init(solver)
5154
end
5255

5356
@inline _btd_solver_symbol(::ALSSolver) = :als
5457
@inline _btd_solver_symbol(solver::AbstractSolver) = solver_symbol(solver)
5558

56-
function _btd_effective_init(
57-
solver::Symbol,
58-
init,
59-
warm_steps::Int,
60-
warm_init,
61-
warm_block_method::Symbol,
62-
warm_block_maxiter::Int,
63-
)
64-
if init == :alswarm
65-
return BTDALSWarmStartInit(
66-
warm_steps;
67-
base_init = warm_init,
68-
block_method = warm_block_method,
69-
block_maxiter = warm_block_maxiter,
70-
)
71-
elseif solver (:als,)
72-
return BTDALSWarmStartInit(
73-
warm_steps;
74-
base_init = init,
75-
block_method = warm_block_method,
76-
block_maxiter = warm_block_maxiter,
77-
)
78-
end
79-
return init
80-
end
59+
_btd_uses_warm_start(::AbstractSolver, _) = false
60+
_btd_uses_warm_start(::ALSSolver, ::BTDALSWarmStartInit) = false
61+
_btd_uses_warm_start(::AbstractSolver, ::BTDALSWarmStartInit) = true
8162

63+
_btd_should_polish(::ALSSolver, ::Integer) = false
64+
_btd_should_polish(::AbstractSolver, polish_n::Integer) = polish_n > 0
8265

8366
function _btd_warm_start_result(
8467
model::JoinModel{T,<:BTDBackend},
@@ -251,7 +234,7 @@ function btd(
251234
) where {T<:AbstractFloat,N}
252235
solver_obj = _solver_object(solver, stepsize; kwargs...)
253236
solver_sym = _btd_solver_symbol(solver_obj)
254-
init_resolved = _resolve_btd_init(init, solver_sym)
237+
init_resolved = _resolve_btd_init(init, solver_obj)
255238
init_eff =
256239
init_resolved == :alswarm ?
257240
BTDALSWarmStartInit(
@@ -269,7 +252,7 @@ function btd(
269252
warm_info = (;)
270253
short_circuited = Ref(false)
271254
result = with_phase_progress() do
272-
if solver_sym (:als,) && init_eff isa BTDALSWarmStartInit
255+
if _btd_uses_warm_start(solver_obj, init_eff)
273256
warm = _btd_warm_start_result(
274257
model,
275258
b,
@@ -320,7 +303,7 @@ function btd(
320303
else
321304
btd_als_polish_maxiter
322305
end
323-
if polish_n > 0 && solver_sym (:als,)
306+
if _btd_should_polish(solver_obj, polish_n)
324307
result = _polish_btd_with_als(
325308
b,
326309
result,

0 commit comments

Comments
 (0)