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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions feectools/ddm/blocking_data_exchanger.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# coding: utf-8

import cunumpy as xp
from cunumpy import synchronize_for_mpi
import numpy as np
from feectools.ddm.mpi import mpi as MPI

Expand Down Expand Up @@ -82,6 +83,10 @@ def start_update_ghost_regions( self, array, requests ):

assert isinstance( array, xp.ndarray )

# MPI reads/writes `array` directly; on a device backend the
# kernels that produced it must have finished first.
synchronize_for_mpi( array )

# Shortcuts
cart = self._cart
comm = self._comm
Expand Down Expand Up @@ -123,6 +128,10 @@ def start_exchange_assembly_data( self, array ):

assert isinstance( array, xp.ndarray )

# MPI reads/writes `array` directly; on a device backend the
# kernels that produced it must have finished first.
synchronize_for_mpi( array )

# Shortcuts
cart = self._cart
comm = self._comm
Expand Down
5 changes: 5 additions & 0 deletions feectools/ddm/interface_data_exchanger.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
# coding: utf-8

from cunumpy import synchronize_for_mpi
from feectools.ddm.mpi import mpi as MPI

from .cart import InterfaceCartDecomposition, find_mpi_type
Expand Down Expand Up @@ -48,6 +49,10 @@ def update_ghost_regions( self, array_minus=None, array_plus=None ):

# ...
def start_update_ghost_regions( self, array_minus=None, array_plus=None ):
# MPI reads/writes these buffers directly; on a device backend the
# kernels that produced them must have finished first.
synchronize_for_mpi( array_minus, array_plus )

send_req = []
recv_req = []
cart = self._cart
Expand Down
8 changes: 3 additions & 5 deletions feectools/ddm/mpi.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,11 +187,9 @@ def _mpi_disabled():

if launched_under_mpi():
try:
# Disable MPI when using CuPy due to known segfault issues with OpenMPI + CUDA
import os
if os.environ.get('ARRAY_BACKEND') == 'cupy':
raise ImportError("MPI disabled when using CuPy backend")

# MPI with the CuPy backend needs a CUDA-aware MPI library: device
# buffers are passed to MPI directly (after synchronize_for_mpi, see
# the data exchangers). A non-CUDA-aware MPI segfaults on them.
if _mpi_disabled():
raise ImportError("MPI disabled (feectools.use_mpi = False or FEECTOOLS_MPI=0)")

Expand Down
8 changes: 8 additions & 0 deletions feectools/ddm/nonblocking_data_exchanger.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# coding: utf-8

import cunumpy as xp
from cunumpy import synchronize_for_mpi
import numpy as np
from itertools import product

Expand Down Expand Up @@ -98,6 +99,9 @@ def prepare_communications(self, u):
return tuple(requests)

def start_update_ghost_regions(self, array, requests ):
# The persistent requests read/write `array` directly; on a device
# backend the kernels that produced it must have finished first.
synchronize_for_mpi( array )
MPI.Prequest.Startall( requests )

def end_update_ghost_regions(self, array, requests):
Expand All @@ -108,6 +112,10 @@ def start_exchange_assembly_data( self, array ):

assert isinstance( array, xp.ndarray )

# MPI reads/writes `array` directly; on a device backend the
# kernels that produced it must have finished first.
synchronize_for_mpi( array )

# Shortcuts
cart = self._cart
comm = self._comm
Expand Down
4 changes: 2 additions & 2 deletions feectools/ddm/tests/test_cart_1d.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,7 +89,7 @@ def run_cart_1d( data_exchanger_type, verbose=False ):
#---------------------------------------------------------------------------

# Fill in true domain with u[i1_loc]=i1_glob
u[p1:-p1] = [i1 for i1 in range(s1,e1+1)]
u[p1:-p1] = xp.asarray([i1 for i1 in range(s1,e1+1)])

request = synchronizer.prepare_communications(u)

