diff --git a/pyproximal/ProxOperator.py b/pyproximal/ProxOperator.py index f32d07c..0a34c6f 100644 --- a/pyproximal/ProxOperator.py +++ b/pyproximal/ProxOperator.py @@ -407,7 +407,7 @@ def __init__(self, f: ProxOperator, a: float, b: float | NDArray) -> None: self.f, self.a, self.b = f, a, b super().__init__(None, f.hasgrad) - def __call__(self, x: NDArray) -> NDArray: + def __call__(self, x: NDArray) -> bool | float | int: return self.f(self.a * x + self.b) @_check_tau diff --git a/pyproximal/proximal/L21.py b/pyproximal/proximal/L21.py index bf59462..03ad8f2 100644 --- a/pyproximal/proximal/L21.py +++ b/pyproximal/proximal/L21.py @@ -28,7 +28,7 @@ class L21(ProxOperator): .. math:: \sigma \|\mathbf{X}\|_{2,1} = \sigma \sum_{j=0}^{N'_x} \|\mathbf{x}_j\|_2 = - \sigma \sum_{j=0}^{N'_x} \sqrt{\sum_{i=0}^{N_{dim}}} |x_{ij}|^2 + \sigma \sum_{j=0}^{N'_x} \sqrt{\sum_{i=0}^{N_{dim}} |x_{ij}|^2} the proximal operator is: @@ -65,13 +65,14 @@ def __init__(self, ndim: int, sigma: float = 1.0) -> None: def __call__(self, x: NDArray) -> float: x = x.reshape(self.ndim, len(x) // self.ndim) - f = self.sigma * np.sum(np.sqrt(np.sum(x**2, axis=0))) + f = self.sigma * np.sum(np.sqrt(np.sum(np.abs(x) ** 2, axis=0))) + print(f"f = {f}") return float(f) @_check_tau def prox(self, x: NDArray, tau: float) -> NDArray: x = x.reshape(self.ndim, len(x) // self.ndim) - aux = np.sqrt(np.sum(x**2, axis=0)) + aux = np.sqrt(np.sum(np.abs(x) ** 2, axis=0)) aux = np.vstack([aux] * self.ndim).ravel() x = (1 - (tau * self.sigma) / np.maximum(aux, tau * self.sigma)) * x.ravel() return x @@ -79,7 +80,7 @@ def prox(self, x: NDArray, tau: float) -> NDArray: @_check_tau def proxdual(self, x: NDArray, tau: float) -> NDArray: x = x.reshape(self.ndim, len(x) // self.ndim) - aux = np.sqrt(np.sum(x**2, axis=0)) + aux = np.sqrt(np.sum(np.abs(x) ** 2, axis=0)) aux = np.vstack([aux] * self.ndim).ravel() x = self.sigma * x.ravel() / np.maximum(aux, self.sigma) return x diff --git a/pyproximal/proximal/L21_plus_L1.py b/pyproximal/proximal/L21_plus_L1.py index 10dcbd4..fe6c90d 100644 --- a/pyproximal/proximal/L21_plus_L1.py +++ b/pyproximal/proximal/L21_plus_L1.py @@ -38,7 +38,9 @@ def __init__(self, sigma: float = 1.0, rho: float = 0.8) -> None: def __call__(self, x: NDArray) -> float: return float( self.rho * self.sigma * np.sum(np.abs(x)) - + (1 - self.rho) * self.sigma * np.sum(np.sqrt(np.sum(x**2, axis=0))) + + (1 - self.rho) + * self.sigma + * np.sum(np.sqrt(np.sum(np.abs(x) ** 2, axis=0))) ) @_check_tau