Skip to content
Draft
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
1 change: 1 addition & 0 deletions docs/source/api/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ Templates
LinearOperator
FunctionOperator
MemoizeOperator
MultiOperator
PyTensorOperator
TorchOperator
JaxOperator
Expand Down
1 change: 1 addition & 0 deletions pylops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@

from .config import *
from .linearoperator import *
from .multioperator import *
from .torchoperator import *
from .pytensoroperator import *
from .jaxoperator import *
Expand Down
18 changes: 18 additions & 0 deletions pylops/_multioperator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,18 @@
import threading
from collections.abc import Callable

from pylops.utils.typing import NDArray


def _matvec_rmatvec_map(op: Callable[[NDArray], NDArray], x: NDArray) -> NDArray:
"""matvec/rmatvec for multiprocessing / multithreading"""
return op(x).squeeze()


def _matvec_rmatvec_map_mt(
op: Callable[[NDArray], NDArray], x: NDArray, y: NDArray, lock: threading.Lock
) -> None:
"""rmatvec for multithreading with lock"""
ylocal = op(x).squeeze()
with lock:
y[:] += ylocal
78 changes: 6 additions & 72 deletions pylops/basicoperators/blockdiag.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,5 @@
__all__ = ["BlockDiag"]

import concurrent.futures as mt
import multiprocessing as mp

import numpy as np
import scipy as sp

Expand All @@ -22,18 +19,14 @@

from collections.abc import Sequence

from pylops import LinearOperator
from pylops import LinearOperator, MultiOperator
from pylops._multioperator import _matvec_rmatvec_map
from pylops.basicoperators import MatrixMult
from pylops.utils.backend import get_array_module, get_module, inplace_set
from pylops.utils.typing import DTypeLike, NDArray, Tinoutengine, Tparallel_kind


def _matvec_rmatvec_map(op, x: NDArray) -> NDArray:
"""matvec/rmatvec for multiprocessing"""
return op(x).squeeze()


class BlockDiag(LinearOperator):
class BlockDiag(MultiOperator):
r"""Block-diagonal operator.

Create a block-diagonal operator from N linear operators.
Expand Down Expand Up @@ -149,9 +142,6 @@ def __init__(
parallel_kind: Tparallel_kind = "multiproc",
dtype: DTypeLike | None = None,
) -> None:
if parallel_kind not in ["multiproc", "multithread"]:
msg = "parallel_kind must be 'multiproc' or 'multithread'"
raise ValueError(msg)
# identify dimensions
self.ops = ops
mops = np.zeros(len(ops), dtype=int)
Expand Down Expand Up @@ -181,15 +171,10 @@ def __init__(
else:
dimsd = (self.nops,)
forceflat = True

# create pool for multithreading / multiprocessing
self.parallel_kind = parallel_kind
self._nproc = nproc
self.pool: mp.pool.Pool | None = None
if self.nproc > 1:
if self.parallel_kind == "multiproc":
self.pool = mp.Pool(processes=nproc)
else:
self.pool = mt.ThreadPoolExecutor(max_workers=nproc)
self._setup_pool(nproc, parallel_kind=parallel_kind)

self.inoutengine = inoutengine
dtype = _get_dtype(ops) if dtype is None else np.dtype(dtype)
clinear = all([getattr(oper, "clinear", True) for oper in self.ops])
Expand All @@ -201,25 +186,6 @@ def __init__(
forceflat=forceflat,
)

@property
def nproc(self) -> int:
return self._nproc

@nproc.setter
def nproc(self, nprocnew: int) -> None:
if self._nproc > 1 and self.pool is not None:
if self.parallel_kind == "multiproc":
self.pool.close()
self.pool.join()
else:
self.pool.shutdown()
if nprocnew > 1:
if self.parallel_kind == "multiproc":
self.pool = mp.Pool(processes=nprocnew)
else:
self.pool = mt.ThreadPoolExecutor(max_workers=nprocnew)
self._nproc = nprocnew

def _matvec_serial(self, x: NDArray) -> NDArray:
ncp = (
get_array_module(x)
Expand Down Expand Up @@ -297,35 +263,3 @@ def _rmatvec_multithread(self, x: NDArray) -> NDArray:
)
y = np.hstack(ys)
return y

def _matvec(self, x: NDArray) -> NDArray:
if self.nproc == 1:
y = self._matvec_serial(x)
else:
if self.parallel_kind == "multiproc":
y = self._matvec_multiproc(x)
else:
y = self._matvec_multithread(x)
return y

def _rmatvec(self, x: NDArray) -> NDArray:
if self.nproc == 1:
y = self._rmatvec_serial(x)
else:
if self.parallel_kind == "multiproc":
y = self._rmatvec_multiproc(x)
else:
y = self._rmatvec_multithread(x)
return y

def close(self):
"""Close the pool of workers used for multiprocessing
/ multithreading.
"""
if self.pool is not None:
if self.parallel_kind == "multiproc":
self.pool.close()
self.pool.join()
else:
self.pool.shutdown()
self.pool = None
91 changes: 8 additions & 83 deletions pylops/basicoperators/hstack.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,5 @@
__all__ = ["HStack"]

import concurrent.futures as mt
import multiprocessing as mp
import threading

import numpy as np
import scipy as sp

Expand All @@ -19,27 +15,16 @@
)
from scipy.sparse.linalg._interface import _get_dtype

from collections.abc import Callable, Sequence
from collections.abc import Sequence

from pylops import LinearOperator
from pylops import LinearOperator, MultiOperator
from pylops._multioperator import _matvec_rmatvec_map, _matvec_rmatvec_map_mt
from pylops.basicoperators import MatrixMult, Zero
from pylops.utils.backend import get_array_module, get_module, inplace_add, inplace_set
from pylops.utils.typing import NDArray, Tinoutengine, Tparallel_kind


def _matvec_rmatvec_map(op, x: NDArray) -> NDArray:
"""matvec/rmatvec for multiprocessing"""
return op(x).squeeze()


def _matvec_map_mt(op: Callable, x: NDArray, y: NDArray, lock: threading.Lock) -> None:
"""rmatvec for multithreading with lock"""
ylocal = op(x).squeeze()
with lock:
y[:] += ylocal


class HStack(LinearOperator):
class HStack(MultiOperator):
r"""Horizontal stacking.