Expand All @@ -101,7 +101,7 @@ def run_cart_1d( data_exchanger_type, verbose=False ):
# CHECK RESULTS
#---------------------------------------------------------------------------
# Verify that ghost cells contain correct data (note periodic domain!)
success = all( u[:] == [i1%n1 for i1 in range(s1-p1,e1+p1+1)] )
success = bool( (u[:] == xp.asarray([i1%n1 for i1 in range(s1-p1,e1+p1+1)])).all() )

# MASTER only: collect information from all processes
success_global = comm.reduce( success, op=MPI.LAND, root=0 )
Expand Down
5 changes: 5 additions & 0 deletions feectools/linalg/kron.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

import numpy as np
import cunumpy as xp
from cunumpy import synchronize_for_mpi
from scipy.sparse import kron
from scipy.sparse import coo_matrix

Expand Down Expand Up @@ -885,6 +886,9 @@ def solve_pass(self, workmem, tempmem):
targetargs = [tempmem[:self._datasize], self._target_transfer, self._mpi_type]

# parts of stripes -> blocked stripes
# (MPI reads/writes the work arrays directly; on a device backend the
# kernels that produced them must have finished first.)
synchronize_for_mpi(workmem, tempmem)
self._comm.Alltoallv(sourceargs, targetargs)

# blocked stripes -> ordered stripes
Expand All @@ -897,6 +901,7 @@ def solve_pass(self, workmem, tempmem):
self._contiguous_to_blocked(workmem, tempmem)

# blocked stripes -> parts of stripes
synchronize_for_mpi(workmem, tempmem)
self._comm.Alltoallv(targetargs, sourceargs)

#==============================================================================
Expand Down
11 changes: 7 additions & 4 deletions feectools/linalg/stencil.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from types import MappingProxyType

import cunumpy as xp
from cunumpy import PyccelKernel
from cunumpy import PyccelKernel, synchronize_for_mpi
from cunumpy.xp import array_backend
from scipy.sparse import coo_matrix, diags as sp_diags

Expand Down Expand Up @@ -221,7 +221,8 @@ def _axpy_python(self, a, x, y):
@staticmethod
def _inner_python(v1, v2, nghost):
index = tuple(slice(ng, -ng) for ng in nghost)
return xp.vdot(v1[index].flat, v2[index].flat)
# ravel, not flat: CuPy's flatiter cannot be used as an array
return xp.vdot(v1[index].ravel(), v2[index].ravel())

#--------------------------------------
# Abstract interface
Expand Down Expand Up @@ -294,6 +295,8 @@ def inner(self, x, y):
if self.parallel:
# Sometimes in the parallel case, we can get an empty vector that breaks our kernel
x._dot_send_data[0] = 0 if x._data.shape[0] == 0 else inner_func(*inner_args)
# The send buffer was just written on the device; MPI reads it directly.
synchronize_for_mpi(x._dot_send_data, x._dot_recv_data)
self.cart.global_comm.Allreduce((x._dot_send_data, self.mpi_type),
(x._dot_recv_data, self.mpi_type),
op=MPI.SUM )
Expand Down Expand Up @@ -2594,7 +2597,7 @@ def _dot(mat, v, out, starts, nrows, nrows_extra, gpads, pads, dm, cm, c_axis, d
ii_kk = tuple( ii + kk )

ii[c_axis] += c_start
out[tuple(ii)] = xp.dot( mat[ii_kk].flat, v[jj].flat )
out[tuple(ii)] = xp.dot( mat[ii_kk].ravel(), v[jj].ravel() )


new_nrows = nrows.copy()
Expand All @@ -2616,7 +2619,7 @@ def _dot(mat, v, out, starts, nrows, nrows_extra, gpads, pads, dm, cm, c_axis, d
kk = [slice(None,n-e) for n,e in zip(ndiags, ee)]
ii_kk = tuple( ii + kk )
ii[c_axis] += c_start
out[tuple(ii)] = xp.dot( mat[ii_kk].flat, v[jj].flat )
out[tuple(ii)] = xp.dot( mat[ii_kk].ravel(), v[jj].ravel() )

new_nrows[d] += er

Expand Down
Loading
Loading