Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
164 changes: 137 additions & 27 deletions feectools/linalg/basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@

"""

import itertools
from abc import ABC, abstractmethod
from types import LambdaType
from inspect import signature
Expand Down Expand Up @@ -280,13 +281,143 @@ def dtype(self):
upon convertion to matrix.
"""

@abstractmethod
def tosparse(self):
""" Convert to a sparse matrix in any of the formats supported by scipy.sparse."""
def toarray(self, out=None, is_sparse=False, format='csr'):
"""
Assemble the global matrix of the linear operator column by column.

@abstractmethod
def toarray(self):
""" Convert to Numpy 2D array. """
Column j is computed as ``self.dot(e_j)``, where e_j is the j-th
canonical basis vector of the domain, in the global numbering of
``Vector.toarray()`` (C-ordered within each StencilVector, blocks
concatenated). Since only ``dot`` is used, this works for any linear
operator, including matrix-free ones, whose domain is a
StencilVectorSpace or a (possibly nested) BlockVectorSpace thereof.

The cost is one call to ``dot`` per global degree of freedom of the
domain, hence this default is meant for testing and small problems.
Subclasses with an explicit matrix representation should override it.

In parallel, all ranks call ``dot`` collectively, each rank fills the
rows it owns, and every rank receives the full matrix.

Parameters
----------
out : numpy.ndarray, optional
Dense array of shape ``self.shape`` into which the result is
written in place. Must be None if ``is_sparse`` is True.

is_sparse : bool
If True, return a scipy.sparse matrix, otherwise a dense
numpy.ndarray.

format : str
Sparse format, only used if ``is_sparse`` is True: one of 'csr'
(default), 'csc', 'bsr', 'lil', 'dok', 'coo' or 'dia'.

Returns
-------
numpy.ndarray or scipy.sparse matrix
The global matrix of shape ``self.shape``, identical on all ranks.
"""
from feectools.linalg.block import BlockVectorSpace
from feectools.linalg.stencil import StencilVectorSpace

# Flatten the (possibly nested) block structure of the domain
def leaves(w):
if isinstance(w.space, StencilVectorSpace):
return [w]
elif isinstance(w.space, BlockVectorSpace):
return [lw for b in w.blocks for lw in leaves(b)]
else:
raise TypeError(f'{type(self).__name__}.toarray() requires a domain made of '
f'StencilVectorSpaces, not {type(w.space).__name__}.')

e_j = self.domain.zeros()
Ae_j = self.codomain.zeros()
e_j_leaves = leaves(e_j)
offsets = xp.cumsum([0] + [lw.space.dimension for lw in e_j_leaves[:-1]])

if is_sparse:
assert out is None, 'out must be None if is_sparse is True.'
assert format in ('csr', 'csc', 'bsr', 'lil', 'dok', 'coo', 'dia'), \
f'Unknown sparse format {format!r}.'
rows, cols, data = [], [], []
elif out is None:
out = xp.zeros(self.shape, dtype=self.dtype)
else:
assert isinstance(out, xp.ndarray)
assert out.shape == self.shape, f'out has shape {out.shape}, expected {self.shape}.'

# Index ranges owned by each rank, for every leaf
bounds = [(lw.starts, lw.ends) for lw in e_j_leaves]
if e_j_leaves[0].space.parallel:
comm = e_j_leaves[0].space.cart.comm
rank = comm.Get_rank()
all_bounds = comm.allgather(bounds)
else:
comm = None
rank = 0
all_bounds = [bounds]

# All ranks loop over all columns, since dot() is collective;
# only the owner of index i sets the entry of e_j to one.
for owner, owner_bounds in enumerate(all_bounds):
for lw, offset, (starts, ends) in zip(e_j_leaves, offsets, owner_bounds):
for i in itertools.product(*(range(s, e + 1) for s, e in zip(starts, ends))):
if rank == owner:
lw[i] = 1
lw.update_ghost_regions()
self.dot(e_j, out=Ae_j)
if rank == owner:
lw[i] = 0

