@@ -46,19 +46,9 @@ $(_note(:OutputSection))
4646
4747@doc " $(_doc_PDSN) "
4848function primal_dual_semismooth_Newton (
49- M:: AbstractManifold ,
50- N:: AbstractManifold ,
51- cost:: TF ,
52- p:: P ,
53- X:: T ,
54- m:: P ,
55- n:: Q ,
56- prox_F:: Function ,
57- diff_prox_F:: Function ,
58- prox_G_dual:: Function ,
59- diff_prox_G_dual:: Function ,
60- linearized_forward_operator:: Function ,
61- adjoint_linearized_operator:: Function ;
49+ M:: AbstractManifold , N:: AbstractManifold , cost:: TF , p:: P , X:: T , m:: P , n:: Q ,
50+ prox_F:: Function , diff_prox_F:: Function , prox_G_dual:: Function , diff_prox_G_dual:: Function ,
51+ linearized_forward_operator:: Function , adjoint_linearized_operator:: Function ;
6252 Λ:: Union{Function, Missing} = missing ,
6353 kwargs... ,
6454 ) where {TF, P, T, Q}
@@ -68,40 +58,19 @@ function primal_dual_semismooth_Newton(
6858 m_res = copy (M, m)
6959 n_res = copy (N, n)
7060 return primal_dual_semismooth_Newton! (
71- M,
72- N,
73- cost,
74- x_res,
75- ξ_res,
76- m_res,
77- n_res,
78- prox_F,
79- diff_prox_F,
80- prox_G_dual,
81- diff_prox_G_dual,
82- linearized_forward_operator,
83- adjoint_linearized_operator;
84- Λ = Λ,
85- kwargs... ,
61+ M, N, cost, x_res, ξ_res, m_res, n_res,
62+ prox_F, diff_prox_F, prox_G_dual, diff_prox_G_dual,
63+ linearized_forward_operator, adjoint_linearized_operator;
64+ Λ = Λ, kwargs... ,
8665 )
8766end
8867calls_with_kwargs (:: typeof (primal_dual_semismooth_Newton)) = (primal_dual_semismooth_Newton!,)
8968
9069@doc " $(_doc_PDSN) "
9170function primal_dual_semismooth_Newton! (
92- M:: mT ,
93- N:: nT ,
94- cost:: Function ,
95- p:: P ,
96- X:: T ,
97- m:: P ,
98- n:: Q ,
99- prox_F:: Function ,
100- diff_prox_F:: Function ,
101- prox_G_dual:: Function ,
102- diff_prox_G_dual:: Function ,
103- linearized_forward_operator:: Function ,
104- adjoint_linearized_operator:: Function ;
71+ M:: mT , N:: nT , cost:: Function , p:: P , X:: T , m:: P , n:: Q ,
72+ prox_F:: Function , diff_prox_F:: Function , prox_G_dual:: Function , diff_prox_G_dual:: Function ,
73+ linearized_forward_operator:: Function , adjoint_linearized_operator:: Function ;
10574 dual_stepsize = 1 / sqrt (8 ),
10675 evaluation:: AbstractEvaluationType = AllocatingEvaluation (),
10776 Λ:: Union{Function, Missing} = missing ,
@@ -115,24 +84,13 @@ function primal_dual_semismooth_Newton!(
11584 vector_transport_method:: VTM = default_vector_transport_method (M, typeof (p)),
11685 kwargs... ,
11786 ) where {
118- mT <: AbstractManifold ,
119- nT <: AbstractManifold ,
120- P,
121- Q,
122- T,
123- RM <: AbstractRetractionMethod ,
124- IRM <: AbstractInverseRetractionMethod ,
125- VTM <: AbstractVectorTransportMethod ,
87+ mT <: AbstractManifold , nT <: AbstractManifold , P, Q, T,
88+ RM <: AbstractRetractionMethod , IRM <: AbstractInverseRetractionMethod , VTM <: AbstractVectorTransportMethod ,
12689 }
12790 keywords_accepted (primal_dual_semismooth_Newton!; kwargs... )
12891 pdmsno = PrimalDualManifoldSemismoothNewtonObjective (
129- cost,
130- prox_F,
131- diff_prox_F,
132- prox_G_dual,
133- diff_prox_G_dual,
134- linearized_forward_operator,
135- adjoint_linearized_operator;
92+ cost, prox_F, diff_prox_F, prox_G_dual, diff_prox_G_dual,
93+ linearized_forward_operator, adjoint_linearized_operator;
13694 Λ = Λ,
13795 evaluation = evaluation,
13896 )
@@ -184,22 +142,17 @@ function primal_dual_step!(tmp::TwoManifoldProblem, pdsn::PrimalDualSemismoothNe
184142 N = get_manifold (tmp, 2 )
185143 # construct X
186144 X = construct_primal_dual_residual_vector (tmp, pdsn)
187-
188145 # construct matrix
189146 ∂X = construct_primal_dual_residual_covariant_derivative_matrix (tmp, pdsn)
190147 ∂X += pdsn. regularization_parameter * sparse (I, size (∂X)) # prevent singular matrix at solution
191-
192148 # solve matrix -> find coordinates
193149 d_coords = ∂X \ - X
194-
195150 dims = manifold_dimension (M)
196151 dx_coords = d_coords[1 : dims]
197152 dξ_coords = d_coords[(dims + 1 ): end ]
198-
199153 # compute step
200154 dx = get_vector (M, pdsn. p, dx_coords, DefaultOrthonormalBasis ())
201155 dξ = get_vector (N, pdsn. n, dξ_coords, DefaultOrthonormalBasis ())
202-
203156 # do step
204157 pdsn. p = retract (M, pdsn. p, dx, pdsn. retraction_method)
205158 return pdsn. X = pdsn. X + dξ
@@ -226,8 +179,7 @@ function construct_primal_dual_residual_vector(
226179 vector_transport_to (
227180 M,
228181 pdsn. m,
229- - pdsn. primal_stepsize *
230- (adjoint_linearized_operator (tmp, pdsn. m, pdsn. n, pdsn. X)),
182+ - pdsn. primal_stepsize * (adjoint_linearized_operator (tmp, pdsn. m, pdsn. n, pdsn. X)),
231183 pdsn. p,
232184 pdsn. vector_transport_method,
233185 ),
@@ -248,17 +200,10 @@ function construct_primal_dual_residual_vector(
248200 pdsn. n,
249201 )
250202 # (2) if p.Λ is missing, assume that n = Λ(m) and do not PT
251- ξ_update = if ! hasproperty (obj, :Λ!! ) || ismissing (obj. Λ!!)
252- ξ_update
253- else
254- vector_transport_to (
255- N,
256- forward_operator (tmp, pdsn. m),
257- ξ_update,
258- pdsn. n,
259- pdsn. vector_transport_method,
203+ noPT = ! hasproperty (obj, :Λ!! ) || ismissing (obj. Λ!!)
204+ ξ_update = noPT ? ξ_update : vector_transport_to (
205+ N, forward_operator (tmp, pdsn. m), ξ_update, pdsn. n, pdsn. vector_transport_method,
260206 )
261- end
262207 # (3) the dual update
263208 ξ_update = get_dual_prox (
264209 tmp, pdsn. n, pdsn. dual_stepsize, pdsn. X + pdsn. dual_stepsize * ξ_update
@@ -304,13 +249,11 @@ function construct_primal_dual_residual_covariant_derivative_matrix(
304249 pdsn. n,
305250 )
306251 # (2) if p.Λ is missing, assume that n = Λ(m) and do not PT
307- η₁ = if ! hasproperty (obj, :Λ!! ) || ismissing (obj. Λ!!)
308- η₁
309- else
310- vector_transport_to (
252+ noPT = ! hasproperty (obj, :Λ!! ) || ismissing (obj. Λ!!)
253+
254+ η₁ = noPT ? η₁ : vector_transport_to (
311255 N, forward_operator (tmp, pdsn. m), η₁, pdsn. n, pdsn. vector_transport_method
312256 )
313- end
314257 # (3) to the dual update
315258 η₁ = pdsn. X + pdsn. dual_stepsize * η₁
316259 # construct ∂X₁₁ and ∂X₂₁
@@ -343,17 +286,13 @@ function construct_primal_dual_residual_covariant_derivative_matrix(
343286 ∂X₁₁[:, j] = sp_∂X₁₁j
344287
345288 Mⱼ = differential_log_argument (M, pdsn. m, pdsn. p, Θⱼ)
346- Kⱼ = if ! hasproperty (obj, :Λ!! ) || ismissing (obj. Λ!!)
347- pdsn. dual_stepsize * linearized_forward_operator (tmp, pdsn. m, Mⱼ, pdsn. n)
348- else
349- pdsn. dual_stepsize * vector_transport_to (
350- N,
351- forward_operator (tmp, pdsn. m),
352- linearized_forward_operator (tmp, pdsn. m, Mⱼ, pdsn. n),
353- pdsn. n,
354- pdsn. vector_transport_method,
355- )
356- end
289+ noPT = ! hasproperty (obj, :Λ!! ) || ismissing (obj. Λ!!)
290+ Kⱼ = pdsn. dual_stepsize * (
291+ noPT ? linearized_forward_operator (tmp, pdsn. m, Mⱼ, pdsn. n) : vector_transport_to (
292+ N, forward_operator (tmp, pdsn. m), linearized_forward_operator (tmp, pdsn. m, Mⱼ, pdsn. n), pdsn. n,
293+ pdsn. vector_transport_method,
294+ )
295+ )
357296 Jⱼ = get_differential_dual_prox (tmp, pdsn. n, pdsn. dual_stepsize, η₁, Kⱼ)
358297 ∂X₂₁j = get_coordinates (N, pdsn. n, - Jⱼ, DefaultOrthonormalBasis ())
359298
0 commit comments