Skip to content

Commit 8bb0acb

Browse files
committed
Make PDSN tests standalone & a bit of code formatting.
1 parent cdbf0e7 commit 8bb0acb

3 files changed

Lines changed: 35 additions & 103 deletions

File tree

src/solvers/primal_dual_semismooth_Newton.jl

Lines changed: 28 additions & 89 deletions
Original file line numberDiff line numberDiff line change
@@ -46,19 +46,9 @@ $(_note(:OutputSection))
4646

4747
@doc "$(_doc_PDSN)"
4848
function 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
)
8766
end
8867
calls_with_kwargs(::typeof(primal_dual_semismooth_Newton)) = (primal_dual_semismooth_Newton!,)
8968

9069
@doc "$(_doc_PDSN)"
9170
function 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
= 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 +
@@ -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

src/solvers/quasi_Newton.jl

Lines changed: 2 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -774,20 +774,10 @@ function update_hessian!(
774774
# the stored vectors are just transported to the new tangent space; `sk` and `yk` are not added
775775
for i in 1:length(d.update.memory_s)
776776
vector_transport_to!(
777-
M,
778-
d.update.memory_s[i],
779-
p_old,
780-
d.update.memory_s[i],
781-
p,
782-
d.update.vector_transport_method,
777+
M, d.update.memory_s[i], p_old, d.update.memory_s[i], p, d.update.vector_transport_method,
783778
)
784779
vector_transport_to!(
785-
M,
786-
d.update.memory_y[i],
787-
p_old,
788-
d.update.memory_y[i],
789-
p,
790-
d.update.vector_transport_method,
780+
M, d.update.memory_y[i], p_old, d.update.memory_y[i], p, d.update.vector_transport_method,
791781
)
792782
end
793783
end

test/solvers/test_primal_dual_semismooth_Newton.jl

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,8 @@
1-
using Manopt, Manifolds, ManifoldsBase, Test, RecursiveArrayTools
2-
using ManifoldDiff: differential_shortest_geodesic_startpoint
1+
s = joinpath(@__DIR__, "..", "ManoptTestSuite.jl")
2+
!(s in LOAD_PATH) && (push!(LOAD_PATH, s))
3+
4+
using Manopt, Manifolds, ManifoldsBase, ManifoldDiff, ManoptTestSuite, Test, RecursiveArrayTools
5+
using ManifoldDiff: differential_shortest_geodesic_startpoint, prox_distance
36

47
@testset "PD-RSSN" begin
58
#

0 commit comments

Comments
 (0)