diff --git a/pyproximal/optimization/cls_primal.py b/pyproximal/optimization/cls_primal.py index 6c8ba99..9a21e83 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,41 +621,48 @@ 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 - epsg_prev = self.epsg + epsg = self.epsg.item() + epsg_prev = self.epsg.item() 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: 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,16 +1082,22 @@ 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 - epsg_prev = self.epsg + epsg = self.epsg.item() + epsg_prev = self.epsg.item() 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) + 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 @@ -1398,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 @@ -1525,9 +1535,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,9 +1925,9 @@ 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] + 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 83bd039..3717cbf 100644 --- a/pyproximal/optimization/cls_primaldual.py +++ b/pyproximal/optimization/cls_primaldual.py @@ -295,15 +295,15 @@ 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] + 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] + mu = self.mu[self.iiter].item() if self.gfirst: y = self.proxg.proxdual(y + mu * self.A.matvec(xhat), mu) diff --git a/pyproximal/optimization/primal.py b/pyproximal/optimization/primal.py index a7c6f07..36abf03 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/pyproximal/proximal/L2.py b/pyproximal/proximal/L2.py index f38f53b..d63bd71 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), ], diff --git a/pytests/test_solver.py b/pytests/test_solver.py index 625c0c9..7200a24 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