# Global vector with nonzeros only in the rows owned by this rank
col_j = Ae_j.toarray()
j = offset + xp.ravel_multi_index(i, lw.space.npts)
if is_sparse:
nz = xp.flatnonzero(col_j)
rows.append(nz)
cols.append(xp.full(nz.size, j))
data.append(col_j[nz])
else:
out[:, j] = col_j
# Clear the ghost regions of the last entry set to one
lw.update_ghost_regions()

if not is_sparse:
if comm is not None:
from feectools.ddm.mpi import mpi as MPI
comm.Allreduce(MPI.IN_PLACE, out, op=MPI.SUM)
return out

rows = xp.concatenate(rows) if rows else xp.zeros(0, dtype=int)
cols = xp.concatenate(cols) if cols else xp.zeros(0, dtype=int)
data = xp.concatenate(data) if data else xp.zeros(0, dtype=self.dtype)
if comm is not None:
rows = xp.concatenate(comm.allgather(rows))
cols = xp.concatenate(comm.allgather(cols))
data = xp.concatenate(comm.allgather(data))

return coo_matrix((data, (rows, cols)), shape=self.shape).asformat(format)

def tosparse(self, format='csr'):
"""
Assemble the global matrix of the linear operator as a scipy.sparse matrix.

Default implementation calling the generic ``LinearOperator.toarray``
with ``is_sparse=True``; see there for cost and parallel behavior.
Subclasses with an explicit matrix representation should override it.

Parameters
----------
format : str
One of 'csr' (default), 'csc', 'bsr', 'lil', 'dok', 'coo' or 'dia'.

Returns
-------
scipy.sparse matrix
The global matrix of shape ``self.shape``, identical on all ranks.
"""
return LinearOperator.toarray(self, is_sparse=True, format=format)

@abstractmethod
def dot(self, v, out=None):
Expand Down Expand Up @@ -977,9 +1108,6 @@ def multiplicants(self):
def dtype(self):
return None

def toarray(self):
raise NotImplementedError('toarray() is not defined for ComposedLinearOperators.')

def tosparse(self):
mats = [M.tosparse() for M in self._multiplicants]
M = mats[0]
Expand Down Expand Up @@ -1084,12 +1212,6 @@ def factorial(self):
""" Returns the power to which the operator is raised. """
return self._factorial

def toarray(self):
raise NotImplementedError('toarray() is not defined for PowerLinearOperators.')

def tosparse(self):
raise NotImplementedError('tosparse() is not defined for PowerLinearOperators.')

def transpose(self, conjugate=False):
return PowerLinearOperator(domain=self.codomain, codomain=self.domain, A=self._operator.transpose(conjugate=conjugate), n=self._factorial)

Expand Down Expand Up @@ -1207,12 +1329,6 @@ def _check_options(self, **kwargs):
elif key == 'verbose':
assert isinstance(value, bool), "verbose must be a bool"

def toarray(self):
raise NotImplementedError('toarray() is not defined for InverseLinearOperators.')

def tosparse(self):
raise NotImplementedError('tosparse() is not defined for InverseLinearOperators.')

def get_info(self):
""" Returns the previous convergence information. """
return self._info
Expand Down Expand Up @@ -1366,12 +1482,6 @@ def dot(self, v, out=None, **kwargs):

return out

def toarray(self):
raise NotImplementedError('toarray() is not defined for MatrixFreeLinearOperator.')

def tosparse(self):
raise NotImplementedError('tosparse() is not defined for MatrixFreeLinearOperator.')

