Skip to content
Merged
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
70 changes: 41 additions & 29 deletions pyproximal/optimization/cls_primal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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":
Expand Down Expand Up @@ -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:
Expand Down
8 changes: 4 additions & 4 deletions pyproximal/optimization/cls_primaldual.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion pyproximal/optimization/primal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
7 changes: 4 additions & 3 deletions pyproximal/proximal/L2.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
from collections.abc import Callable
from math import sqrt
from typing import TYPE_CHECKING, Any, Optional

import numpy as np
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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),
],
Expand Down
4 changes: 2 additions & 2 deletions pytests/test_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]),
)


Expand Down Expand Up @@ -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
Expand Down
Loading