From 828a8e488c405ca6f07e497bace0d94ea34e8887 Mon Sep 17 00:00:00 2001 From: Tamas Hakkel Date: Tue, 22 Apr 2025 21:15:39 +0200 Subject: [PATCH 1/4] add RecursiveArrayToolsExt to support ArrayPartitions # Conflicts: # Project.toml --- Project.toml | 7 +++++++ ext/RecursiveArrayToolsExt.jl | 19 +++++++++++++++++++ 2 files changed, 26 insertions(+) create mode 100644 ext/RecursiveArrayToolsExt.jl diff --git a/Project.toml b/Project.toml index 643559de..9c1a305c 100644 --- a/Project.toml +++ b/Project.toml @@ -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" diff --git a/ext/RecursiveArrayToolsExt.jl b/ext/RecursiveArrayToolsExt.jl new file mode 100644 index 00000000..0a92839e --- /dev/null +++ b/ext/RecursiveArrayToolsExt.jl @@ -0,0 +1,19 @@ +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) +gradient(g::SeparableSum, xs::ArrayPartition) = gradient(g, xs.x) + +end # module RecursiveArrayToolsExt From f85e1b137dd94d8ff2eaa853eb64c76bf85983d4 Mon Sep 17 00:00:00 2001 From: Tamas Hakkel Date: Wed, 1 Jul 2026 00:36:57 +0200 Subject: [PATCH 2/4] Add tests and fix bugs for RecursiveArrayToolsExt --- ext/RecursiveArrayToolsExt.jl | 25 ++++++-- test/Project.toml | 1 + test/runtests.jl | 1 + test/test_recursivearraytools.jl | 99 ++++++++++++++++++++++++++++++++ 4 files changed, 121 insertions(+), 5 deletions(-) create mode 100644 test/test_recursivearraytools.jl diff --git a/ext/RecursiveArrayToolsExt.jl b/ext/RecursiveArrayToolsExt.jl index 0a92839e..dfc0cfd1 100644 --- a/ext/RecursiveArrayToolsExt.jl +++ b/ext/RecursiveArrayToolsExt.jl @@ -4,16 +4,31 @@ 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) +function prox!(y::ArrayPartition, f::PrecomposedSlicedSeparableSum, x::ArrayPartition, gamma) + _, fy = prox!(y.x, f, x.x, gamma) + return y, fy +end (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!(ys::ArrayPartition, g::SeparableSum, xs::ArrayPartition, gamma::Number) + _, fy = prox!(ys.x, g, xs.x, gamma) + return ys, fy +end +function prox!(ys::ArrayPartition, g::SeparableSum, xs::ArrayPartition, gammas::Tuple) + _, fy = prox!(ys.x, g, xs.x, gammas) + return ys, fy +end 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) -gradient(g::SeparableSum, xs::ArrayPartition) = gradient(g, xs.x) +function gradient!(grads::ArrayPartition, g::SeparableSum, xs::ArrayPartition) + _, f_val = gradient!(grads.x, g, xs.x) + return grads, f_val +end +function gradient(g::SeparableSum, xs::ArrayPartition) + grad, f_val = gradient(g, xs.x) + return ArrayPartition(grad), f_val +end end # module RecursiveArrayToolsExt diff --git a/test/Project.toml b/test/Project.toml index e064bdc3..07c19fe7 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -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" diff --git a/test/runtests.jl b/test/runtests.jl index 04474040..845daec2 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -156,6 +156,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 diff --git a/test/test_recursivearraytools.jl b/test/test_recursivearraytools.jl new file mode 100644 index 00000000..2786d34a --- /dev/null +++ b/test/test_recursivearraytools.jl @@ -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 From bb3229d2b270f9f315a8995e690ff0a37d34d2b4 Mon Sep 17 00:00:00 2001 From: Tamas Hakkel Date: Wed, 1 Jul 2026 09:30:20 +0200 Subject: [PATCH 3/4] Fix Aqua tests for recursive-array-tools-ext --- test/runtests.jl | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/test/runtests.jl b/test/runtests.jl index 845daec2..99dfdd5d 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -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 From b9f6cf88186fef2ab6234cd2915ae24abea54833 Mon Sep 17 00:00:00 2001 From: Tamas Hakkel Date: Thu, 2 Jul 2026 00:54:55 +0200 Subject: [PATCH 4/4] Revert incorrect prox! and gradient! wrapper fixes, keep gradient wrapper --- ext/RecursiveArrayToolsExt.jl | 20 ++++---------------- 1 file changed, 4 insertions(+), 16 deletions(-) diff --git a/ext/RecursiveArrayToolsExt.jl b/ext/RecursiveArrayToolsExt.jl index dfc0cfd1..c46aa027 100644 --- a/ext/RecursiveArrayToolsExt.jl +++ b/ext/RecursiveArrayToolsExt.jl @@ -4,28 +4,16 @@ using ProximalOperators import ProximalCore: prox, prox!, gradient, gradient! (f::PrecomposedSlicedSeparableSum)(x::ArrayPartition) = f(x.x) -function prox!(y::ArrayPartition, f::PrecomposedSlicedSeparableSum, x::ArrayPartition, gamma) - _, fy = prox!(y.x, f, x.x, gamma) - return y, fy -end +prox!(y::ArrayPartition, f::PrecomposedSlicedSeparableSum, x::ArrayPartition, gamma) = prox!(y.x, f, x.x, gamma) (g::SeparableSum)(xs::ArrayPartition) = g(xs.x) -function prox!(ys::ArrayPartition, g::SeparableSum, xs::ArrayPartition, gamma::Number) - _, fy = prox!(ys.x, g, xs.x, gamma) - return ys, fy -end -function prox!(ys::ArrayPartition, g::SeparableSum, xs::ArrayPartition, gammas::Tuple) - _, fy = prox!(ys.x, g, xs.x, gammas) - return ys, fy -end +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 -function gradient!(grads::ArrayPartition, g::SeparableSum, xs::ArrayPartition) - _, f_val = gradient!(grads.x, g, xs.x) - return grads, f_val -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