def transpose(self, conjugate=False):
if self._dot_transpose is None:
raise NotImplementedError('no transpose dot method was given -- cannot create the transpose operator')
Expand Down
6 changes: 0 additions & 6 deletions feectools/linalg/fft.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,12 +21,6 @@ class DistributedFFTBase(LinearOperator):
The function at position i is applied to the i-th tensor direction.
If only a single callable is given, it is used for all directions.
"""
def toarray(self):
raise NotImplementedError('toarray() is not defined for DistributedFFTBase.')

def tosparse(self):
raise NotImplementedError('tosparse() is not defined for DistributedFFTBase.')

# Possible additions for the future:
# * split off the LinearSolver class when used with the space ndarray (as used in the KroneckerLinearSolver),
# and make it state if it works in-place (or if it needs temporary memory), and what its optimal
Expand Down
6 changes: 0 additions & 6 deletions feectools/linalg/kron.py
Original file line number Diff line number Diff line change
Expand Up @@ -528,12 +528,6 @@ def codomain(self):
@property
def dtype(self):
return None

def toarray(self):
raise NotImplementedError('toarray() is not defined for KroneckerLinearSolvers.')

def tosparse(self):
raise NotImplementedError('tosparse() is not defined for KroneckerLinearSolvers.')

def transpose(self, conjugate=False):
new_domain = self._codomain
Expand Down
162 changes: 162 additions & 0 deletions feectools/linalg/tests/test_toarray.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
# -*- coding: UTF-8 -*-
#
# Tests for the generic LinearOperator.toarray() and LinearOperator.tosparse(),
# which assemble the matrix of any linear operator from its dot() method.
#
import pytest
import cunumpy as xp

from feectools.ddm.mpi import mpi as MPI
from feectools.ddm.cart import DomainDecomposition, CartDecomposition
from feectools.linalg.basic import LinearOperator, MatrixFreeLinearOperator, IdentityOperator
from feectools.linalg.stencil import StencilVectorSpace, StencilMatrix
from feectools.linalg.block import BlockVectorSpace, BlockLinearOperator
from feectools.linalg.solvers import inverse
from feectools.linalg.tests.test_block import compute_global_starts_ends

SPARSE_FORMATS = ['csr', 'csc', 'bsr', 'lil', 'dok', 'coo', 'dia']

#===============================================================================
# HELPERS
#===============================================================================
def get_space(npts, pads, periods, comm=None):
D = DomainDecomposition(npts, periods=periods, comm=comm)
global_starts, global_ends = compute_global_starts_ends(D, npts)
C = CartDecomposition(D, npts, global_starts, global_ends, pads=pads, shifts=[1] * len(npts))
return StencilVectorSpace(C)

def get_random_matrix(V, W, seed):
""" Random StencilMatrix, identical on all ranks in the global numbering. """
rng = xp.random.default_rng(seed)
M = StencilMatrix(V, W)
# Draw the full (global) band and keep the local part, so the matrix
# does not depend on the domain decomposition.
shape = tuple(W.npts) + M._data.shape[W.ndim:]
band = rng.random(shape) - 0.5
idx = tuple(slice(s, e + 1) for s, e in zip(W.starts, W.ends))
M[idx] = band[idx]
M.remove_spurious_entries()
return M

def matrix_free(A):
""" Hide the explicit matrix of A, so that the generic toarray() is used. """
return MatrixFreeLinearOperator(A.domain, A.codomain, lambda v, out=None: A.dot(v, out=out))

def reference(A):
""" Global dense matrix of A, from its own (row-local in parallel) tosparse(). """
local = A.tosparse().toarray()
comm = A.domain.cart.comm if isinstance(A.domain, StencilVectorSpace) else A.domain.spaces[0].cart.comm
if A.domain.parallel:
glob = xp.zeros_like(local)
comm.Allreduce(local, glob, op=MPI.SUM)
return glob
return local

def check_all(O, ref):
""" Check dense, in-place and all sparse outputs of the generic toarray()/tosparse(). """
assert xp.allclose(O.toarray(), ref)

out = xp.full(O.shape, 7.0)
res = O.toarray(out=out)
assert res is out
assert xp.allclose(out, ref)

for fmt in SPARSE_FORMATS:
S = O.toarray(is_sparse=True, format=fmt)
assert S.format == fmt
assert S.shape == O.shape
assert xp.allclose(S.toarray(), ref)

S = O.tosparse()
assert S.format == 'csr'
assert xp.allclose(S.toarray(), ref)
assert O.tosparse('csc').format == 'csc'

def stencil_case(periods, comm):
V = get_space([6, 5], [2, 1], periods, comm=comm)
return get_random_matrix(V, V, seed=0)

def block_case(periods, comm):
""" 2x2 block operator, and a nested block operator [[B, 0], [0, A]]. """
V = get_space([6, 5], [2, 1], periods, comm=comm)
A = get_random_matrix(V, V, seed=0)
A01 = get_random_matrix(V, V, seed=1)
VV = BlockVectorSpace(V, V)
B = BlockLinearOperator(VV, VV, blocks=[[A, A01], [None, A]])
Vn = BlockVectorSpace(VV, V)
N = BlockLinearOperator(Vn, Vn, blocks=[[B, None], [None, A]])
return A, B, N

#===============================================================================
# SERIAL TESTS
#===============================================================================
@pytest.mark.parametrize('periods', [[False, False], [True, False], [True, True]])
def test_toarray_stencil(periods):
A = stencil_case(periods, comm=None)
check_all(matrix_free(A), reference(A))

@pytest.mark.parametrize('periods', [[False, False], [True, True]])
def test_toarray_block(periods):
A, B, N = block_case(periods, comm=None)
refA, refB = reference(A), reference(B)
check_all(matrix_free(B), refB)

Z = xp.zeros((refB.shape[0], refA.shape[1]))
refN = xp.block([[refB, Z], [Z.T, refA]])
check_all(matrix_free(N), refN)

def test_toarray_composite_operators():
""" Operators that used to raise NotImplementedError in toarray()/tosparse(). """
A = stencil_case([True, False], comm=None)
ref = reference(A)
V = A.domain

assert xp.allclose((A @ A).toarray(), ref @ ref)
assert xp.allclose((A ** 2).toarray(), ref @ ref)
assert xp.allclose((A ** 2).tosparse().toarray(), ref @ ref)

M = A.T @ A + IdentityOperator(V)
refM = ref.T @ ref + xp.eye(ref.shape[0])
Minv = inverse(M, 'cg', tol=1e-13, maxiter=1000)
assert xp.allclose(Minv.toarray(), xp.linalg.inv(refM), atol=1e-8)

def test_toarray_invalid_input():
O = matrix_free(stencil_case([False, False], comm=None))
with pytest.raises(AssertionError):
O.toarray(out=xp.zeros((O.shape[0], O.shape[1] + 1)))
with pytest.raises(AssertionError):
O.toarray(out=xp.zeros(O.shape), is_sparse=True)
with pytest.raises(AssertionError):
O.toarray(is_sparse=True, format='xyz')

#===============================================================================
# PARALLEL TESTS
#===============================================================================
@pytest.mark.parametrize('periods', [[False, False], [True, False], [True, True]])
@pytest.mark.parallel
def test_toarray_stencil_parallel(periods):
A = stencil_case(periods, comm=MPI.COMM_WORLD)
check_all(matrix_free(A), reference(A))

@pytest.mark.parametrize('periods', [[False, False], [True, True]])
@pytest.mark.parallel
def test_toarray_block_parallel(periods):
A, B, N = block_case(periods, comm=MPI.COMM_WORLD)
refA, refB = reference(A), reference(B)
check_all(matrix_free(B), refB)

Z = xp.zeros((refB.shape[0], refA.shape[1]))
refN = xp.block([[refB, Z], [Z.T, refA]])
check_all(matrix_free(N), refN)

@pytest.mark.parallel
def test_toarray_composite_operators_parallel():
A = stencil_case([True, False], comm=MPI.COMM_WORLD)
ref = reference(A)
assert xp.allclose((A @ A).toarray(), ref @ ref)
assert xp.allclose((A ** 2).tosparse().toarray(), ref @ ref)

#===============================================================================
if __name__ == '__main__':
import sys
pytest.main(sys.argv)
Loading
Loading