diff --git a/feectools/accelerate/compile_psydac.mk b/feectools/accelerate/compile_psydac.mk index fc09fe48b..78d6b7804 100644 --- a/feectools/accelerate/compile_psydac.mk +++ b/feectools/accelerate/compile_psydac.mk @@ -37,7 +37,7 @@ all: $(OUTPUTS) @for dep in $^ ; do \ echo $$dep ; \ done - pyccel compile -v $(FLAGS)$(FLAGS_openmp) $< + pyccel compile $(FLAGS)$(FLAGS_openmp) $< @echo "" #-------------------------------------- diff --git a/feectools/api/fem_bilinear_form.py b/feectools/api/fem_bilinear_form.py index 684e74113..8457e9bbc 100644 --- a/feectools/api/fem_bilinear_form.py +++ b/feectools/api/fem_bilinear_form.py @@ -594,9 +594,9 @@ def allocate_matrices(self, backend=None): if is_conformal: matrix[k1, k2] = global_mats[k1, k2] elif use_restriction: - matrix.multiplicants[-1][k1, k2] = global_mats[k1, k2] + matrix.multiplicands[-1][k1, k2] = global_mats[k1, k2] elif use_prolongation: - matrix.multiplicants[0][k1, k2] = global_mats[k1, k2] + matrix.multiplicands[0][k1, k2] = global_mats[k1, k2] else: # case of scalar equation if is_broken: # multi-patch diff --git a/feectools/api/fem_common.py b/feectools/api/fem_common.py index a2e849267..f9b2c477f 100644 --- a/feectools/api/fem_common.py +++ b/feectools/api/fem_common.py @@ -277,7 +277,7 @@ def extract_stencil_mats(mats): if isinstance(M, (StencilInterfaceMatrix, StencilMatrix)): new_mats.append(M) elif isinstance(M, ComposedLinearOperator): - new_mats += [i for i in M.multiplicants if isinstance(i, (StencilInterfaceMatrix, StencilMatrix))] + new_mats += [i for i in M.multiplicands if isinstance(i, (StencilInterfaceMatrix, StencilMatrix))] return new_mats #============================================================================== diff --git a/feectools/feec/derivatives.py b/feectools/feec/derivatives.py index 6b12a7a43..4d69b1e7f 100644 --- a/feectools/feec/derivatives.py +++ b/feectools/feec/derivatives.py @@ -136,6 +136,21 @@ def codomain(self): def dtype( self ): return self.domain.dtype + @property + def diffdir(self) -> int: + """Direction (axis) of the derivative.""" + return self._diffdir + + @property + def negative(self) -> bool: + """Whether the operator is the negative derivative.""" + return self._negative + + @property + def transposed(self) -> bool: + """Whether the operator is the transposed derivative.""" + return self._transposed + def __truediv__(self, a): """ Divide by scalar. """ return self * (1.0 / a) diff --git a/feectools/linalg/basic.py b/feectools/linalg/basic.py index 0ddf67a9c..13eb76e97 100644 --- a/feectools/linalg/basic.py +++ b/feectools/linalg/basic.py @@ -8,6 +8,7 @@ """ import itertools +import warnings from abc import ABC, abstractmethod from types import LambdaType from inspect import signature @@ -1057,27 +1058,27 @@ def __init__(self, domain, codomain, *args): for i in range(len(args)-1): assert args[i].domain == args[i+1].codomain - multiplicants = () + multiplicands = () tmp_vectors = [] for a in args[:-1]: if isinstance(a, ComposedLinearOperator): - multiplicants = (*multiplicants, *a.multiplicants) + multiplicands = (*multiplicands, *a.multiplicands) tmp_vectors.extend(a.tmp_vectors) tmp_vectors.append(a.domain.zeros()) else: - multiplicants = (*multiplicants, a) + multiplicands = (*multiplicands, a) tmp_vectors.append(a.domain.zeros()) last = args[-1] if isinstance(last, ComposedLinearOperator): - multiplicants = (*multiplicants, *last.multiplicants) + multiplicands = (*multiplicands, *last.multiplicands) tmp_vectors.extend(last.tmp_vectors) else: - multiplicants = (*multiplicants, last) + multiplicands = (*multiplicands, last) self._domain = domain self._codomain = codomain - self._multiplicants = multiplicants + self._multiplicands = multiplicands self._tmp_vectors = tuple(tmp_vectors) @property @@ -1098,34 +1099,41 @@ def codomain(self): return self._codomain @property - def multiplicants(self): + def multiplicands(self): r""" - A tuple $(A_1,\dots,A_n)$ containing the multiplicants of the linear operator + A tuple $(A_1,\dots,A_n)$ containing the multiplicands of the linear operator $self = A_n\circ\dots\circ A_1$. """ - return self._multiplicants + return self._multiplicands + + @property + def multiplicants(self): + """Deprecated alias of ``multiplicands``.""" + warnings.warn("ComposedLinearOperator.multiplicants is deprecated, use multiplicands instead.", + DeprecationWarning, stacklevel=2) + return self._multiplicands @property def dtype(self): return None def tosparse(self): - mats = [M.tosparse() for M in self._multiplicants] + mats = [M.tosparse() for M in self._multiplicands] M = mats[0] for Mi in mats[1:]: M = M @ Mi return coo_matrix(M) def transpose(self, conjugate=False): - t_multiplicants = () - for a in self._multiplicants: - t_multiplicants = (a.transpose(conjugate=conjugate), *t_multiplicants) + t_multiplicands = () + for a in self._multiplicands: + t_multiplicands = (a.transpose(conjugate=conjugate), *t_multiplicands) new_dom = self.codomain new_cod = self.domain assert isinstance(new_dom, VectorSpace) assert isinstance(new_cod, VectorSpace) - return ComposedLinearOperator(self.codomain, self.domain, *t_multiplicants) + return ComposedLinearOperator(self.codomain, self.domain, *t_multiplicands) def dot(self, v, out=None): assert isinstance(v, Vector) @@ -1137,11 +1145,11 @@ def dot(self, v, out=None): x = v for i in range(len(self._tmp_vectors)): y = self._tmp_vectors[-1-i] - A = self._multiplicants[-1-i] + A = self._multiplicands[-1-i] A.dot(x, out=y) x = y - A = self._multiplicants[0] + A = self._multiplicands[0] if out is not None: A.dot(x, out=out) @@ -1150,11 +1158,11 @@ def dot(self, v, out=None): return out def exchange_assembly_data(self): - for op in self._multiplicants: + for op in self._multiplicands: op.exchange_assembly_data() def set_backend(self, backend, precompiled=False): - for op in self._multiplicants: + for op in self._multiplicands: op.set_backend(backend) #=============================================================================== diff --git a/feectools/linalg/kron.py b/feectools/linalg/kron.py index e2d76d771..343d88511 100644 --- a/feectools/linalg/kron.py +++ b/feectools/linalg/kron.py @@ -1,4 +1,7 @@ #coding = utf-8 +from __future__ import annotations + +from collections.abc import Sequence from functools import reduce import numpy as np @@ -7,19 +10,40 @@ from scipy.sparse import kron from scipy.sparse import coo_matrix -from feectools.linalg.basic import LinearOperator, LinearSolver +from feectools.ddm.cart import CartDecomposition +from feectools.linalg.basic import ComposedLinearOperator, LinearOperator, LinearSolver from feectools.linalg.stencil import StencilVectorSpace, StencilVector, StencilMatrix from feectools.linalg.direct_solvers import DenseInverse __all__ = ('KroneckerStencilMatrix', + 'ComposedKroneckerStencilMatrix', 'KroneckerLinearSolver', + 'KroneckerSumSolver', 'KroneckerDenseMatrix', 'kronecker_solve') #============================================================================== class KroneckerStencilMatrix(LinearOperator): - """ - Kronecker product of 1D stencil matrices. + r""" + Kronecker product $M = A_1 \otimes A_2 \otimes \dots \otimes A_m$ of stencil matrices. + + Each factor $A_k$ is a StencilMatrix acting on a group of consecutive + axes of the domain and codomain; its number of axes is ``A_k.domain.ndim``. + The axes of all factors must add up to ``V.ndim``. The usual case is one + 1d factor per axis, but other groupings are allowed, e.g. a 2d x 1d + product on a 3d space:: + + M = KroneckerStencilMatrix(V, W, A_xy, A_z) # M.axes == ((0, 1), (2,)) + + The factors are process-local: they live on their own spaces, without a + communicator, and along its axes the factor $A_k$ owns (at least) the rows + (``starts``/``ends``) of the codomain ``W`` on this process, e.g. exactly + these rows or all rows. The pads of the + factors must not exceed the pads of the domain ``V``, whose ghost regions + are read by ``dot``. + + The product ``A @ B`` of two Kronecker matrices with the same axis groups + is a :class:`ComposedKroneckerStencilMatrix`. Parameters ---------- @@ -29,93 +53,156 @@ class KroneckerStencilMatrix(LinearOperator): W : StencilVectorSpace The codomain. - args : list of StencilMatrix - Factors of the Kronecker product (one for each dimension). - + *args : StencilMatrix + Factors of the Kronecker product, ordered by axis. For each factor + ``A_k``, ``A_k.domain.npts`` and ``A_k.codomain.npts`` must equal the + ``npts`` of ``V`` and ``W`` along the axes of that factor. The pads + of ``A_k`` must not exceed the pads of ``V``. """ - def __init__(self, V, W, *args): + def __init__(self, V: StencilVectorSpace, W: StencilVectorSpace, *args: StencilMatrix): assert isinstance(V, StencilVectorSpace) assert isinstance(W, StencilVectorSpace) - - for i,A in enumerate(args): - assert isinstance(A, LinearOperator) - assert A.domain.ndim == 1 - assert A.domain.npts[0] == V.npts[i] + assert V.ndim == W.ndim + assert len(args) > 0, 'A KroneckerStencilMatrix needs at least one factor.' + + # group the axes of V and W by factor + axes = [] + d = 0 + for A in args: + assert isinstance(A, StencilMatrix), \ + f'Factors must be of type StencilMatrix, got {type(A)}.' + n = A.domain.ndim + grp = tuple(range(d, d + n)) + assert d + n <= V.ndim, \ + f'The factors have more axes than the domain ({V.ndim}).' + assert tuple(A.domain.npts) == tuple(V.npts[d:d+n]), \ + f'Domain npts {A.domain.npts} of factor on axes {grp} do not match {V.npts}.' + assert tuple(A.codomain.npts) == tuple(W.npts[d:d+n]), \ + f'Codomain npts {A.codomain.npts} of factor on axes {grp} do not match {W.npts}.' + axes.append(grp) + d += n + assert d == V.ndim, \ + f'The factors cover {d} axes, but the domain has {V.ndim}.' + + # dot reads the rows of W on this process from the factors + for A, grp in zip(args, axes): + for k, a in enumerate(grp): + if not A.codomain.starts[k] <= W.starts[a] <= W.ends[a] <= A.codomain.ends[k]: + raise ValueError(f'The factor on axes {grp} owns the rows {A.codomain.starts[k]}..' + f'{A.codomain.ends[k]} along axis {a}, which do not contain the rows ' + f'{W.starts[a]}..{W.ends[a]} of the codomain on this process.') + + # dot reads the ghost regions of x, so the band must fit into them + for A, grp in zip(args, axes): + for p, a in zip(A.pads, grp): + if p > V.pads[a]: + raise ValueError(f'Pads {A.pads} of factor on axes {grp} exceed the domain pads {V.pads}. ' + 'Products of Kronecker matrices are ComposedKroneckerStencilMatrix (A @ B).') self._domain = V self._codomain = W - self._mats = args - self._ndim = len(args) + self._mats = tuple(args) + self._axes = tuple(axes) #-------------------------------------- # Abstract interface #-------------------------------------- @property - def domain( self ): + def domain(self) -> StencilVectorSpace: return self._domain # ... @property - def codomain( self ): + def codomain(self) -> StencilVectorSpace: return self._codomain # ... @property - def dtype( self ): + def dtype(self): return self.domain.dtype # ... @property - def ndim( self ): - return self._ndim - + def ndim(self) -> int: + """Number of axes of the domain (not the number of factors, see ``axes``).""" + return self._domain.ndim + # ... @property - def mats( self ): + def mats(self) -> tuple[StencilMatrix, ...]: + """Factors of the Kronecker product, ordered by axis.""" return self._mats # ... @property - def nbytes( self ): - """Local (per-MPI-rank) memory footprint of the 1d factor matrices, in bytes.""" + def axes(self) -> tuple[tuple[int, ...], ...]: + """Axes of the domain/codomain on which each factor acts, e.g. ``((0, 1), (2,))``.""" + return self._axes + + # ... + @property + def nbytes(self) -> int: + """Local (per-MPI-rank) memory footprint of the factor matrices, in bytes.""" return int(sum(getattr(mat, 'nbytes', 0) for mat in self._mats)) # ... - def dot(self, x, out=None): + def dot(self, x: StencilVector, out: StencilVector | None = None) -> StencilVector: + """ + Matrix-vector product ``M @ x``. - dot = xp.dot + Parameters + ---------- + x : StencilVector + Vector in the domain. + + out : StencilVector, optional + Vector in the codomain, in which the result is stored. + + Returns + ------- + StencilVector + The result, in the codomain (``out`` if given). + """ assert isinstance(x, StencilVector) assert x.space is self.domain - # Necessary if vector space is periodic or distributed across processes - if not x.ghost_regions_in_sync: - x.update_ghost_regions() - if out is not None: assert isinstance(out, StencilVector) assert out.space is self.codomain else: out = StencilVector(self.codomain) + # Necessary if vector space is periodic or distributed across processes + if not x.ghost_regions_in_sync: + x.update_ghost_regions() + starts = self._codomain.starts ends = self._codomain.ends pads = self._codomain.pads shifts = self._codomain.shifts mats = self.mats + axes = self.axes nrows = tuple(e-s+1 for s,e in zip(starts, ends)) - pnrows = tuple(2*p+1 for p in pads) + + # per axis: band of the factor, row offset in its data array and ghost offset of x + mpads = tuple(p for A in mats for p in A.pads) + row_off = self._row_offsets() + x_off = tuple(p*m for p,m in zip(self._domain.pads, self._domain.shifts)) + pnrows = tuple(2*p+1 for p in mpads) for ii in xp.ndindex(*nrows): v = 0. xx = tuple(i+p*s for i,p,s in zip(ii, pads, shifts)) + rr = tuple(i+o for i,o in zip(ii, row_off)) for jj in xp.ndindex(*pnrows): - i_mats = [mat._data[s, j] for s,j,mat in zip(xx, jj, mats)] - ii_jj = tuple(i+j+(s-1)*p for i,j,p,s in zip(ii, jj, pads, shifts)) + i_mats = [mat._data[(*(rr[a] for a in grp), *(jj[a] for a in grp))] + for mat,grp in zip(mats, axes)] + ii_jj = tuple(i+j-p+o for i,j,p,o in zip(ii, jj, mpads, x_off)) # ``array_api_compat.cupy`` does not accept a Python list in # ``prod``; multiplying the scalar factors also avoids a # temporary device array in this innermost loop. @@ -128,45 +215,86 @@ def dot(self, x, out=None): return out # ... - def copy(self): + def copy(self) -> KroneckerStencilMatrix: mats = [m.copy() for m in self.mats] return KroneckerStencilMatrix(self.domain, self.codomain, *mats) # ... - def __neg__(self): + def __neg__(self) -> KroneckerStencilMatrix: mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])] return KroneckerStencilMatrix(self.domain, self.codomain, *mats) # ... - def __mul__(self, a): + def __mul__(self, a) -> KroneckerStencilMatrix: mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a] return KroneckerStencilMatrix(self.domain, self.codomain, *mats) # ... - def __imul__(self, a): - self.mats[-1] *= a + def __imul__(self, a) -> KroneckerStencilMatrix: + last = self._mats[-1] + last *= a return self + # ... + def __matmul__(self, B): + """ + Product ``self @ B``. + + If ``B`` is a KroneckerStencilMatrix or a ComposedKroneckerStencilMatrix + with the same axis groups, the result is a + :class:`ComposedKroneckerStencilMatrix`. In all other cases (a different + operator, other axis groups, or a vector) the call is passed to + ``LinearOperator.__matmul__``. + + Parameters + ---------- + B : LinearOperator | Vector + Right operand. Its codomain must be the domain of ``self``. + + Returns + ------- + ComposedKroneckerStencilMatrix | LinearOperator | Vector + The product. + """ + if _is_kronecker_with_axes(B, self.axes): + return ComposedKroneckerStencilMatrix(B.domain, self.codomain, self, B) + return super().__matmul__(B) + #-------------------------------------- # Other properties/methods #-------------------------------------- def __getitem__(self, key): - pads = self._codomain.pads + """ + Entry ``M[i_1, ..., i_d, k_1, ..., k_d]`` for row indices ``i`` and + diagonal offsets ``k``, i.e. the product of the corresponding entries + of the factors. + """ rows = key[:self.ndim] cols = key[self.ndim:] - mats = self.mats - elements = [A[i,j] for A,i,j in zip(mats, rows, cols)] + elements = [A[(*(rows[a] for a in grp), *(cols[a] for a in grp))] + for A,grp in zip(self.mats, self.axes)] return reduce(lambda a, b: a * b, elements, 1) - def tostencil(self): + def _row_offsets(self) -> tuple[int, ...]: + """ + Per axis: index in the data array of its factor of the first local row of the codomain + (ghost region of the factor plus the offset between the first rows of codomain and factor). + """ + return tuple(p*m + s - fs + for A, grp in zip(self.mats, self.axes) + for p, m, fs, s in zip(A.codomain.pads, A.codomain.shifts, A.codomain.starts, + (self._codomain.starts[a] for a in grp))) + + def tostencil(self) -> StencilMatrix: + """Convert to a StencilMatrix on the domain and codomain.""" mats = self.mats ssc = self.codomain.starts eec = self.codomain.ends ssd = self.domain.starts eed = self.domain.ends - pads = [A.pads[0] for A in self.mats] + pads = [p for A in self.mats for p in A.pads] xpads = self.domain.pads # Number of rows in matrix (along each dimension) @@ -176,13 +304,16 @@ def tostencil(self): # create the stencil matrix M = StencilMatrix(self.domain, self.codomain, pads=tuple(pads)) + # row offset of each axis in the data array of its factor + row_off = list(self._row_offsets()) + mats = [mat._data for mat in mats] - self._tostencil(M._data, mats, nrows, nrows_extra, pads, xpads) + self._tostencil(M._data, mats, self.axes, nrows, nrows_extra, pads, xpads, row_off) return M @staticmethod - def _tostencil(M, mats, nrows, nrows_extra, pads, xpads): + def _tostencil(M, mats, axes, nrows, nrows_extra, pads, xpads, row_off): ndiags = [2*p + 1 for p in pads] diff = [xp-p for xp,p in zip(xpads, pads)] @@ -191,12 +322,14 @@ def _tostencil(M, mats, nrows, nrows_extra, pads, xpads): for xx in xp.ndindex( *nrows ): ii = tuple(xp + x for xp, x in zip(xpads, xx) ) + rr = tuple(o + x for o, x in zip(row_off, xx) ) for kk in xp.ndindex( *ndiags ): - values = [mat[i,k] for mat,i,k in zip(mats, ii, kk)] + values = [mat[(*(rr[a] for a in grp), *(kk[a] for a in grp))] + for mat,grp in zip(mats, axes)] M[(*ii, *kk)] = reduce(lambda a, b: a * b, values, 1) - + # handle partly-multiplied rows new_nrows = nrows.copy() for d,er in enumerate(nrows_extra): @@ -210,6 +343,7 @@ def _tostencil(M, mats, nrows, nrows_extra, pads, xpads): xx.insert(d, nrows[d]+n) ii = tuple(x+xp for x,xp in zip(xx, xpads)) + rr = tuple(x+o for x,o in zip(xx, row_off)) ee = [max(x-l+1,0) for x,l in zip(xx, nrows)] jj = tuple( slice(x+d, x+d+2*p+1-e) for x,p,d,e in zip(xx, pads, diff, ee) ) ndiags = [2*p + 1-e for p,e in zip(pads,ee)] @@ -217,20 +351,246 @@ def _tostencil(M, mats, nrows, nrows_extra, pads, xpads): ii_kk = tuple( list(ii) + kk ) for kk in xp.ndindex( *ndiags ): - values = [mat[i,k] for mat,i,k in zip(mats, ii, kk)] + values = [mat[(*(rr[a] for a in grp), *(kk[a] for a in grp))] + for mat,grp in zip(mats, axes)] M[(*ii, *kk)] = reduce(lambda a, b: a * b, values, 1) new_nrows[d] += er def tosparse(self): + """Convert the local rows to a scipy sparse matrix (Kronecker product of the factors' ``tosparse``).""" return reduce(kron, (m.tosparse() for m in self.mats)) def toarray(self): + """Convert the local rows to a dense array.""" return self.tosparse().toarray() - def transpose(self, conjugate=False): + def transpose(self, conjugate: bool = False) -> KroneckerStencilMatrix: + """Transpose of the matrix (Hermitian transpose if ``conjugate`` is True); the factors are transposed.""" mats_tr = [Mi.transpose(conjugate=conjugate) for Mi in self.mats] return KroneckerStencilMatrix(self.codomain, self.domain, *mats_tr) +#============================================================================== +class ComposedKroneckerStencilMatrix(ComposedLinearOperator): + r""" + Product $M = F_1 F_2 \cdots F_n$ of Kronecker matrices with the same axis groups. + + Created by ``A @ B`` of two KroneckerStencilMatrix (or of products of them). + As for any :class:`ComposedLinearOperator`, ``multiplicands`` are the + operands $F_1, \dots, F_n$ (chains are flattened, so ``(A @ B) @ C`` has + the multiplicands ``(A, B, C)``) and ``dot`` applies them from right to left. + + In addition, the product is again a Kronecker product, + $M = C_1 \otimes \dots \otimes C_m$ with $C_k = F_{1,k} F_{2,k} \cdots F_{n,k}$, + and ``mats`` holds its exact factors $C_k$, computed at construction. Their + band is wider than the ghost regions of the domain, so they live on + process-local spaces with larger pads. They are used by ``tosparse`` and can + be used to build a :class:`KroneckerLinearSolver` for $M$. ``mats`` is a + snapshot: changing an operand in place afterwards changes ``dot`` but not + ``mats``. + + Constructing the product is collective if the operands are distributed: + each process needs all rows of the factors of the right operand that its + rows of the left factor couple to, so they are gathered. + + Parameters + ---------- + domain : StencilVectorSpace + Domain of the last operand. + + codomain : StencilVectorSpace + Codomain of the first operand. + + *args : KroneckerStencilMatrix | ComposedKroneckerStencilMatrix + The operands, with the same axis groups. + """ + + def __init__(self, + domain: StencilVectorSpace, + codomain: StencilVectorSpace, + *args: KroneckerStencilMatrix | ComposedKroneckerStencilMatrix): + + assert len(args) >= 2, 'A ComposedKroneckerStencilMatrix needs at least two operands.' + axes = args[0].axes + for a in args: + assert _is_kronecker_with_axes(a, axes), \ + 'All operands must be Kronecker matrices with the same axis groups.' + + super().__init__(domain, codomain, *args) + + # exact factors of the product, multiplied from the right + mats = list(args[-1].mats) + for a in reversed(args[:-1]): + mats = [_multiply_factors(A_k, C_k, a.domain) for A_k, C_k in zip(a.mats, mats)] + + self._mats = tuple(mats) + self._axes = axes + + @classmethod + def _from_parts(cls, domain, codomain, multiplicands, mats) -> ComposedKroneckerStencilMatrix: + """Build from known operands and factors of the product (no multiplication, not collective).""" + obj = cls.__new__(cls) + ComposedLinearOperator.__init__(obj, domain, codomain, *multiplicands) + obj._mats = tuple(mats) + obj._axes = multiplicands[0].axes + return obj + + #-------------------------------------- + # Kronecker structure + #-------------------------------------- + @property + def mats(self) -> tuple[StencilMatrix, ...]: + """Exact factors $C_k$ of the product, ordered by axis (process-local, wide band).""" + return self._mats + + @property + def axes(self) -> tuple[tuple[int, ...], ...]: + """Axes of the domain/codomain on which each factor acts, e.g. ``((0, 1), (2,))``.""" + return self._axes + + @property + def ndim(self) -> int: + """Number of axes of the domain.""" + return self.domain.ndim + + @property + def dtype(self): + return self.domain.dtype + + @property + def nbytes(self) -> int: + """Local (per-MPI-rank) memory footprint of the factors of the product and of the operands, in bytes.""" + nbytes = sum(getattr(mat, 'nbytes', 0) for mat in self._mats) + nbytes += sum(F.nbytes for F in self.multiplicands) + return int(nbytes) + + #-------------------------------------- + # Operations that keep the type + #-------------------------------------- + def copy(self) -> ComposedKroneckerStencilMatrix: + return self._from_parts(self.domain, self.codomain, + [F.copy() for F in self.multiplicands], + [m.copy() for m in self.mats]) + + def __neg__(self) -> ComposedKroneckerStencilMatrix: + return self * -1 + + def __mul__(self, a) -> ComposedKroneckerStencilMatrix: + multiplicands = [*(F.copy() for F in self.multiplicands[:-1]), self.multiplicands[-1] * a] + mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a] + return self._from_parts(self.domain, self.codomain, multiplicands, mats) + + def __rmul__(self, a) -> ComposedKroneckerStencilMatrix: + return self * a + + def __matmul__(self, B): + """ + Product ``self @ B``: a ComposedKroneckerStencilMatrix if ``B`` is a Kronecker + matrix with the same axis groups, else ``LinearOperator.__matmul__``. + """ + if _is_kronecker_with_axes(B, self.axes): + return ComposedKroneckerStencilMatrix(B.domain, self.codomain, self, B) + return super().__matmul__(B) + + def transpose(self, conjugate: bool = False) -> ComposedKroneckerStencilMatrix: + """Transpose: the operands are transposed in reverse order, and so are the factors of the product.""" + multiplicands = [F.transpose(conjugate=conjugate) for F in reversed(self.multiplicands)] + mats = [m.transpose(conjugate=conjugate) for m in self.mats] + return self._from_parts(self.codomain, self.domain, multiplicands, mats) + + #-------------------------------------- + # Conversion + #-------------------------------------- + def tosparse(self): + """Convert the local rows to a scipy sparse matrix (Kronecker product of the factors of the product).""" + return reduce(kron, (m.tosparse() for m in self.mats)) + + def toarray(self): + """Convert the local rows to a dense array.""" + return self.tosparse().toarray() + +#============================================================================== +def _is_kronecker_with_axes(B, axes) -> bool: + """Whether B is a (composed) Kronecker stencil matrix with the given axis groups.""" + return isinstance(B, (KroneckerStencilMatrix, ComposedKroneckerStencilMatrix)) and B.axes == axes + + +def _multiply_factors(A: StencilMatrix, B: StencilMatrix, U: StencilVectorSpace) -> StencilMatrix: + """ + Compute the process-local factor ``C = A @ B`` of a Kronecker product. + + The product is computed in scipy sparse format. Its band is wider than + the ones of ``A`` and ``B``, so ``C`` is stored on new spaces with + the decomposition of ``B.domain`` / ``A.codomain`` and larger pads. + + Parameters + ---------- + A, B : StencilMatrix + Process-local factors acting on the same axes. + + U : StencilVectorSpace + Codomain of the Kronecker matrix of ``B``. If it is distributed, + the rows of ``B`` owned by other processes are gathered over its + communicator. + + Returns + ------- + StencilMatrix + The product, with rows ``A.codomain.starts`` to ``A.codomain.ends``. + """ + for S in (A.domain, A.codomain, B.domain, B.codomain): + if any(m != 1 for m in S.shifts): + raise NotImplementedError('Products of factors with shifts != 1 are not supported.') + assert tuple(A.domain.npts) == tuple(B.codomain.npts) + assert tuple(A.codomain.periods) == tuple(B.domain.periods) + + # all rows of B that rows of A on this process can couple to; first sum duplicate entries + # (periodic factors with 2p + 1 > n have two diagonals in the same column), so that only rows + # replicated on several processes are removed below + B_sp = B.tosparse().tocoo() + B_sp.sum_duplicates() + if U.parallel: + parts = U.cart.comm.allgather((B_sp.row, B_sp.col, B_sp.data)) + rows = np.concatenate([p[0] for p in parts]).astype(np.int64) + cols = np.concatenate([p[1] for p in parts]).astype(np.int64) + data = np.concatenate([p[2] for p in parts]) + # processes with the same rows along these axes send them twice; keep one copy + _, idx = np.unique(rows * B_sp.shape[1] + cols, return_index=True) + B_sp = coo_matrix((data[idx], (rows[idx], cols[idx])), shape=B_sp.shape) + + C_sp = (A.tosparse().tocsr() @ B_sp.tocsr()).tocoo() + + # band of C: diagonal offset of each entry along each axis + cod, dom = A.codomain, B.domain + periods = cod.periods + rr = np.unravel_index(C_sp.row, tuple(cod.npts)) + cc = np.unravel_index(C_sp.col, tuple(dom.npts)) + kk = [] + pads = [] + for d, (pA, pB, n, P) in enumerate(zip(A.pads, B.pads, dom.npts, periods)): + k = cc[d] - rr[d] + if P: + k = (k + n//2) % n - n//2 + pads.append(min(pA + pB, n//2)) + else: + pads.append(min(pA + pB, n - 1)) + kk.append(k) + assert k.size == 0 or np.abs(k).max() <= pads[-1] + + # new spaces with the same decomposition and wider pads + def widen(S): + cart = CartDecomposition(S.cart.domain_decomposition, S.npts, + S.cart.global_starts, S.cart.global_ends, + pads=pads, shifts=list(S.shifts)) + return StencilVectorSpace(cart, dtype=C_sp.dtype) + + C = StencilMatrix(widen(dom), widen(cod)) + + index = (*(xp.asarray(r - s + p) for r,s,p in zip(rr, cod.starts, pads)), + *(xp.asarray(k + p) for k,p in zip(kk, pads))) + C._data[index] = xp.asarray(C_sp.data) + + return C + #============================================================================== class KroneckerDenseMatrix(LinearOperator): """ @@ -384,9 +744,19 @@ def set_backend(self, backend, precompiled=False): pass #============================================================================== class KroneckerLinearSolver(LinearOperator): - """ - A solver for Ax=b, where A is a Kronecker matrix from arbirary dimension d, - defined by d solvers. We also need information about the space of b. + r""" + Solver for $A x = b$, where $A = A_1 \otimes A_2 \otimes \dots \otimes A_m$ + is a Kronecker product given by one solver per factor. + + Each factor acts on a group of consecutive axes of the space; by default + every factor is 1d (one solver per axis). For other groupings, pass + ``factor_ndims``, e.g. ``factor_ndims=(2, 1)`` for a 2d x 1d product on a + 3d space. A solver of a factor with several axes receives the vectors + flattened in C order over these axes (as in ``StencilMatrix.tosparse``). + + Factors are solved in parallel (with MPI_Alltoallv) along a distributed + axis. This is only implemented for 1d factors: a factor with several + axes must not be distributed across processes along any of its axes. Parameters ---------- @@ -397,10 +767,13 @@ class KroneckerLinearSolver(LinearOperator): W : StencilVectorSpace The space x will live in; i.e. which gives us information about the distribution of the unknown vector x. - - solvers : list of LinearSolver - The components of A in each dimension. - + + solvers : sequence of LinearSolver + Solvers for the factors of A, ordered by axis. + + factor_ndims : sequence of int, optional + Number of axes of each factor. Defaults to 1 for every factor. + Attributes ---------- domain : StencilVectorSpace @@ -409,28 +782,49 @@ class KroneckerLinearSolver(LinearOperator): codomain : StencilVectorSpace The space of the unknown vector x. """ - def __init__(self, V, W, solvers): + def __init__(self, + V: StencilVectorSpace, + W: StencilVectorSpace, + solvers: Sequence[LinearSolver], + factor_ndims: Sequence[int] | None = None): assert isinstance(V, StencilVectorSpace) assert isinstance(W, StencilVectorSpace) assert hasattr( solvers, '__iter__' ) for solver in solvers: assert isinstance(solver, LinearSolver) - assert V.ndim == len(solvers) - assert W.ndim == len(solvers) + if factor_ndims is None: + factor_ndims = (1,) * len(solvers) + factor_ndims = tuple(int(n) for n in factor_ndims) + assert len(factor_ndims) == len(solvers), \ + f'Got {len(solvers)} solvers but {len(factor_ndims)} factor_ndims.' + assert all(n >= 1 for n in factor_ndims) + + assert V.ndim == sum(factor_ndims) + assert W.ndim == sum(factor_ndims) assert V.npts == W.npts + # axes of each factor + axes = [] + d = 0 + for n in factor_ndims: + axes.append(tuple(range(d, d + n))) + d += n + # general arguments self._domain = V self._codomain = W self._solvers = solvers + self._factor_ndims = factor_ndims + self._axes = tuple(axes) self._parallel = self._domain.parallel self._dtype = self._codomain._dtype if self._parallel: self._mpi_type = self._domain._mpi_type else: self._mpi_type = None - self._ndim = self._codomain.ndim + # number of factors (= number of solve passes) + self._ndim = len(solvers) # compute and setup solver arguments self._setup_solvers() @@ -462,9 +856,11 @@ def _setup_solvers(self): ends = np.array(self._domain.ends) + 1 self._slice = tuple([slice(s, e) for s,e in zip(starts, ends)]) - # local and global sizes - nglobals = self._domain.npts - nlocals = ends - starts + # local and global sizes, per factor (axes of a factor are flattened into one) + npts = self._domain.npts + nlocals_axis = ends - starts + nglobals = [int(np.prod([npts[a] for a in grp])) for grp in self._axes] + nlocals = np.array([np.prod(nlocals_axis[list(grp)]) for grp in self._axes], dtype=int) self._localsize = np.prod(nlocals) mglobals = self._localsize // nlocals self._nlocals = nlocals @@ -474,20 +870,25 @@ def _setup_solvers(self): tempsize = self._localsize self._allserial = True - for i in range(self._ndim): + for i, grp in enumerate(self._axes): # decide for each direction individually, if we should # use a serial or a parallel/distributed solver # useful e.g. if we have little data in some directions # (and thus no data distributed there) - if not self._parallel or self._domain.cart.subcomm[i].size <= 1: + distributed = self._parallel and any(self._domain.cart.subcomm[a].size > 1 for a in grp) + + if not distributed: # serial solve solver_passes[i] = KroneckerLinearSolver.KroneckerSolverSerialPass( self._solvers[i], nglobals[i], mglobals[i]) + elif len(grp) > 1: + raise NotImplementedError(f'The factor on axes {grp} is distributed across processes; ' + 'parallel solves are only implemented for 1d factors.') else: # for the parallel case, use Alltoallv solver_passes[i] = KroneckerLinearSolver.KroneckerSolverParallelPass( - self._solvers[i], self._domain._mpi_type, i, + self._solvers[i], self._domain._mpi_type, grp[0], self._domain.cart, mglobals[i], nglobals[i], nlocals[i], self._localsize) # we have a parallel solve pass now, so we are not completely local any more @@ -537,37 +938,55 @@ def _allocate_temps(self): return temp1, temp2 @property - def domain(self): + def domain(self) -> StencilVectorSpace: return self._domain @property - def codomain(self): + def codomain(self) -> StencilVectorSpace: return self._codomain @property def dtype(self): return None - def transpose(self, conjugate=False): + @property + def factor_ndims(self) -> tuple[int, ...]: + """Number of axes of each factor.""" + return self._factor_ndims + + def transpose(self, conjugate: bool = False) -> KroneckerLinearSolver: new_domain = self._codomain new_codomain = self._domain new_solvers = [solver.transpose() for solver in self._solvers] - return KroneckerLinearSolver(new_domain, new_codomain, new_solvers) + return KroneckerLinearSolver(new_domain, new_codomain, new_solvers, factor_ndims=self._factor_ndims) - def dot(self, v, out=None): + def dot(self, v: StencilVector, out: StencilVector | None = None) -> StencilVector: return self.solve(v, out=out) @property - def solvers(self): + def solvers(self) -> tuple[LinearSolver, ...]: """ - Returns an immutable view onto references to the one-dimensional solvers. + Returns an immutable view onto references to the solvers of the factors. """ return tuple(self._solvers) - def solve(self, rhs, out=None): + def solve(self, rhs: StencilVector, out: StencilVector | None = None) -> StencilVector: """ Solves Ax=b where A is a Kronecker product matrix (and represented as such), and b is a suitable vector. + + Parameters + ---------- + rhs : StencilVector + The right-hand side b, in the domain. + + out : StencilVector, optional + Vector in the codomain, in which the solution is stored. + + Returns + ------- + StencilVector + The solution x (``out`` if given). """ # type checks @@ -651,7 +1070,11 @@ def _reorder_temp_to_outslice(self, source, outslice): # outslice[:] = sourceview.transpose(self._perm) perm = tuple(int(p) for p in self._perm) - outslice[:] = sourceview.transpose(perm) + if len(self._axes) == outslice.ndim: + outslice[:] = sourceview.transpose(perm) + else: + # factors with several axes: back from one axis per factor to all axes + outslice[:] = sourceview.transpose(perm).reshape(outslice.shape) class KroneckerSolverSerialPass: """ @@ -944,21 +1367,195 @@ def solve_pass(self, workmem, tempmem): self._comm.Alltoallv(targetargs, sourceargs) #============================================================================== -def kronecker_solve(solvers, rhs, out=None): +class KroneckerSumSolver(LinearOperator): + r""" + Exact inverse of a sum of Kronecker products by fast diagonalization, + + .. math:: + + A = \sum_{d=1}^n M_1 \otimes \dots \otimes M_{d-1} \otimes S_d \otimes M_{d+1} \otimes \dots \otimes M_n + + \sigma \, M_1 \otimes \dots \otimes M_n \,, + + with symmetric 1d matrices $S_d$ and symmetric positive definite 1d matrices $M_d$ + (e.g. stiffness and mass matrices of a Laplacian on a tensor-product grid). + + In each direction, the generalized eigenproblem $S_d U_d = M_d U_d \Lambda_d$ is solved, + with $U_d^T M_d U_d = I$. Then $U^T A U = \Lambda$ with $U = U_1 \otimes \dots \otimes U_n$ + and the diagonal $\Lambda = \Lambda_1 \oplus \dots \oplus \Lambda_n + \sigma$ (entries + $\lambda_{1,i_1} + \dots + \lambda_{n,i_n} + \sigma$), hence + + .. math:: + + A^{-1} = U \, \Lambda^{-1} \, U^T \,. + + The dense matrices $U_d^T$ and $U_d$ are applied with KroneckerLinearSolver (also along + distributed axes), $\Lambda^{-1}$ locally. Entries of $\Lambda$ that vanish (relative to its + largest entry, e.g. constants for a periodic Laplacian with $\sigma = 0$) are skipped, which + gives the pseudo-inverse. + + Parameters + ---------- + V : StencilVectorSpace + Domain and codomain. + + stiffness : sequence of array | None + Dense global 1d matrices $S_d$ (shape ``(npts[d], npts[d])``); None for no term in direction d. + + mass : sequence of array + Dense global 1d matrices $M_d$ (symmetric positive definite). + + sigma : float + Coefficient of the mass term. + + rtol : float + Entries of $\Lambda$ with ``|lambda| <= rtol * max|Lambda|`` are treated as zero (pseudo-inverse). """ - Solve linear system Ax=b with A=kron( A_n, A_{n-1}, ..., A_2, A_1 ), given - $n$ separate linear solvers $L_n$ for the 1D problems $A_n x_n = b_n$: - x_n = L_n.solve( b_n ) + class _DenseApply(LinearSolver): + """'Solver' for KroneckerLinearSolver that applies a dense matrix to each right-hand side.""" + + def __init__(self, mat): + self._mat = mat + + @property + def space(self): + return xp.ndarray + + def transpose(self): + return KroneckerSumSolver._DenseApply(self._mat.T) + + def solve(self, rhs, out=None): + # rows of rhs are the vectors: out_i = mat @ rhs_i + result = rhs @ self._mat.T + if out is None: + return result + out[...] = result + return out + + def __init__(self, + V: StencilVectorSpace, + stiffness: Sequence, + mass: Sequence, + sigma: float = 0.0, + rtol: float = 1e-12): + + from scipy.linalg import eigh # deferred: scipy.linalg is slow to import + + assert isinstance(V, StencilVectorSpace) + assert len(stiffness) == len(mass) == V.ndim + + U = [] + lam = [] + for d, (S, M) in enumerate(zip(stiffness, mass)): + M = np.asarray(xp.to_numpy(xp.asarray(M)), dtype=float) + assert M.shape == (V.npts[d], V.npts[d]), f'Mass matrix of direction {d} has shape {M.shape}.' + if S is None: + # U^T M U = I with eigenvalue 0 + mu, Q = eigh(M) + assert mu.min() > 0, f'Mass matrix of direction {d} is not positive definite.' + U.append(Q / np.sqrt(mu)) + lam.append(np.zeros(V.npts[d])) + else: + S = np.asarray(xp.to_numpy(xp.asarray(S)), dtype=float) + assert S.shape == M.shape, f'Stiffness matrix of direction {d} has shape {S.shape}.' + l_d, U_d = eigh(S, M) + U.append(U_d) + lam.append(l_d) + + # eigenvalues of A on the rows owned by this process + local_lam = [l[s:e+1] for l, s, e in zip(lam, V.starts, V.ends)] + Lam = reduce(np.add.outer, local_lam) + sigma if V.ndim > 1 else local_lam[0] + sigma + Lam = np.asarray(Lam).reshape(tuple(e - s + 1 for s, e in zip(V.starts, V.ends))) + + # pseudo-inverse: skip (numerically) vanishing eigenvalues + lam_max = max(abs(l).max() for l in lam) * V.ndim + abs(sigma) + inv_Lam = np.zeros_like(Lam) + nonzero = np.abs(Lam) > rtol * lam_max + inv_Lam[nonzero] = 1.0 / Lam[nonzero] + + self._space = V + self._sigma = sigma + self._eigenvectors = tuple(U) + self._eigenvalues = tuple(lam) + self._inv_Lam = xp.asarray(inv_Lam) + self._UT = KroneckerLinearSolver(V, V, [KroneckerSumSolver._DenseApply(xp.asarray(u.T)) for u in U]) + self._U = KroneckerLinearSolver(V, V, [KroneckerSumSolver._DenseApply(xp.asarray(u)) for u in U]) + self._slice = tuple(slice(p*m, p*m + e - s + 1) for p, m, s, e in zip(V.pads, V.shifts, V.starts, V.ends)) + self._tmp = V.zeros() + + #-------------------------------------- + # Abstract interface + #-------------------------------------- + @property + def domain(self) -> StencilVectorSpace: + return self._space + + @property + def codomain(self) -> StencilVectorSpace: + return self._space + + @property + def dtype(self): + return self._space.dtype + + @property + def eigenvalues(self) -> tuple: + r"""1d generalized eigenvalues $\Lambda_d$ of $(S_d, M_d)$ in each direction.""" + return self._eigenvalues + + @property + def eigenvectors(self) -> tuple: + """1d generalized eigenvectors $U_d$ (columns, $U_d^T M_d U_d = I$) in each direction.""" + return self._eigenvectors + + def transpose(self, conjugate: bool = False) -> KroneckerSumSolver: + """A is symmetric, so is its inverse.""" + return self + + def dot(self, v: StencilVector, out: StencilVector | None = None) -> StencilVector: + r"""Apply $A^{-1} = U \Lambda^{-1} U^T$ to ``v``.""" + assert isinstance(v, StencilVector) + assert v.space is self._space + if out is not None: + assert isinstance(out, StencilVector) + assert out.space is self._space + else: + out = StencilVector(self._space) + + tmp = self._UT.solve(v, out=self._tmp) + tmp._data[self._slice] *= self._inv_Lam + return self._U.solve(tmp, out=out) + + def solve(self, rhs: StencilVector, out: StencilVector | None = None) -> StencilVector: + return self.dot(rhs, out=out) + +#============================================================================== +def kronecker_solve(solvers: Sequence[LinearSolver], + rhs: StencilVector, + out: StencilVector | None = None, + factor_ndims: Sequence[int] | None = None) -> StencilVector: + """ + Solve the linear system Ax=b with A = kron(A_1, A_2, ..., A_m), given + one linear solver L_k per factor ($L_k$ solves $A_k x_k = b_k$). Parameters ---------- - solvers : list( LinearSolver ) - List of linear solvers along each direction: [L_1, L_2, ..., L_n]. + solvers : sequence of LinearSolver + Solvers for the factors, ordered by axis: [L_1, L_2, ..., L_m]. rhs : StencilVector Right hand side vector of linear system Ax=b. + out : StencilVector, optional + Vector in the space of rhs, in which the solution is stored. + + factor_ndims : sequence of int, optional + Number of axes of each factor (see KroneckerLinearSolver). Defaults to 1 for every factor. + + Returns + ------- + StencilVector + The solution x (``out`` if given). """ # all these feasability checks are again performed in the KroneckerLinearSolver class assert hasattr(solvers, '__iter__') @@ -966,7 +1563,10 @@ def kronecker_solve(solvers, rhs, out=None): assert isinstance(solver, LinearSolver) assert isinstance(rhs, StencilVector) - assert rhs.space.ndim == len(solvers) + if factor_ndims is None: + assert rhs.space.ndim == len(solvers) + else: + assert rhs.space.ndim == sum(factor_ndims) if out is not None: assert isinstance(out, StencilVector) @@ -974,5 +1574,5 @@ def kronecker_solve(solvers, rhs, out=None): else: out = StencilVector(rhs.space) - kronsolver = KroneckerLinearSolver(rhs.space, rhs.space, solvers) + kronsolver = KroneckerLinearSolver(rhs.space, rhs.space, solvers, factor_ndims=factor_ndims) return kronsolver.solve(rhs, out=out) diff --git a/feectools/linalg/tests/test_kron_stencil_matrix.py b/feectools/linalg/tests/test_kron_stencil_matrix.py index d160394fb..81b29b4ce 100644 --- a/feectools/linalg/tests/test_kron_stencil_matrix.py +++ b/feectools/linalg/tests/test_kron_stencil_matrix.py @@ -113,3 +113,414 @@ def test_KroneckerStencilMatrix(dtype, npts, pads, periodic): # Test dot product expected = M_sp.dot(xp.to_numpy(w.toarray())) assert xp.array_equal(xp.asarray(expected), M.dot(w).toarray()) + +#=============================================================================== +# Kronecker matrices with factors on groups of axes, products and solvers +#=============================================================================== +import numpy as np +from scipy.sparse import csr_matrix + +from maybempi import MPI +from feectools.linalg.basic import ComposedLinearOperator +from feectools.linalg.direct_solvers import SparseSolver +from feectools.linalg.kron import KroneckerLinearSolver, kronecker_solve +from feectools.linalg.kron import ComposedKroneckerStencilMatrix + + +def make_space(comm, npts, pads, periods, mpi_dims_mask=None): + """Distributed StencilVectorSpace (serial if comm is None).""" + D = DomainDecomposition([n-1 for n in npts], periods=periods, comm=comm, + mpi_dims_mask=mpi_dims_mask) + global_starts, global_ends = compute_global_starts_ends(D, npts) + cart = CartDecomposition(D, npts, global_starts, global_ends, pads=pads, shifts=[1]*len(npts)) + return StencilVectorSpace(cart) + + +def make_factor(W, grp, mpads, seed, full_rows=False, rows=None): + """ + Process-local factor on the axes grp of W (rows owned by this process, all rows if + full_rows, or the given (starts, ends)), with band mpads and random entries; also + returns the global dense matrix. + """ + npts = [W.npts[a] for a in grp] + periods = [W.periods[a] for a in grp] + spads = list(mpads) # factor space pads may be smaller than the ones of W + starts = [W.starts[a] for a in grp] + ends = [W.ends[a] for a in grp] + if full_rows: + starts, ends = [0] * len(grp), [n - 1 for n in npts] + if rows is not None: + starts, ends = rows + + D = DomainDecomposition([n-1 for n in npts], periods=periods) + cart = CartDecomposition(D, npts, [np.array([s]) for s in starts], [np.array([e]) for e in ends], + pads=spads, shifts=[1]*len(grp)) + V = StencilVectorSpace(cart) + M = StencilMatrix(V, V, pads=tuple(mpads)) + + # same random values on all processes; diagonally dominant + rng = np.random.default_rng(seed) + ndiags = tuple(2*p+1 for p in mpads) + vals = rng.random(tuple(npts) + ndiags) + vals[(Ellipsis, *mpads)] += 2 * np.prod(ndiags) + + N = int(np.prod(npts)) + dense = np.zeros((N, N)) + for i in np.ndindex(*npts): + for kk in np.ndindex(*ndiags): + col = [] + for i_d, k_d, p, n, P in zip(i, kk, mpads, npts, periods): + c = i_d + k_d - p + if P: + c %= n + elif not 0 <= c < n: + break + col.append(c) + else: + v = vals[(*i, *kk)] + dense[np.ravel_multi_index(i, npts), np.ravel_multi_index(col, npts)] += v + if all(s <= i_d <= e for i_d, s, e in zip(i, starts, ends)): + M._data[(*(i_d - s + sp for i_d, s, sp in zip(i, starts, spads)), *kk)] = v + return M, dense + + +def local_slices(W): + return tuple(slice(s, e+1) for s, e in zip(W.starts, W.ends)) + + +def local_rows(W): + """Flattened (C order) global indices of the rows owned by this process.""" + grids = np.meshgrid(*[np.arange(s, e+1) for s, e in zip(W.starts, W.ends)], indexing='ij') + return np.ravel_multi_index(tuple(g.ravel() for g in grids), tuple(W.npts)) + + +def random_vector(W, seed): + wglob = np.random.default_rng(seed).random(tuple(W.npts)) + w = StencilVector(W) + w[local_slices(W)] = xp.asarray(wglob[local_slices(W)]) + return w, wglob.ravel() + + +def assert_local_equal(W, y, expected): + y_loc = xp.to_numpy(y[local_slices(W)]) + assert np.allclose(y_loc, expected.reshape(tuple(W.npts))[local_slices(W)], rtol=1e-12, atol=1e-12) + + +def make_kron(W, factor_ndims, mpads, seed, full_rows=False): + axes, d = [], 0 + for n in factor_ndims: + axes.append(tuple(range(d, d + n))) + d += n + mats, dense = [], [] + for k, grp in enumerate(axes): + M, Md = make_factor(W, grp, [mpads[a] for a in grp], seed + k, full_rows=full_rows) + mats.append(M) + dense.append(Md) + from functools import reduce as _reduce + return KroneckerStencilMatrix(W, W, *mats), _reduce(np.kron, dense), mats, dense + + +NPTS = [6, 7, 8] +PADS = [2, 2, 3] +PERIODS = [True, False, True] +GROUPS = [(1, 1, 1), (2, 1), (1, 2), (3,)] + + +def check_grouped_kron(comm, factor_ndims): + W = make_space(comm, NPTS, PADS, PERIODS) + M, Md, _, _ = make_kron(W, factor_ndims, [1, 2, 2], seed=10) + w, wglob = random_vector(W, seed=1) + + assert M.ndim == 3 + assert len(M.axes) == len(factor_ndims) + + # dot + assert_local_equal(W, M.dot(w), Md @ wglob) + + # tosparse (local rows) + rows = local_rows(W) + assert np.allclose(M.tosparse().tocsr()[rows].toarray(), Md[rows]) + + # __getitem__: row i, diagonal offset k + i = tuple(W.starts) + k = (1, -1, 2) + col = [(ii + kk) % n for ii, kk, n in zip(i, k, NPTS)] + assert np.isclose(float(M[(*i, *k)]), + Md[np.ravel_multi_index(i, NPTS), np.ravel_multi_index(col, NPTS)]) + + # scaling + assert_local_equal(W, (M * 3.).dot(w), 3. * Md @ wglob) + assert_local_equal(W, (-M).dot(w), -Md @ wglob) + + # tostencil, transpose (process-local factors are only complete in serial) + if comm is None: + assert np.allclose(M.tostencil().toarray(), Md) + assert np.allclose(M.T.toarray(), Md.T) + + +def check_matmul(comm, factor_ndims): + W = make_space(comm, NPTS, PADS, PERIODS) + A, Ad, _, _ = make_kron(W, factor_ndims, [1, 2, 2], seed=20) + B, Bd, _, _ = make_kron(W, factor_ndims, [2, 1, 3], seed=30) + w, wglob = random_vector(W, seed=2) + rows = local_rows(W) + + C = A @ B + assert type(C) is ComposedKroneckerStencilMatrix + assert isinstance(C, ComposedLinearOperator) + assert C.domain is B.domain and C.codomain is A.codomain + assert C.axes == A.axes and C.ndim == 3 + assert C.multiplicands == (A, B) + for Ck, Ak, Bk in zip(C.mats, A.mats, B.mats): + assert Ck.pads == tuple(min(pa + pb, n//2 if P else n-1) for pa, pb, n, P + in zip(Ak.pads, Bk.pads, Ak.domain.npts, Ak.domain.periods)) + + # exact factors and dot through the operands + assert np.allclose(C.tosparse().tocsr()[rows].toarray(), (Ad @ Bd)[rows]) + assert_local_equal(W, C.dot(w), Ad @ Bd @ wglob) + out = StencilVector(W) + assert C.dot(w, out=out) is out + assert_local_equal(W, out, Ad @ Bd @ wglob) + + # chains are flattened, on both sides + for D in (C @ A, A @ (B @ A), (A @ B) @ A): + assert type(D) is ComposedKroneckerStencilMatrix + assert D.multiplicands == (A, B, A) + assert np.allclose(D.tosparse().tocsr()[rows].toarray(), (Ad @ Bd @ Ad)[rows]) + assert_local_equal(W, D.dot(w), Ad @ Bd @ Ad @ wglob) + + # scaling, copy (keep the type, leave C unchanged) + for C2, f in ((C * 2., 2.), (2. * C, 2.), (-C, -1.), (C.copy(), 1.)): + assert type(C2) is ComposedKroneckerStencilMatrix + assert_local_equal(W, C2.dot(w), f * Ad @ Bd @ wglob) + assert np.allclose(C2.tosparse().tocsr()[rows].toarray(), f * (Ad @ Bd)[rows]) + C2 = C.copy() + C2 *= 0.5 + assert_local_equal(W, C2.dot(w), 0.5 * Ad @ Bd @ wglob) + assert_local_equal(W, C.dot(w), Ad @ Bd @ wglob) + assert np.allclose(C.tosparse().tocsr()[rows].toarray(), (Ad @ Bd)[rows]) + + # the factors of the product do not fit into the ghost regions of the domain + with pytest.raises(ValueError): + KroneckerStencilMatrix(W, W, *C.mats) + + # transpose (process-local factors are only complete in serial) + if comm is None: + assert type(C.T) is ComposedKroneckerStencilMatrix + assert np.allclose(C.T.toarray(), (Ad @ Bd).T) + assert_local_equal(W, C.T.dot(w), (Ad @ Bd).T @ wglob) + + # other operands fall back to LinearOperator.__matmul__ + E, Ed, _, _ = make_kron(W, (1, 1, 1) if factor_ndims != (1, 1, 1) else (2, 1), [1, 1, 1], seed=40) + for AE in (A @ E, C @ E): + assert type(AE) is ComposedLinearOperator + assert_local_equal(W, (A @ E).dot(w), Ad @ Ed @ wglob) + assert_local_equal(W, A @ w, Ad @ wglob) + + +def check_solver(comm, factor_ndims, mpi_dims_mask=None): + W = make_space(comm, NPTS, PADS, PERIODS, mpi_dims_mask=mpi_dims_mask) + M, Md, _, dense = make_kron(W, factor_ndims, [1, 2, 2], seed=50) + solvers = [SparseSolver(csr_matrix(D)) for D in dense] + b, bglob = random_vector(W, seed=3) + + S = KroneckerLinearSolver(W, W, solvers, factor_ndims=factor_ndims) + assert S.factor_ndims == factor_ndims + assert_local_equal(W, S.solve(b), np.linalg.solve(Md, bglob)) + assert_local_equal(W, kronecker_solve(solvers, b, factor_ndims=factor_ndims), np.linalg.solve(Md, bglob)) + assert_local_equal(W, S.T.solve(b), np.linalg.solve(Md.T, bglob)) + + # solver for a product of Kronecker matrices (the factors hold all rows only in serial) + C = M @ M + assert type(C) is ComposedKroneckerStencilMatrix + if comm is None: + S2 = KroneckerLinearSolver(W, W, [SparseSolver(Ck.tosparse().tocsr()) for Ck in C.mats], + factor_ndims=factor_ndims) + assert_local_equal(W, S2.solve(b), np.linalg.solve(Md @ Md, bglob)) + + +#=============================================================================== +@pytest.mark.parametrize('factor_ndims', GROUPS) +def test_grouped_kron_ser(factor_ndims): + check_grouped_kron(None, factor_ndims) + +@pytest.mark.parametrize('factor_ndims', GROUPS) +def test_kron_matmul_ser(factor_ndims): + check_matmul(None, factor_ndims) + +@pytest.mark.parametrize('factor_ndims', GROUPS) +def test_kron_solver_groups_ser(factor_ndims): + check_solver(None, factor_ndims) + +def test_kron_constructor_checks(): + W = make_space(None, NPTS, PADS, PERIODS) + M, _, mats, _ = make_kron(W, (2, 1), [1, 1, 1], seed=0) + with pytest.raises(AssertionError): + KroneckerStencilMatrix(W, W, mats[0]) # too few axes + with pytest.raises(AssertionError): + KroneckerStencilMatrix(W, W, mats[1], mats[0]) # npts do not match + with pytest.raises(AssertionError): + KroneckerStencilMatrix(W, W, *mats, mats[1]) # too many axes + with pytest.raises(AssertionError): + KroneckerLinearSolver(W, W, [SparseSolver(csr_matrix(np.eye(2)))] * 2, factor_ndims=(1, 1)) + +@pytest.mark.mpi +@pytest.mark.parametrize('factor_ndims', GROUPS) +def test_grouped_kron_par(factor_ndims): + check_grouped_kron(MPI.COMM_WORLD, factor_ndims) + +@pytest.mark.mpi +@pytest.mark.parametrize('factor_ndims', GROUPS) +def test_kron_matmul_par(factor_ndims): + check_matmul(MPI.COMM_WORLD, factor_ndims) + +@pytest.mark.mpi +@pytest.mark.parametrize('factor_ndims', [(1, 1, 1), (2, 1)]) +def test_kron_solver_groups_par(factor_ndims): + # grouped axes are not distributed: only the last axis is + check_solver(MPI.COMM_WORLD, factor_ndims, mpi_dims_mask=[False, False, True]) + +@pytest.mark.mpi +def test_kron_solver_distributed_group_par(): + comm = MPI.COMM_WORLD + if comm.Get_size() == 1: + pytest.skip('needs more than one process') + W = make_space(comm, NPTS, PADS, PERIODS, mpi_dims_mask=[True, False, False]) + _, _, _, dense = make_kron(W, (2, 1), [1, 1, 1], seed=0) + with pytest.raises(NotImplementedError): + KroneckerLinearSolver(W, W, [SparseSolver(csr_matrix(D)) for D in dense], factor_ndims=(2, 1)) + + +def test_multiplicants_deprecated(): + W = make_space(None, NPTS, PADS, PERIODS) + A = make_kron(W, (1, 1, 1), [1, 1, 1], seed=0)[0] + C = ComposedLinearOperator(W, W, A, A) + with pytest.warns(DeprecationWarning): + assert C.multiplicants == C.multiplicands == (A, A) + + +#=============================================================================== +# KroneckerSumSolver (fast diagonalization) +#=============================================================================== +from feectools.linalg.kron import KroneckerSumSolver + + +def laplace_1d(n, periodic, seed): + """1d stiffness (random positive weights on a difference operator) and mass matrix.""" + rng = np.random.default_rng(seed) + nd = n if periodic else n - 1 + D = np.zeros((nd, n)) + for r in range(nd): + D[r, r] = -1. + D[r, (r + 1) % n] = 1. + S = D.T @ np.diag(1. + rng.random(nd)) @ D + B = rng.random((n, n)) * (np.abs(np.subtract.outer(np.arange(n), np.arange(n))) <= 2) + M = B @ B.T + n * np.eye(n) + return S, M + + +def check_sum_solver(comm, sigma, with_none): + npts, periods = [6, 7, 8], [True, False, True] + W = make_space(comm, npts, [2, 2, 3], periods) + S, M = zip(*(laplace_1d(n, P, seed=d) for d, (n, P) in enumerate(zip(npts, periods)))) + S = list(S) + if with_none: + S[1] = None + + def kron_term(d): + mats = [M[e] if e != d else S[d] for e in range(3)] + return reduce(np.kron, mats) + + A = sum(kron_term(d) for d in range(3) if S[d] is not None) + sigma * reduce(np.kron, M) + solver = KroneckerSumSolver(W, S, M, sigma=sigma) + + if sigma > 0: + b, bglob = random_vector(W, seed=4) + assert_local_equal(W, solver.dot(b), np.linalg.solve(A, bglob)) + else: + # singular (constants of the periodic directions are in the kernel): pseudo-inverse, + # exact for right-hand sides in the range of A + bglob = A @ np.random.default_rng(5).random(A.shape[0]) + b = StencilVector_from(W, bglob) + x = solver.dot(b) + assert_local_equal(W, StencilVector_from(W, A @ local_to_global(W, x)), bglob) + assert solver.transpose() is solver + + +def local_to_global(W, x): + """Gather a distributed StencilVector into a global array (all processes).""" + glob = np.zeros(tuple(W.npts)) + glob[local_slices(W)] = xp.to_numpy(x[local_slices(W)]) + comm = W.cart.comm if W.parallel else None + if comm is not None: + glob = comm.allreduce(glob) + return glob.ravel() + + +def StencilVector_from(W, glob): + v = StencilVector(W) + v[local_slices(W)] = xp.asarray(glob.reshape(tuple(W.npts))[local_slices(W)]) + return v + + +@pytest.mark.parametrize('sigma', [0.7, 0.]) +@pytest.mark.parametrize('with_none', [False, True]) +def test_kron_sum_solver_ser(sigma, with_none): + check_sum_solver(None, sigma, with_none) + +@pytest.mark.mpi +@pytest.mark.parametrize('sigma', [0.7, 0.]) +@pytest.mark.parametrize('with_none', [False, True]) +def test_kron_sum_solver_par(sigma, with_none): + check_sum_solver(MPI.COMM_WORLD, sigma, with_none) + + +#=============================================================================== +# Row layout of the factors, duplicate entries of periodic factors +#=============================================================================== +def check_full_row_factors(comm, factor_ndims): + """Factors owning all rows (not only the ones of W on this process).""" + W = make_space(comm, NPTS, PADS, PERIODS) + M, Md, _, _ = make_kron(W, factor_ndims, [1, 2, 2], seed=60, full_rows=True) + w, wglob = random_vector(W, seed=6) + assert_local_equal(W, M.dot(w), Md @ wglob) + rows = local_rows(W) + assert np.allclose(M.tostencil().tosparse().tocsr()[rows].toarray(), Md[rows]) + + +@pytest.mark.parametrize('factor_ndims', [(1, 1, 1), (2, 1)]) +def test_kron_full_row_factors_ser(factor_ndims): + check_full_row_factors(None, factor_ndims) + +@pytest.mark.mpi +@pytest.mark.parametrize('factor_ndims', [(1, 1, 1), (2, 1)]) +def test_kron_full_row_factors_par(factor_ndims): + check_full_row_factors(MPI.COMM_WORLD, factor_ndims) + + +def test_kron_factor_rows_must_contain_local_rows(): + W = make_space(None, NPTS, PADS, PERIODS) + good = [make_factor(W, (a,), [1], seed=a)[0] for a in range(3)] + # factor on axis 0 owning only the rows 1, ..., n-1 (row 0 of W is missing) + bad = make_factor(W, (0,), [1], seed=0, rows=([1], [NPTS[0] - 1]))[0] + with pytest.raises(ValueError): + KroneckerStencilMatrix(W, W, bad, *good[1:]) + + +def check_matmul_periodic_duplicates(comm): + """Periodic direction with 2p + 1 > n: two diagonals hit the same column.""" + npts, pads, periods = [4, 7, 8], [2, 2, 3], [True, False, True] + W = make_space(comm, npts, pads, periods, mpi_dims_mask=[False, True, True]) + A, Ad, _, _ = make_kron(W, (1, 1, 1), [2, 1, 2], seed=70) + B, Bd, _, _ = make_kron(W, (1, 1, 1), [2, 2, 1], seed=80) + rows = local_rows(W) + C = A @ B + assert np.allclose(C.tosparse().tocsr()[rows].toarray(), (Ad @ Bd)[rows]) + + +def test_kron_matmul_periodic_duplicates_ser(): + check_matmul_periodic_duplicates(None) + +@pytest.mark.mpi +def test_kron_matmul_periodic_duplicates_par(): + check_matmul_periodic_duplicates(MPI.COMM_WORLD) diff --git a/pyproject.toml b/pyproject.toml index e80c4f4cb..d2059be47 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "feectools" -version = "0.6.0" +version = "0.7.0" description = "Slimmed-down fork of Psydac (https://github.com/pyccel/psydac) with less functionality and fewer dependencies." readme = "README.md" requires-python = ">= 3.10"