From 3b54fda201c678b7eb8914f362cc1d2e8bfb9521 Mon Sep 17 00:00:00 2001 From: mrava87 Date: Fri, 28 Aug 2026 22:12:30 +0100 Subject: [PATCH 1/4] fix: ensure step sizes and tau are passed as scalars --- pyproximal/optimization/cls_primal.py | 10 +++++----- pyproximal/optimization/cls_primaldual.py | 4 ++-- pyproximal/proximal/L2.py | 7 ++++--- 3 files changed, 11 insertions(+), 10 deletions(-) diff --git a/pyproximal/optimization/cls_primal.py b/pyproximal/optimization/cls_primal.py index 6c8ba991..287680e8 100644 --- a/pyproximal/optimization/cls_primal.py +++ b/pyproximal/optimization/cls_primal.py @@ -626,8 +626,8 @@ def step( epsg = self.epsg epsg_prev = self.epsg else: - epsg = self.epsg[self.iiter] - epsg_prev = self.epsg[self.iiter - 1] + epsg = self.epsg[self.iiter].item() + epsg_prev = self.epsg[self.iiter - 1].item() # proximal step if not self.backtracking: @@ -1080,8 +1080,8 @@ def step( epsg = self.epsg epsg_prev = self.epsg else: - epsg = self.epsg[self.iiter] - epsg_prev = self.epsg[self.iiter - 1] + epsg = self.epsg[self.iiter].item() + epsg_prev = self.epsg[self.iiter - 1].item() # update fix point g = x - self.tau * self.proxf.grad(x) @@ -1915,7 +1915,7 @@ def step( if self.tau.ndim == 0: tau = self.tau else: - tau = self.tau[self.iiter] + tau = self.tau[self.iiter].item() # proximal steps if self.gfirst: diff --git a/pyproximal/optimization/cls_primaldual.py b/pyproximal/optimization/cls_primaldual.py index 83bd039b..412776d0 100644 --- a/pyproximal/optimization/cls_primaldual.py +++ b/pyproximal/optimization/cls_primaldual.py @@ -297,13 +297,13 @@ def step( if self.tau.ndim == 0: tau = self.tau else: - tau = self.tau[self.iiter] + tau = self.tau[self.iiter].item() # define mu for current iteration if self.mu.ndim == 0: mu = self.mu else: - mu = self.mu[self.iiter] + mu = self.mu[self.iiter].item() if self.gfirst: y = self.proxg.proxdual(y + mu * self.A.matvec(xhat), mu) diff --git a/pyproximal/proximal/L2.py b/pyproximal/proximal/L2.py index f38f53b0..d63bd719 100644 --- a/pyproximal/proximal/L2.py +++ b/pyproximal/proximal/L2.py @@ -1,4 +1,5 @@ from collections.abc import Callable +from math import sqrt from typing import TYPE_CHECKING, Any, Optional import numpy as np @@ -37,7 +38,7 @@ class L2(ProxOperator): Data vector q : :obj:`numpy.ndarray`, optional Dot vector - sigma : :obj:`int`, optional + sigma : :obj:`float`, optional Multiplicative coefficient of L2 norm alpha : :obj:`float`, optional Multiplicative coefficient of dot product @@ -267,8 +268,8 @@ def prox(self, x: NDArray, tau: float) -> NDArray: if self.q is not None: y -= tau * self.alpha * self.q x = regularized_inversion( - np.sqrt(tau * self.sigma) * self.Op, - np.sqrt(tau * self.sigma) * self.b, + sqrt(tau * self.sigma) * self.Op, + sqrt(tau * self.sigma) * self.b, [ Identity(self.Op.shape[1], dtype=self.Op.dtype), ], From ab2e7af63b4a72e9b2a08de4e19d92b08d46e75f Mon Sep 17 00:00:00 2001 From: mrava87 Date: Fri, 28 Aug 2026 22:31:02 +0100 Subject: [PATCH 2/4] minor: added more .item() --- pyproximal/optimization/cls_primal.py | 16 +++++++++------- pyproximal/optimization/cls_primaldual.py | 4 ++-- 2 files changed, 11 insertions(+), 9 deletions(-) diff --git a/pyproximal/optimization/cls_primal.py b/pyproximal/optimization/cls_primal.py index 287680e8..1d91df01 100644 --- a/pyproximal/optimization/cls_primal.py +++ b/pyproximal/optimization/cls_primal.py @@ -623,8 +623,8 @@ def step( # define epsg for current iteration if self.epsg.ndim == 0: - epsg = self.epsg - epsg_prev = self.epsg + epsg = self.epsg.item() + epsg_prev = self.epsg.item() else: epsg = self.epsg[self.iiter].item() epsg_prev = self.epsg[self.iiter - 1].item() @@ -1077,8 +1077,8 @@ def step( """ # define epsg for current iteration if self.epsg.ndim == 0: - epsg = self.epsg - epsg_prev = self.epsg + epsg = self.epsg.item() + epsg_prev = self.epsg.item() else: epsg = self.epsg[self.iiter].item() epsg_prev = self.epsg[self.iiter - 1].item() @@ -1525,9 +1525,11 @@ def step( x = np.zeros_like(x) for i, proxg in enumerate(self.proxgs): ztmp = 2 * y - self.zs[i] - self.tau * grad - ztmp = proxg.prox(ztmp, self.tau * self.epsg[i] / self.weights[i]) + ztmp = proxg.prox( + ztmp, self.tau * self.epsg[i].item() / self.weights[i].item() + ) self.zs[i] += self.eta * (ztmp - y) - x += self.weights[i] * self.zs[i] + x += self.weights[i].item() * self.zs[i] # update y if self.acceleration == "vandenberghe": @@ -1913,7 +1915,7 @@ def step( """ # define tau for current iteration if self.tau.ndim == 0: - tau = self.tau + tau = self.tau.item() else: tau = self.tau[self.iiter].item() diff --git a/pyproximal/optimization/cls_primaldual.py b/pyproximal/optimization/cls_primaldual.py index 412776d0..3717cbfd 100644 --- a/pyproximal/optimization/cls_primaldual.py +++ b/pyproximal/optimization/cls_primaldual.py @@ -295,13 +295,13 @@ def step( # define tau for current iteration if self.tau.ndim == 0: - tau = self.tau + tau = self.tau.item() else: tau = self.tau[self.iiter].item() # define mu for current iteration if self.mu.ndim == 0: - mu = self.mu + mu = self.mu.item() else: mu = self.mu[self.iiter].item() From f9540500a8269ba814ebd95e1949e947e12453aa Mon Sep 17 00:00:00 2001 From: mrava87 Date: Fri, 28 Aug 2026 22:57:50 +0100 Subject: [PATCH 3/4] fix: fixed .item() issue also for proximal gradient solvers --- pyproximal/optimization/cls_primal.py | 42 +++++++++++++++++---------- 1 file changed, 26 insertions(+), 16 deletions(-) diff --git a/pyproximal/optimization/cls_primal.py b/pyproximal/optimization/cls_primal.py index 1d91df01..5f12e5b3 100644 --- a/pyproximal/optimization/cls_primal.py +++ b/pyproximal/optimization/cls_primal.py @@ -15,7 +15,7 @@ from collections.abc import Sequence from math import sqrt -from typing import TYPE_CHECKING, Any, Optional, cast +from typing import TYPE_CHECKING, Any, Optional import numpy as np import pylops @@ -621,6 +621,12 @@ def step( """ xold = x.copy() + # define tau for current iteration + if self.tau.size == 1: + tau = self.tau.item() + else: + tau = self.tau + # define epsg for current iteration if self.epsg.ndim == 0: epsg = self.epsg.item() @@ -632,30 +638,31 @@ def step( # proximal step if not self.backtracking: if self.eta == 1.0: - x = self.proxg.prox(y - self.tau * self.proxf.grad(y), epsg * self.tau) + x = self.proxg.prox(y - tau * self.proxf.grad(y), epsg * tau) else: x = x + self.eta * ( self.proxg.prox( - x - self.tau * self.proxf.grad(x), - epsg * self.tau, + x - tau * self.proxf.grad(x), + epsg * tau, ) - x ) else: - x, self.tau = _backtracking( + x, tau = _backtracking( y, - cast(float, self.tau), + tau, self.proxf, self.proxg, epsg, beta=self.beta, niterback=self.niterback, ) + self.tau = self.ncp.atleast_1d(self.ncp.asarray(tau, dtype=np.float32)) if self.eta != 1.0: x = x + self.eta * ( self.proxg.prox( - x - self.tau * self.proxf.grad(x), - epsg * self.tau, + x - tau * self.proxf.grad(x), + epsg * tau, ) - x ) @@ -1075,6 +1082,12 @@ def step( Updated additional model vector """ + # define tau for current iteration + if self.tau.size == 1: + tau = self.tau.item() + else: + tau = self.tau + # define epsg for current iteration if self.epsg.ndim == 0: epsg = self.epsg.item() @@ -1084,7 +1097,7 @@ def step( epsg_prev = self.epsg[self.iiter - 1].item() # update fix point - g = x - self.tau * self.proxf.grad(x) + g = x - tau * self.proxf.grad(x) r = g - y # update history vectors @@ -1108,24 +1121,21 @@ def step( y = np.vstack(self.G).T @ alpha # update main variable - x = self.proxg.prox(y, epsg * self.tau) + x = self.proxg.prox(y, epsg * tau) else: # update auxiliary variable ytest = np.vstack(self.G).T @ alpha # update main variable - xtest = self.proxg.prox(ytest, epsg * self.tau) + xtest = self.proxg.prox(ytest, epsg * tau) # check if function is decreased, otherwise do basic PG step pfold, self.pf = self.pf, self.proxf(xtest) - if ( - self.pf - <= pfold - self.tau * np.linalg.norm(self.proxf.grad(x)) ** 2 / 2 - ): + if self.pf <= pfold - tau * np.linalg.norm(self.proxf.grad(x)) ** 2 / 2: y = ytest x = xtest else: - x = self.proxg.prox(g, epsg * self.tau) + x = self.proxg.prox(g, epsg * tau) y = g # tolerance check: break iterations if overall From d634983d1d1efad88b2f6d8cc291ed7372bf0038 Mon Sep 17 00:00:00 2001 From: mrava87 Date: Fri, 28 Aug 2026 23:11:38 +0100 Subject: [PATCH 4/4] doc: clarify type of weights in GeneralizedProximalGradient --- pyproximal/optimization/cls_primal.py | 2 +- pyproximal/optimization/primal.py | 2 +- pytests/test_solver.py | 4 ++-- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/pyproximal/optimization/cls_primal.py b/pyproximal/optimization/cls_primal.py index 5f12e5b3..9a21e83a 100644 --- a/pyproximal/optimization/cls_primal.py +++ b/pyproximal/optimization/cls_primal.py @@ -1408,7 +1408,7 @@ def setup( # type: ignore[override] epsg : :obj:`float` or :obj:`numpy.ndarray`, optional Scaling factor(s) of ``g`` function(s). If a scalar is provided the same scaling factor is applied to every ``g`` function. - weights : :obj:`float`, optional + weights : :obj:`numpy.ndarray`, optional Weighting factors of ``g`` functions. Must sum to 1. eta : :obj:`float`, optional Relaxation parameter (must be between 0 and 1, 0 excluded). Note that diff --git a/pyproximal/optimization/primal.py b/pyproximal/optimization/primal.py index a7c6f07d..36abf037 100644 --- a/pyproximal/optimization/primal.py +++ b/pyproximal/optimization/primal.py @@ -461,7 +461,7 @@ def GeneralizedProximalGradient( epsg : :obj:`float` or :obj:`numpy.ndarray`, optional Scaling factor(s) of ``g`` function(s). If a scalar is provided the same scaling factor is applied to every ``g`` function. - weights : :obj:`float`, optional + weights : :obj:`numpy.ndarray`, optional Weighting factors of ``g`` functions. Must sum to 1. eta : :obj:`float`, optional Relaxation parameter (must be between 0 and 1, 0 excluded). Note that diff --git a/pytests/test_solver.py b/pytests/test_solver.py index 625c0c90..7200a248 100644 --- a/pytests/test_solver.py +++ b/pytests/test_solver.py @@ -147,7 +147,7 @@ def test_GPG_weights(par): ], x0=np.zeros(m), tau=1.0, - weights=[1.0, 1.0], + weights=np.array([1.0, 1.0]), ) @@ -701,7 +701,7 @@ def test_ADMM_DRS(par): @pytest.mark.parametrize("par", [(par1), (par2), (par3)]) -@pytest.mark.parametrize("weights", [None, (0.5, 0.5)]) +@pytest.mark.parametrize("weights", [None, np.array([0.5, 0.5])]) def test_PPXA_with_ADMM(par, weights) -> None: """Check equivalency of PPXA and ADMM when using a single regularization term