Stack a set of N linear operators horizontally. Note that in case
Expand Down Expand Up @@ -157,9 +142,6 @@ def __init__(
parallel_kind: Tparallel_kind = "multiproc",
dtype: str | None = None,
) -> None:
if parallel_kind not in ["multiproc", "multithread"]:
msg = "parallel_kind must be 'multiproc' or 'multithread'"
raise ValueError(msg)
# identify dimensions
self.ops = ops
mops = np.zeros(len(ops), dtype=int)
Expand All @@ -182,16 +164,10 @@ def __init__(
else:
dimsd = (self.nops,)
forceflat = True

# create pool for multithreading / multiprocessing
self.parallel_kind = parallel_kind
self._nproc = nproc
self.pool = None
if self.nproc > 1:
if self.parallel_kind == "multiproc":
self.pool = mp.Pool(processes=nproc)
else:
self.pool = mt.ThreadPoolExecutor(max_workers=nproc)
self.lock = threading.Lock()
self._setup_pool(nproc, parallel_kind=parallel_kind)

self.inoutengine = inoutengine
dtype = _get_dtype(self.ops) if dtype is None else np.dtype(dtype)
clinear = all([getattr(oper, "clinear", True) for oper in self.ops])
Expand All @@ -203,25 +179,6 @@ def __init__(
forceflat=forceflat,
)

@property
def nproc(self) -> int:
return self._nproc

@nproc.setter
def nproc(self, nprocnew: int):
if self._nproc > 1 and self.pool is not None:
if self.parallel_kind == "multiproc":
self.pool.close()
self.pool.join()
else:
self.pool.shutdown()
if nprocnew > 1:
if self.parallel_kind == "multiproc":
self.pool = mp.Pool(processes=nprocnew)
else:
self.pool = mt.ThreadPoolExecutor(max_workers=nprocnew)
self._nproc = nprocnew

def _matvec_serial(self, x: NDArray) -> NDArray:
ncp = (
get_array_module(x)
Expand Down Expand Up @@ -277,7 +234,7 @@ def _matvec_multithread(self, x: NDArray) -> NDArray:
y = np.zeros(self.nops, dtype=self.dtype)
list(
self.pool.map(
lambda args: _matvec_map_mt(*args),
lambda args: _matvec_rmatvec_map_mt(*args),
[
(
oper._matvec,
Expand All @@ -300,35 +257,3 @@ def _rmatvec_multithread(self, x: NDArray) -> NDArray:
)
y = np.hstack(ys)
return y

def _matvec(self, x: NDArray) -> NDArray:
if self.nproc == 1:
y = self._matvec_serial(x)
else:
if self.parallel_kind == "multiproc":
y = self._matvec_multiproc(x)
else:
y = self._matvec_multithread(x)
return y

def _rmatvec(self, x: NDArray) -> NDArray:
if self.nproc == 1:
y = self._rmatvec_serial(x)
else:
if self.parallel_kind == "multiproc":
y = self._rmatvec_multiproc(x)
else:
y = self._rmatvec_multithread(x)
return y

def close(self):
"""Close the pool of workers used for multiprocessing /
multithreading.
"""
if self.pool is not None:
if self.parallel_kind == "multiproc":
self.pool.close()
self.pool.join()
else:
self.pool.shutdown()
self.pool = None
41 changes: 32 additions & 9 deletions pylops/basicoperators/kronecker.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,17 @@
__all__ = ["Kronecker"]

from typing import TYPE_CHECKING

import numpy as np

from pylops import LinearOperator
from pylops.utils.typing import DTypeLike, NDArray
from pylops import MultiOperator
from pylops.utils.typing import DTypeLike, NDArray, Tparallel_kind

if TYPE_CHECKING:
from pylops.linearoperator import LinearOperator


class Kronecker(LinearOperator):
class Kronecker(MultiOperator):
r"""Kronecker operator.

Perform Kronecker product of two operators. Note that the combined operator
Expand All @@ -22,6 +27,18 @@ class Kronecker(LinearOperator):
Second operator
dtype : :obj:`str`, optional
Type of elements in input array.
nproc : :obj:`int`, optional
.. versionadded:: 2.9.0

Number of processes/threads used to evaluate the N operators in parallel
using ``multiprocessing``/``concurrent.futures``. If ``nproc=1``, work in serial mode.
parallel_kind : :obj:`str`, optional
.. versionadded:: 2.9.0

Parallelism kind when ``nproc>1``. Can be ``multiproc`` (using
:mod:`multiprocessing`) or ``multithread`` (using
:class:`concurrent.futures.ThreadPoolExecutor`). Defaults
to ``multiproc``.
name : :obj:`str`, optional
.. versionadded:: 2.0.0

Expand Down Expand Up @@ -65,15 +82,21 @@ class Kronecker(LinearOperator):

def __init__(
self,
Op1: LinearOperator,
Op2: LinearOperator,
Op1: "LinearOperator",
Op2: "LinearOperator",
nproc: int = 1,
parallel_kind: Tparallel_kind = "multiproc",
dtype: DTypeLike = "float64",
name: str = "K",
) -> None:
self.Op1 = Op1
self.Op2 = Op2
self.Op1H = self.Op1.H
self.Op2H = self.Op2.H

# create pool for multithreading / multiprocessing
self._setup_pool(nproc, parallel_kind=parallel_kind)

shape = (
self.Op1.shape[0] * self.Op2.shape[0],
self.Op1.shape[1] * self.Op2.shape[1],
Expand All @@ -82,12 +105,12 @@ def __init__(

def _matvec(self, x: NDArray) -> NDArray:
x = x.reshape(self.Op1.shape[1], self.Op2.shape[1])
y = self.Op2.matmat(x.T).T
y = self.Op1.matmat(y).ravel()
y = self.Op2.matmat(x.T, pool=self.pool).T
y = self.Op1.matmat(y, pool=self.pool).ravel()
return y

def _rmatvec(self, x: NDArray) -> NDArray:
x = x.reshape(self.Op1.shape[0], self.Op2.shape[0])
y = self.Op2H.matmat(x.T).T
y = self.Op1H.matmat(y).ravel()
y = self.Op2H.matmat(x.T, pool=self.pool).T
y = self.Op1H.matmat(y, pool=self.pool).ravel()
return y
Loading
Loading