From 6ee8f082c5e1f9cdfd25ed726f6f99f4b5ecd97b Mon Sep 17 00:00:00 2001 From: Stefan Possanner Date: Thu, 8 Oct 2026 08:52:05 +0200 Subject: [PATCH 1/6] KroneckerStencilMatrix: matmul, factors on groups of axes, docs - KroneckerStencilMatrix factors may act on several consecutive axes (e.g. 2d x 1d on a 3d space); new `axes` property, stricter checks (StencilMatrix factors, codomain npts, number of axes). - `A @ B` of two Kronecker matrices with the same axis groups returns a KroneckerStencilMatrix with factors C_k = A_k @ B_k (computed in sparse format on process-local spaces with wider pads, rows of B gathered across processes). Domain/codomain stay the original spaces; copies of the operands are kept in `factors` and used by `dot`. - dot/tostencil use the factor's own pads; tostencil raises ValueError if the band does not fit into the domain pads. - KroneckerLinearSolver and kronecker_solve take `factor_ndims` for solvers of factors with several axes (serial along grouped axes). - Fix __imul__ on the tuple of factors; docstrings and type annotations. - Bump version to 0.6.0. Co-Authored-By: Claude Opus 5.5 --- feectools/linalg/kron.py | 542 +++++++++++++++--- .../linalg/tests/test_kron_stencil_matrix.py | 271 +++++++++ pyproject.toml | 2 +- 3 files changed, 722 insertions(+), 93 deletions(-) diff --git a/feectools/linalg/kron.py b/feectools/linalg/kron.py index eb98c99a3..42e5618dc 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,6 +10,7 @@ from scipy.sparse import kron from scipy.sparse import coo_matrix +from feectools.ddm.cart import CartDecomposition from feectools.linalg.basic import LinearOperator, LinearSolver from feectools.linalg.stencil import StencilVectorSpace, StencilVector, StencilMatrix @@ -17,8 +21,26 @@ #============================================================================== 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: along its axes, the factor $A_k$ owns the + same rows (``starts``/``ends``) as the codomain ``W`` on this process, + but it lives on its own spaces, without a communicator. + + A product ``M = A @ B`` of two Kronecker matrices with the same axis + groups is again a KroneckerStencilMatrix with factors $C_k = A_k B_k$ + (see ``__matmul__``). The band of $C_k$ is wider than the ghost regions of + the domain, so the operands are stored in ``factors`` and ``dot`` applies + them one after another. Parameters ---------- @@ -28,93 +50,190 @@ 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. Unless + ``factors`` is given, the pads of ``A_k`` must not exceed the pads + of ``V``. + factors : sequence of KroneckerStencilMatrix, optional + Operands $F_1, \dots, F_n$ with $M = F_1 F_2 \cdots F_n$, as created + by ``__matmul__``. If given, ``dot`` applies them from right to left. """ - def __init__(self, V, W, *args): + def __init__(self, + V: StencilVectorSpace, + W: StencilVectorSpace, + *args: StencilMatrix, + factors: Sequence[KroneckerStencilMatrix] | None = None): 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] - - self._domain = V - self._codomain = W - self._mats = args - self._ndim = len(args) + 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}.' + + if factors is None: + # 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 ' + f'the domain pads {V.pads}; pass the operands ' + 'as factors (as done by A @ B).') + tmp_vectors = () + else: + factors = tuple(factors) + assert len(factors) > 0 + for F in factors: + assert isinstance(F, KroneckerStencilMatrix) + assert F.axes == tuple(axes) + assert factors[0].codomain == W + assert factors[-1].domain == V + for F, G in zip(factors[:-1], factors[1:]): + assert F.domain == G.codomain + # intermediate results for dot + tmp_vectors = tuple(F.codomain.zeros() for F in factors[1:]) + + self._domain = V + self._codomain = W + self._mats = tuple(args) + self._axes = tuple(axes) + self._factors = factors + self._tmp_vectors = tmp_vectors #-------------------------------------- # 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.""" - return int(sum(getattr(mat, 'nbytes', 0) for mat in self._mats)) + def axes(self) -> tuple[tuple[int, ...], ...]: + """Axes of the domain/codomain on which each factor acts, e.g. ``((0, 1), (2,))``.""" + return self._axes # ... - def dot(self, x, out=None): + @property + def factors(self) -> tuple[KroneckerStencilMatrix, ...] | None: + """Operands $F_1, \\dots, F_n$ with ``self == F_1 @ ... @ F_n``, or None if ``self`` is not a product.""" + return self._factors - dot = xp.dot + # ... + @property + def nbytes(self) -> int: + """Local (per-MPI-rank) memory footprint of the factor matrices (and of the stored operands), in bytes.""" + nbytes = sum(getattr(mat, 'nbytes', 0) for mat in self._mats) + if self._factors is not None: + nbytes += sum(F.nbytes for F in self._factors) + return int(nbytes) + + # ... + def dot(self, x: StencilVector, out: StencilVector | None = None) -> StencilVector: + """ + Matrix-vector product ``M @ x``. + + 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) + # product of Kronecker matrices: apply the operands one after another + if self._factors is not None: + y = x + for F, tmp in zip(reversed(self._factors[1:]), reversed(self._tmp_vectors)): + y = F.dot(y, out=tmp) + return self._factors[0].dot(y, out=out) + + # 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 = tuple(p*m for A in mats for p,m in zip(A.codomain.pads, A.codomain.shifts)) + 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. @@ -127,47 +246,192 @@ def dot(self, x, out=None): return out # ... - def copy(self): - mats = [m.copy() for m in self.mats] - return KroneckerStencilMatrix(self.domain, self.codomain, *mats) + def copy(self) -> KroneckerStencilMatrix: + mats = [m.copy() for m in self.mats] + factors = None if self._factors is None else [F.copy() for F in self._factors] + return KroneckerStencilMatrix(self.domain, self.codomain, *mats, factors=factors) # ... - def __neg__(self): - mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])] - return KroneckerStencilMatrix(self.domain, self.codomain, *mats) + def __neg__(self) -> KroneckerStencilMatrix: + mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])] + factors = None if self._factors is None else \ + [-self._factors[0], *(F.copy() for F in self._factors[1:])] + return KroneckerStencilMatrix(self.domain, self.codomain, *mats, factors=factors) # ... - def __mul__(self, a): - mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a] - return KroneckerStencilMatrix(self.domain, self.codomain, *mats) + def __mul__(self, a) -> KroneckerStencilMatrix: + mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a] + factors = None if self._factors is None else \ + [*(F.copy() for F in self._factors[:-1]), self._factors[-1] * a] + return KroneckerStencilMatrix(self.domain, self.codomain, *mats, factors=factors) # ... - def __imul__(self, a): - self.mats[-1] *= a + def __imul__(self, a) -> KroneckerStencilMatrix: + last = self._mats[-1] + last *= a + if self._factors is not None: + last = self._factors[-1] + last *= a return self + # ... + def __matmul__(self, B): + """ + Product ``self @ B``. + + If ``B`` is a KroneckerStencilMatrix with the same axis groups, the + result is a KroneckerStencilMatrix with factors + ``C_k = self.mats[k] @ B.mats[k]``, domain ``B.domain`` and codomain + ``self.codomain``. Copies of the operands are stored in + ``factors`` (products are flattened, so ``(A @ B) @ C`` has the + factors ``(A, B, C)``) and used by ``dot``. + + This is a collective operation if ``B.codomain`` is distributed: each + process needs all rows of ``B.mats[k]`` that its rows of + ``self.mats[k]`` couple to, so they are gathered. + + 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 + ------- + KroneckerStencilMatrix | LinearOperator | Vector + The product. + """ + if not isinstance(B, KroneckerStencilMatrix) or B.axes != self.axes: + return super().__matmul__(B) + + assert self.domain == B.codomain, \ + 'The domain of the left operand must be the codomain of the right operand.' + + mats = [self._multiply_factors(A_k, B_k, B.codomain) + for A_k, B_k in zip(self.mats, B.mats)] + + factors = (*(self._factors or (self,)), *(B._factors or (B,))) + factors = [F.copy() for F in factors] + + return KroneckerStencilMatrix(B.domain, self.codomain, *mats, factors=factors) + + @staticmethod + 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 + B_sp = B.tosparse().tocoo() + 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 + #-------------------------------------- # 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 tostencil(self) -> StencilMatrix: + """ + Convert to a StencilMatrix on the domain and codomain. + + Raises + ------ + ValueError + If the band of a factor exceeds the domain pads, e.g. for a product + created by ``A @ B``; use ``tosparse`` instead. + """ 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 + if any(p > xp_ for p,xp_ in zip(pads, xpads)): + raise ValueError(f'The pads {tuple(pads)} of the factors exceed the domain pads ' + f'{xpads}, so the matrix does not fit into a StencilMatrix on ' + 'its domain; use tosparse instead.') + # Number of rows in matrix (along each dimension) nrows = [ed-s+1 for s,ed in zip(ssd, eed)] nrows_extra = [0 if ec<=ed else ec-ed for ec,ed in zip(eec,eed)] @@ -175,13 +439,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 = [p*m for A in mats for p,m in zip(A.codomain.pads, A.codomain.shifts)] + 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)] @@ -190,12 +457,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): @@ -209,6 +478,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)] @@ -216,19 +486,30 @@ 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; for a product the order of the operands + is reversed. + """ mats_tr = [Mi.transpose(conjugate=conjugate) for Mi in self.mats] - return KroneckerStencilMatrix(self.codomain, self.domain, *mats_tr) + factors = None if self._factors is None else \ + [F.transpose(conjugate=conjugate) for F in reversed(self._factors)] + return KroneckerStencilMatrix(self.codomain, self.domain, *mats_tr, factors=factors) #============================================================================== class KroneckerDenseMatrix(LinearOperator): @@ -383,9 +664,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 ---------- @@ -396,10 +687,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 @@ -408,28 +702,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() @@ -454,9 +769,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 @@ -466,20 +783,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 @@ -529,37 +851,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 @@ -643,7 +983,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: """ @@ -899,21 +1243,32 @@ def solve_pass(self, workmem, tempmem): self._comm.Alltoallv(targetargs, sourceargs) #============================================================================== -def kronecker_solve(solvers, rhs, out=None): +def kronecker_solve(solvers: Sequence[LinearSolver], + rhs: StencilVector, + out: StencilVector | None = None, + factor_ndims: Sequence[int] | None = None) -> StencilVector: """ - 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 ) + 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__') @@ -921,7 +1276,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) @@ -929,5 +1287,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..1ae484fa3 100644 --- a/feectools/linalg/tests/test_kron_stencil_matrix.py +++ b/feectools/linalg/tests/test_kron_stencil_matrix.py @@ -113,3 +113,274 @@ 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 + + +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): + """ + Process-local factor on the axes grp of W (rows owned by this process), + 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] + + 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): + 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) + 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) + assert M.factors is None + + # 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 isinstance(C, KroneckerStencilMatrix) + assert C.domain is B.domain and C.codomain is A.codomain + assert C.axes == A.axes + assert len(C.factors) == 2 + for F, G in zip(C.factors, (A, B)): + assert F is not G + assert all(np.array_equal(xp.to_numpy(f._data), xp.to_numpy(g._data)) for f, g in zip(F.mats, G.mats)) + 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 + D = C @ A + assert len(D.factors) == 3 + 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 + assert_local_equal(W, (C * 2.).dot(w), 2. * Ad @ Bd @ wglob) + assert_local_equal(W, (-C).dot(w), -Ad @ Bd @ wglob) + assert_local_equal(W, C.copy().dot(w), Ad @ Bd @ wglob) + C2 = C.copy() + C2 *= 0.5 + assert_local_equal(W, C2.dot(w), 0.5 * Ad @ Bd @ wglob) + assert np.allclose(C2.tosparse().tocsr()[rows].toarray(), 0.5 * (Ad @ Bd)[rows]) + assert_local_equal(W, C.dot(w), Ad @ Bd @ wglob) + + # the band does not fit into the ghost regions + with pytest.raises(ValueError): + C.tostencil() + with pytest.raises(ValueError): + KroneckerStencilMatrix(W, W, *C.mats) + + # transpose (process-local factors are only complete in serial) + if comm is None: + 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) + AE = A @ E + assert isinstance(AE, ComposedLinearOperator) + assert_local_equal(W, AE.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 + 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)) diff --git a/pyproject.toml b/pyproject.toml index f378db48e..e80c4f4cb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "feectools" -version = "0.5.0" +version = "0.6.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" From fa6fe90403cf34935709616712b6a199b2554036 Mon Sep 17 00:00:00 2001 From: Stefan Possanner Date: Thu, 8 Oct 2026 11:02:35 +0200 Subject: [PATCH 2/6] ComposedKroneckerStencilMatrix for products; rename multiplicands - A @ B of Kronecker matrices with the same axis groups now returns a ComposedKroneckerStencilMatrix (subclass of ComposedLinearOperator): `multiplicands` are the operands (dot goes through them), `mats` the exact factors C_k of the product (process-local, wide band; used by tosparse and for KroneckerLinearSolver). Scaling, copy, transpose and further products keep the type. - KroneckerStencilMatrix loses the `factors` argument: its factor pads always fit into the domain pads, so dot and tostencil always work. - ComposedLinearOperator.multiplicants -> multiplicands (correct spelling); `multiplicants` remains as a deprecated alias. Co-Authored-By: Claude Opus 5.5 --- feectools/api/fem_bilinear_form.py | 4 +- feectools/api/fem_common.py | 2 +- feectools/linalg/basic.py | 44 +- feectools/linalg/kron.py | 462 ++++++++++-------- .../linalg/tests/test_kron_stencil_matrix.py | 56 ++- 5 files changed, 320 insertions(+), 248 deletions(-) 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/linalg/basic.py b/feectools/linalg/basic.py index 1fdcc4431..4ef0dc39b 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 @@ -1055,27 +1056,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 @@ -1096,34 +1097,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) @@ -1135,11 +1143,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) @@ -1148,11 +1156,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 42e5618dc..667f2565e 100644 --- a/feectools/linalg/kron.py +++ b/feectools/linalg/kron.py @@ -11,10 +11,11 @@ from scipy.sparse import coo_matrix from feectools.ddm.cart import CartDecomposition -from feectools.linalg.basic import LinearOperator, LinearSolver +from feectools.linalg.basic import ComposedLinearOperator, LinearOperator, LinearSolver from feectools.linalg.stencil import StencilVectorSpace, StencilVector, StencilMatrix __all__ = ('KroneckerStencilMatrix', + 'ComposedKroneckerStencilMatrix', 'KroneckerLinearSolver', 'KroneckerDenseMatrix', 'kronecker_solve') @@ -34,13 +35,12 @@ class KroneckerStencilMatrix(LinearOperator): The factors are process-local: along its axes, the factor $A_k$ owns the same rows (``starts``/``ends``) as the codomain ``W`` on this process, - but it lives on its own spaces, without a communicator. + but it lives on its own spaces, without a communicator. The pads of the + factors must not exceed the pads of the domain ``V``, whose ghost regions + are read by ``dot``. - A product ``M = A @ B`` of two Kronecker matrices with the same axis - groups is again a KroneckerStencilMatrix with factors $C_k = A_k B_k$ - (see ``__matmul__``). The band of $C_k$ is wider than the ghost regions of - the domain, so the operands are stored in ``factors`` and ``dot`` applies - them one after another. + The product ``A @ B`` of two Kronecker matrices with the same axis groups + is a :class:`ComposedKroneckerStencilMatrix`. Parameters ---------- @@ -53,20 +53,11 @@ class KroneckerStencilMatrix(LinearOperator): *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. Unless - ``factors`` is given, the pads of ``A_k`` must not exceed the pads - of ``V``. - - factors : sequence of KroneckerStencilMatrix, optional - Operands $F_1, \dots, F_n$ with $M = F_1 F_2 \cdots F_n$, as created - by ``__matmul__``. If given, ``dot`` applies them from right to left. + ``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: StencilVectorSpace, - W: StencilVectorSpace, - *args: StencilMatrix, - factors: Sequence[KroneckerStencilMatrix] | None = None): + def __init__(self, V: StencilVectorSpace, W: StencilVectorSpace, *args: StencilMatrix): assert isinstance(V, StencilVectorSpace) assert isinstance(W, StencilVectorSpace) @@ -92,34 +83,17 @@ def __init__(self, assert d == V.ndim, \ f'The factors cover {d} axes, but the domain has {V.ndim}.' - if factors is None: - # 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 ' - f'the domain pads {V.pads}; pass the operands ' - 'as factors (as done by A @ B).') - tmp_vectors = () - else: - factors = tuple(factors) - assert len(factors) > 0 - for F in factors: - assert isinstance(F, KroneckerStencilMatrix) - assert F.axes == tuple(axes) - assert factors[0].codomain == W - assert factors[-1].domain == V - for F, G in zip(factors[:-1], factors[1:]): - assert F.domain == G.codomain - # intermediate results for dot - tmp_vectors = tuple(F.codomain.zeros() for F in factors[1:]) - - self._domain = V - self._codomain = W - self._mats = tuple(args) - self._axes = tuple(axes) - self._factors = factors - self._tmp_vectors = tmp_vectors + # 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 = tuple(args) + self._axes = tuple(axes) #-------------------------------------- # Abstract interface @@ -156,20 +130,11 @@ 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 factors(self) -> tuple[KroneckerStencilMatrix, ...] | None: - """Operands $F_1, \\dots, F_n$ with ``self == F_1 @ ... @ F_n``, or None if ``self`` is not a product.""" - return self._factors - # ... @property def nbytes(self) -> int: - """Local (per-MPI-rank) memory footprint of the factor matrices (and of the stored operands), in bytes.""" - nbytes = sum(getattr(mat, 'nbytes', 0) for mat in self._mats) - if self._factors is not None: - nbytes += sum(F.nbytes for F in self._factors) - return int(nbytes) + """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: StencilVector, out: StencilVector | None = None) -> StencilVector: @@ -199,13 +164,6 @@ def dot(self, x: StencilVector, out: StencilVector | None = None) -> StencilVect else: out = StencilVector(self.codomain) - # product of Kronecker matrices: apply the operands one after another - if self._factors is not None: - y = x - for F, tmp in zip(reversed(self._factors[1:]), reversed(self._tmp_vectors)): - y = F.dot(y, out=tmp) - return self._factors[0].dot(y, out=out) - # Necessary if vector space is periodic or distributed across processes if not x.ghost_regions_in_sync: x.update_ghost_regions() @@ -247,31 +205,23 @@ def dot(self, x: StencilVector, out: StencilVector | None = None) -> StencilVect # ... def copy(self) -> KroneckerStencilMatrix: - mats = [m.copy() for m in self.mats] - factors = None if self._factors is None else [F.copy() for F in self._factors] - return KroneckerStencilMatrix(self.domain, self.codomain, *mats, factors=factors) + mats = [m.copy() for m in self.mats] + return KroneckerStencilMatrix(self.domain, self.codomain, *mats) # ... def __neg__(self) -> KroneckerStencilMatrix: - mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])] - factors = None if self._factors is None else \ - [-self._factors[0], *(F.copy() for F in self._factors[1:])] - return KroneckerStencilMatrix(self.domain, self.codomain, *mats, factors=factors) + mats = [-self.mats[0], *(m.copy() for m in self.mats[1:])] + return KroneckerStencilMatrix(self.domain, self.codomain, *mats) # ... def __mul__(self, a) -> KroneckerStencilMatrix: - mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a] - factors = None if self._factors is None else \ - [*(F.copy() for F in self._factors[:-1]), self._factors[-1] * a] - return KroneckerStencilMatrix(self.domain, self.codomain, *mats, factors=factors) + mats = [*(m.copy() for m in self.mats[:-1]), self.mats[-1] * a] + return KroneckerStencilMatrix(self.domain, self.codomain, *mats) # ... def __imul__(self, a) -> KroneckerStencilMatrix: last = self._mats[-1] last *= a - if self._factors is not None: - last = self._factors[-1] - last *= a return self # ... @@ -279,19 +229,11 @@ def __matmul__(self, B): """ Product ``self @ B``. - If ``B`` is a KroneckerStencilMatrix with the same axis groups, the - result is a KroneckerStencilMatrix with factors - ``C_k = self.mats[k] @ B.mats[k]``, domain ``B.domain`` and codomain - ``self.codomain``. Copies of the operands are stored in - ``factors`` (products are flattened, so ``(A @ B) @ C`` has the - factors ``(A, B, C)``) and used by ``dot``. - - This is a collective operation if ``B.codomain`` is distributed: each - process needs all rows of ``B.mats[k]`` that its rows of - ``self.mats[k]`` couple to, so they are gathered. - - In all other cases (a different operator, other axis groups, or a - vector) the call is passed to ``LinearOperator.__matmul__``. + 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 ---------- @@ -300,97 +242,12 @@ def __matmul__(self, B): Returns ------- - KroneckerStencilMatrix | LinearOperator | Vector + ComposedKroneckerStencilMatrix | LinearOperator | Vector The product. """ - if not isinstance(B, KroneckerStencilMatrix) or B.axes != self.axes: - return super().__matmul__(B) - - assert self.domain == B.codomain, \ - 'The domain of the left operand must be the codomain of the right operand.' - - mats = [self._multiply_factors(A_k, B_k, B.codomain) - for A_k, B_k in zip(self.mats, B.mats)] - - factors = (*(self._factors or (self,)), *(B._factors or (B,))) - factors = [F.copy() for F in factors] - - return KroneckerStencilMatrix(B.domain, self.codomain, *mats, factors=factors) - - @staticmethod - 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 - B_sp = B.tosparse().tocoo() - 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 + if _is_kronecker_with_axes(B, self.axes): + return ComposedKroneckerStencilMatrix(B.domain, self.codomain, self, B) + return super().__matmul__(B) #-------------------------------------- # Other properties/methods @@ -409,15 +266,7 @@ def __getitem__(self, key): return reduce(lambda a, b: a * b, elements, 1) def tostencil(self) -> StencilMatrix: - """ - Convert to a StencilMatrix on the domain and codomain. - - Raises - ------ - ValueError - If the band of a factor exceeds the domain pads, e.g. for a product - created by ``A @ B``; use ``tosparse`` instead. - """ + """Convert to a StencilMatrix on the domain and codomain.""" mats = self.mats ssc = self.codomain.starts @@ -427,11 +276,6 @@ def tostencil(self) -> StencilMatrix: pads = [p for A in self.mats for p in A.pads] xpads = self.domain.pads - if any(p > xp_ for p,xp_ in zip(pads, xpads)): - raise ValueError(f'The pads {tuple(pads)} of the factors exceed the domain pads ' - f'{xpads}, so the matrix does not fit into a StencilMatrix on ' - 'its domain; use tosparse instead.') - # Number of rows in matrix (along each dimension) nrows = [ed-s+1 for s,ed in zip(ssd, eed)] nrows_extra = [0 if ec<=ed else ec-ed for ec,ed in zip(eec,eed)] @@ -500,16 +344,228 @@ def toarray(self): return self.tosparse().toarray() def transpose(self, conjugate: bool = False) -> KroneckerStencilMatrix: - """ - Transpose of the matrix (Hermitian transpose if ``conjugate`` is True). + """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) - The factors are transposed; for a product the order of the operands - is reversed. +#============================================================================== +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): """ - mats_tr = [Mi.transpose(conjugate=conjugate) for Mi in self.mats] - factors = None if self._factors is None else \ - [F.transpose(conjugate=conjugate) for F in reversed(self._factors)] - return KroneckerStencilMatrix(self.codomain, self.domain, *mats_tr, factors=factors) + 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 + B_sp = B.tosparse().tocoo() + 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): diff --git a/feectools/linalg/tests/test_kron_stencil_matrix.py b/feectools/linalg/tests/test_kron_stencil_matrix.py index 1ae484fa3..8d489298f 100644 --- a/feectools/linalg/tests/test_kron_stencil_matrix.py +++ b/feectools/linalg/tests/test_kron_stencil_matrix.py @@ -124,6 +124,7 @@ def test_KroneckerStencilMatrix(dtype, npts, pads, periodic): 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): @@ -227,7 +228,6 @@ def check_grouped_kron(comm, factor_ndims): assert M.ndim == 3 assert len(M.axes) == len(factor_ndims) - assert M.factors is None # dot assert_local_equal(W, M.dot(w), Md @ wglob) @@ -261,13 +261,11 @@ def check_matmul(comm, factor_ndims): rows = local_rows(W) C = A @ B - assert isinstance(C, KroneckerStencilMatrix) + 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 - assert len(C.factors) == 2 - for F, G in zip(C.factors, (A, B)): - assert F is not G - assert all(np.array_equal(xp.to_numpy(f._data), xp.to_numpy(g._data)) for f, g in zip(F.mats, G.mats)) + 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)) @@ -279,38 +277,39 @@ def check_matmul(comm, factor_ndims): assert C.dot(w, out=out) is out assert_local_equal(W, out, Ad @ Bd @ wglob) - # chains are flattened - D = C @ A - assert len(D.factors) == 3 - 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 - assert_local_equal(W, (C * 2.).dot(w), 2. * Ad @ Bd @ wglob) - assert_local_equal(W, (-C).dot(w), -Ad @ Bd @ wglob) - assert_local_equal(W, C.copy().dot(w), 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 np.allclose(C2.tosparse().tocsr()[rows].toarray(), 0.5 * (Ad @ Bd)[rows]) assert_local_equal(W, C.dot(w), Ad @ Bd @ wglob) + assert np.allclose(C.tosparse().tocsr()[rows].toarray(), (Ad @ Bd)[rows]) - # the band does not fit into the ghost regions - with pytest.raises(ValueError): - C.tostencil() + # 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) - AE = A @ E - assert isinstance(AE, ComposedLinearOperator) - assert_local_equal(W, AE.dot(w), Ad @ Ed @ wglob) + 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) @@ -328,6 +327,7 @@ def check_solver(comm, factor_ndims, mpi_dims_mask=None): # 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) @@ -384,3 +384,11 @@ def test_kron_solver_distributed_group_par(): _, _, _, 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) From 0546a028023d9cf036e1532f5000414840f19f43 Mon Sep 17 00:00:00 2001 From: Stefan Possanner Date: Thu, 8 Oct 2026 12:05:46 +0200 Subject: [PATCH 3/6] Add KroneckerSumSolver (fast diagonalization) Exact inverse of sum_d M_1 x ... x S_d x ... x M_n + sigma M_1 x ... x M_n (e.g. a Laplacian on a tensor-product grid): per direction the generalized eigenproblem S_d U_d = M_d U_d Lambda_d is solved, and A^{-1} = U Lambda^{-1} U^T with U = U_1 x ... x U_n. U^T and U are applied with KroneckerLinearSolver (also along distributed axes), Lambda^{-1} locally; vanishing eigenvalues are skipped (pseudo-inverse). A direction without stiffness term is given as None. Tests against dense solves, regular and singular, serial and with MPI. Co-Authored-By: Claude Opus 5.5 --- feectools/linalg/kron.py | 164 ++++++++++++++++++ .../linalg/tests/test_kron_stencil_matrix.py | 76 ++++++++ 2 files changed, 240 insertions(+) diff --git a/feectools/linalg/kron.py b/feectools/linalg/kron.py index 667f2565e..20859a8ad 100644 --- a/feectools/linalg/kron.py +++ b/feectools/linalg/kron.py @@ -17,6 +17,7 @@ __all__ = ('KroneckerStencilMatrix', 'ComposedKroneckerStencilMatrix', 'KroneckerLinearSolver', + 'KroneckerSumSolver', 'KroneckerDenseMatrix', 'kronecker_solve') @@ -1298,6 +1299,169 @@ def solve_pass(self, workmem, tempmem): synchronize_for_mpi(workmem, tempmem) self._comm.Alltoallv(targetargs, sourceargs) +#============================================================================== +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). + """ + + 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, diff --git a/feectools/linalg/tests/test_kron_stencil_matrix.py b/feectools/linalg/tests/test_kron_stencil_matrix.py index 8d489298f..709deb408 100644 --- a/feectools/linalg/tests/test_kron_stencil_matrix.py +++ b/feectools/linalg/tests/test_kron_stencil_matrix.py @@ -392,3 +392,79 @@ def test_multiplicants_deprecated(): 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) From c774ecf0985c23d8f6f851fb1606840eb9494726 Mon Sep 17 00:00:00 2001 From: Stefan Possanner Date: Thu, 8 Oct 2026 12:05:54 +0200 Subject: [PATCH 4/6] DirectionalDerivativeOperator: diffdir, negative and transposed properties Co-Authored-By: Claude Opus 5.5 --- feectools/feec/derivatives.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) 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) From e000fb438bc3c03edd22e83bb63e281f79b5bf86 Mon Sep 17 00:00:00 2001 From: Stefan Possanner Date: Thu, 8 Oct 2026 13:56:38 +0200 Subject: [PATCH 5/6] Kronecker matrices: factor row layout and periodic duplicates (review) - KroneckerStencilMatrix: the rows owned by a factor must contain the rows of the codomain on this process (ValueError otherwise). dot and tostencil index the factor rows by global row, so both process-local factors and factors owning all rows (e.g. from tokronstencil) work in parallel; before, other rows were silently used. - Products: sum duplicate COO entries of the local factor (periodic factors with 2p + 1 > n have two diagonals in the same column) before gathering the rows of other processes; before, np.unique dropped one of them in parallel. - Tests for full-row factors, the row check and periodic duplicates (serial and MPI). Co-Authored-By: Claude Opus 5.5 --- feectools/linalg/kron.py | 34 ++++++++-- .../linalg/tests/test_kron_stencil_matrix.py | 66 +++++++++++++++++-- 2 files changed, 89 insertions(+), 11 deletions(-) diff --git a/feectools/linalg/kron.py b/feectools/linalg/kron.py index 20859a8ad..214529ca9 100644 --- a/feectools/linalg/kron.py +++ b/feectools/linalg/kron.py @@ -34,9 +34,10 @@ class KroneckerStencilMatrix(LinearOperator): M = KroneckerStencilMatrix(V, W, A_xy, A_z) # M.axes == ((0, 1), (2,)) - The factors are process-local: along its axes, the factor $A_k$ owns the - same rows (``starts``/``ends``) as the codomain ``W`` on this process, - but it lives on its own spaces, without a communicator. The pads of the + 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``. @@ -84,6 +85,14 @@ def __init__(self, V: StencilVectorSpace, W: StencilVectorSpace, *args: StencilM 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): @@ -180,7 +189,7 @@ def dot(self, x: StencilVector, out: StencilVector | None = None) -> StencilVect # 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 = tuple(p*m for A in mats for p,m in zip(A.codomain.pads, A.codomain.shifts)) + 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) @@ -266,6 +275,16 @@ def __getitem__(self, key): for A,grp in zip(self.mats, self.axes)] return reduce(lambda a, b: a * b, elements, 1) + 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.""" @@ -285,7 +304,7 @@ def tostencil(self) -> StencilMatrix: M = StencilMatrix(self.domain, self.codomain, pads=tuple(pads)) # row offset of each axis in the data array of its factor - row_off = [p*m for A in mats for p,m in zip(A.codomain.pads, A.codomain.shifts)] + row_off = list(self._row_offsets()) mats = [mat._data for mat in mats] @@ -523,8 +542,11 @@ def _multiply_factors(A: StencilMatrix, B: StencilMatrix, U: StencilVectorSpace) 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 + # 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) diff --git a/feectools/linalg/tests/test_kron_stencil_matrix.py b/feectools/linalg/tests/test_kron_stencil_matrix.py index 709deb408..81b29b4ce 100644 --- a/feectools/linalg/tests/test_kron_stencil_matrix.py +++ b/feectools/linalg/tests/test_kron_stencil_matrix.py @@ -136,16 +136,21 @@ def make_space(comm, npts, pads, periods, mpi_dims_mask=None): return StencilVectorSpace(cart) -def make_factor(W, grp, mpads, seed): +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), - with band mpads and random entries; also returns the global dense matrix. + 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], @@ -201,14 +206,14 @@ def assert_local_equal(W, y, expected): 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): +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) + 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 @@ -468,3 +473,54 @@ def test_kron_sum_solver_ser(sigma, with_none): @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) From 9502142dde2bf8739e7b93c2e972a9d9221a7d13 Mon Sep 17 00:00:00 2001 From: Stefan Possanner Date: Thu, 8 Oct 2026 17:02:08 +0200 Subject: [PATCH 6/6] Bump version to 0.7.0; pyccel without -v in psydac_compile Co-Authored-By: Claude Opus 5.5 --- feectools/accelerate/compile_psydac.mk | 2 +- pyproject.toml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) 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/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"