@@ -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
8989dispatch 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+
136199function 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+ )
161211end
162212
163213function 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+ )
188225end
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... )
221261end
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... )
239293end
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." ))
257336end
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... )
278353end
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... )
298369end
299370
0 commit comments