Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,18 @@ ProximalCore = "dc4f5ac2-75d1-4f31-931e-60435d74994b"
SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf"
TSVD = "9449cd9e-2762-5aa3-a617-5413e99d722e"

[weakdeps]
RecursiveArrayTools = "731186ca-8d62-57ce-b412-fbd966d074cd"

[extensions]
RecursiveArrayToolsExt = "RecursiveArrayTools"

[compat]
IterativeSolvers = "0.8 - 0.9"
LinearAlgebra = "1.4"
OSQP = "0.3 - 0.8"
ProximalCore = "0.2.0"
RecursiveArrayTools = "2, 3"
SparseArrays = "1.4"
TSVD = "0.3 - 0.4"
julia = "1.9"
22 changes: 22 additions & 0 deletions ext/RecursiveArrayToolsExt.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
module RecursiveArrayToolsExt
using RecursiveArrayTools
using ProximalOperators
import ProximalCore: prox, prox!, gradient, gradient!

(f::PrecomposedSlicedSeparableSum)(x::ArrayPartition) = f(x.x)
prox!(y::ArrayPartition, f::PrecomposedSlicedSeparableSum, x::ArrayPartition, gamma) = prox!(y.x, f, x.x, gamma)

(g::SeparableSum)(xs::ArrayPartition) = g(xs.x)
prox!(ys::ArrayPartition, g::SeparableSum, xs::ArrayPartition, gamma::Number) = prox!(ys.x, g, xs.x, gamma)
prox!(ys::ArrayPartition, g::SeparableSum, xs::ArrayPartition, gammas::Tuple) = prox!(ys.x, g, xs.x, gammas)
function prox(g::SeparableSum, xs::ArrayPartition, gamma=1)
y, fy = prox(g, xs.x, gamma)
return ArrayPartition(y), fy
end
gradient!(grads::ArrayPartition, g::SeparableSum, xs::ArrayPartition) = gradient!(grads.x, g, xs.x)
function gradient(g::SeparableSum, xs::ArrayPartition)
grad, f_val = gradient(g, xs.x)
return ArrayPartition(grad), f_val
end

end # module RecursiveArrayToolsExt
1 change: 1 addition & 0 deletions test/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -5,5 +5,6 @@ LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"
ProximalCore = "dc4f5ac2-75d1-4f31-931e-60435d74994b"
ProximalOperators = "a725b495-10eb-56fe-b38b-717eba820537"
Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c"
RecursiveArrayTools = "731186ca-8d62-57ce-b412-fbd966d074cd"
SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf"
Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40"
4 changes: 3 additions & 1 deletion test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -114,7 +114,8 @@ function predicates_test(f)
end

@testset "Aqua" begin
Aqua.test_all(ProximalOperators; ambiguities=false)
Aqua.test_all(ProximalOperators; ambiguities=false, stale_deps=false, persistent_tasks=false)
Aqua.test_stale_deps(ProximalOperators, ignore=[:OSQP])
end

@testset "Documentation" begin
Expand Down Expand Up @@ -156,6 +157,7 @@ end
include("test_separableSum.jl")
include("test_slicedSeparableSum.jl")
include("test_precomposedSlicedSeparableSum.jl")
include("test_recursivearraytools.jl")
include("test_sum.jl")
include("test_reshapeInput.jl")
end
Expand Down
99 changes: 99 additions & 0 deletions test/test_recursivearraytools.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
using Test
using Random
using LinearAlgebra
using ProximalOperators
using RecursiveArrayTools

Random.seed!(1234)

@testset "RecursiveArrayToolsExt SeparableSum" begin
x = randn(10)
X = randn(10,10) .+ im*randn(10,10)

lambdas = (abs.(randn(size(x))), 0.1)
prox_col = (NormL1(lambdas[1]), NormL2(lambdas[2]))

f = SeparableSum(prox_col)

# Test on ArrayPartition
xs = ArrayPartition(x, X)

# evaluation
@test f(xs) == f((x, X))

# prox
ys, fys = prox(f, xs, 1.0)
y, fy = prox(f, (x, X), 1.0)
@test fys == fy
@test ys.x[1] == y[1]
@test ys.x[2] == y[2]
@test typeof(ys) <: ArrayPartition

# prox!
ys_mut = ArrayPartition(similar(x), similar(X))
prox!(ys_mut, f, xs, 1.0)
@test ys_mut.x[1] == y[1]
@test ys_mut.x[2] == y[2]

# prox! with multiple gammas
ys_mut2 = ArrayPartition(similar(x), similar(X))
gammas = (0.5, 1.3)
prox!(ys_mut2, f, xs, gammas)
y_g, fy_g = prox(f, (x, X), gammas)
@test ys_mut2.x[1] == y_g[1]
@test ys_mut2.x[2] == y_g[2]

# gradient
fs = (SqrNormL2(), LeastSquares(randn(5,10), randn(5)))
f2 = SeparableSum(fs)
x1, x2 = randn(10), randn(10)
xs2 = ArrayPartition(x1, x2)

grad_xs2, f_xs2 = gradient(f2, xs2)
grad_x, f_x = gradient(f2, (x1, x2))

@test f_xs2 == f_x
@test grad_xs2.x[1] == grad_x[1]
@test grad_xs2.x[2] == grad_x[2]
@test typeof(grad_xs2) <: ArrayPartition

# gradient!
grad_xs2_mut = ArrayPartition(similar(x1), similar(x2))
gradient!(grad_xs2_mut, f2, xs2)
@test grad_xs2_mut.x[1] == grad_x[1]
@test grad_xs2_mut.x[2] == grad_x[2]
end

@testset "RecursiveArrayToolsExt PrecomposedSlicedSeparableSum" begin
fs = (NormL1(), NormL2(), SqrNormL2())

A1 = (Diagonal(ones(10)), nothing)
F = qr(randn(5, 5))
A2 = (nothing, Matrix(F.Q))
F = qr(randn(5, 5))
A3 = (nothing, Matrix(F.Q))
mu = rand(5)
A3[2] .*= reshape(mu, 5, 1)
ops = (A1, A2, A3)

idxs = ((Colon(), nothing), (nothing, 1:5), (nothing, 6:10))
μs = (1.0, 1.0, mu)

f = PrecomposedSlicedSeparableSum(fs, idxs, ops, μs)
x_tuple = (randn(10), rand(10))
xs = ArrayPartition(x_tuple...)

# evaluation
@test f(xs) == f(x_tuple)

# prox!
ys_mut = ArrayPartition(zeros(10), zeros(10))
ys_tuple_mut = (zeros(10), zeros(10))

fy_xs = prox!(ys_mut, f, xs, 1.0)
fy_tuple = prox!(ys_tuple_mut, f, x_tuple, 1.0)

@test fy_xs == fy_tuple
@test ys_mut.x[1] == ys_tuple_mut[1]
@test ys_mut.x[2] == ys_tuple_mut[2]
end
Loading