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
35 changes: 33 additions & 2 deletions CUDA_STRATEGY.md
Original file line number Diff line number Diff line change
Expand Up @@ -234,6 +234,34 @@ Conventions, as in struphy:
- GPU tests are skipped without CuPy and a GPU (`cunumpy.kernel_testing.requires_cupy`). Until a GPU runner exists,
they are run by hand on an H100 before a PR that touches CUDA code is merged, and the PR description says so.

## Inner products on the device

`StencilVectorSpace.inner` (and so `StencilVector.inner`, `BlockVectorSpace.inner`, `dot_inner`) returns a 0-d
CuPy array for device data, in the serial and the MPI case (the `Allreduce` works on the device buffers, after
`synchronize_for_mpi`); before, the serial case copied the 8-byte result to the host. The result is a copy of the
reduction buffer, which the next inner product with the same vector overwrites. On NumPy it is a NumPy scalar, as
before.

- `StencilVectorSpace.axpy` with a 0-d device array computes `y += a * x` with array operations (the axpy kernel
takes `alpha` by value, so calling it would copy `a` to the host and wait for the device).
- CG, PCG, BiCG, BiCGStab and PBiCGStab keep their step sizes on the device; the only copy per iteration is the
residual norm of the convergence test (`solvers._host`). BiCGStab tested the residual twice per iteration (once
redundantly); it now tests it once. MINRES, LSMR and the Uzawa solver do their scalar recurrences on the host
and copy every inner product, as before; GMRES keeps its Arnoldi inner products on the device.
- Copies to the host per iteration, counted under the fake CuPy (`linalg/tests/test_inner_on_device.py`, which
also counts implicit conversions such as `float()`, invisible to `count_transfers`):

| Solver | before | after |
| --- | --- | --- |
| CG | 2 | 1 |
| PCG | 3 | 1 |
| BiCG | 4 | 1 |
| BiCGStab | 6 | 1 |
| PBiCGStab | 5 | 1 |
| MINRES, LSMR | 3 | 3 |

Iteration counts and solutions are identical to NumPy.

## Open questions

- **Setup kernels on the GPU.** B-spline, field evaluation and DOF kernels still run on the host through
Expand All @@ -245,8 +273,11 @@ Conventions, as in struphy:
layout first.
- **Complex data on the device.** Not needed by struphy so far; would need a second CUDA kernel per folder (or
dtype dispatch in `cunumpy.kernels.Kernel`).
- **`inner` reduction.** One `atomicAdd` per block of 256 threads, then a copy of the 8-byte result to the host in
the serial case (in the parallel case it goes into the MPI reduction). To be measured on the H100.
- **`inner` reduction.** One `atomicAdd` per block of 256 threads; the result stays on the device (see
[Inner products on the device](#inner-products-on-the-device)). To be measured on the H100.
- **Fewer convergence tests.** The Krylov solvers still copy the residual norm to the host once per iteration.
Testing it every k iterations would remove most of these synchronizations, at the cost of up to k - 1 extra
iterations (which must not divide by a zero residual) and of iteration counts that differ from NumPy.
- **Interface matrices** (`StencilInterfaceMatrix`) and the remaining stencil kernels (`stencil2coo`, ...) still use
`PyccelKernel` with host copies.
- **FFT on the device.** `DistributedFFT` and friends still stage their data through the host (SciPy FFT); they
Expand Down
4 changes: 3 additions & 1 deletion feectools/linalg/basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,9 +93,11 @@ def inner(self, x, y):

Returns
-------
float | complex
float | complex | cupy.ndarray
The scalar product of the two vectors. Note that inner(x, x) is
a non-negative real number which is zero if and only if x = 0.
For vectors with device (CuPy) data, a 0-d device array: the
result stays on the device.

"""

Expand Down
7 changes: 5 additions & 2 deletions feectools/linalg/block.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,9 +112,11 @@ def inner(self, x, y):

Returns
-------
float | complex
float | complex | cupy.ndarray
The scalar product of the two vectors. Note that inner(x, x) is
a non-negative real number which is zero if and only if x = 0.
For vectors with device (CuPy) data, a 0-d device array: the
result stays on the device.

"""

Expand All @@ -134,7 +136,8 @@ def axpy(self, a, x, y):
Parameters
----------
a : scalar
The scaling coefficient needed for the operation.
The scaling coefficient needed for the operation (a 0-d device
array is accepted, see `StencilVectorSpace.axpy`).

x : BlockVector
The vector which is not modified by this function.
Expand Down
94 changes: 62 additions & 32 deletions feectools/linalg/solvers.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,27 @@
'GMRES',
)

#===============================================================================
# Scalars on the device
#
# On the CuPy backend `Vector.inner` returns a 0-d device array. The Krylov
# solvers below (CG, PCG, BiCG, BiCGStab, PBiCGStab) keep their step sizes
# (alpha, beta, omega, ...) on the device: arithmetic with 0-d device arrays
# and `axpy` with a device scalar do not copy to the host. The only copy per
# iteration is the 8-byte residual norm of the convergence test, made explicit
# with `_host`. MINRES, LSMR and the Uzawa solver do their scalar recurrences
# on the host and copy every inner product (as before). On the NumPy backend
# nothing changes: both helpers return their argument as it is.
#===============================================================================
def _host(s):
"""A device scalar (0-d CuPy array) as a NumPy scalar (one copy to the host); anything else unchanged."""
return xp.to_numpy(s)[()] if xp.is_gpu(s) else s


def _sqrt(s):
"""Square root that keeps a device scalar on the device; `math.sqrt` for host scalars."""
return xp.sqrt(s) if xp.is_gpu(s) else sqrt(s)

#===============================================================================
def inverse(A, solver, **kwargs):
"""
Expand Down Expand Up @@ -205,11 +226,12 @@ def solve(self, b, out=None):
print( "+ Iter. # | L2-norm of residual |")
print( "+---------+---------------------+")
template = "| {:7d} | {:19.2e} |"
print(template.format(1, sqrt(am)))
print(template.format(1, sqrt(_host(am))))

# Iterate to convergence
for m in range(2, maxiter+1):
if am < tol_sqr:
# the only copy to the host per iteration (with device vectors)
if _host(am) < tol_sqr:
m -= 1
break
A.dot(p, out=v)
Expand All @@ -223,12 +245,13 @@ def solve(self, b, out=None):
p += r
am = am1
if verbose:
print(template.format(m, sqrt(am)))
print(template.format(m, sqrt(_host(am))))

if verbose:
print( "+---------+---------------------+")

# Convergence information
am = _host(am)
self._info = {'niter': m, 'success': am < tol_sqr, 'res_norm': sqrt(am) }

if recycle:
Expand Down Expand Up @@ -368,12 +391,13 @@ def solve(self, b, out=None):
print( "+ Iter. # | L2-norm of residual |")
print( "+---------+---------------------+")
template = "| {:7d} | {:19.2e} |"
print( template.format(1, sqrt(nrmr_sqr)))
print( template.format(1, sqrt(_host(nrmr_sqr))))

# Iterate to convergence
for k in range(2, maxiter+1):

if nrmr_sqr < tol_sqr:
# the only copy to the host per iteration (with device vectors)
if _host(nrmr_sqr) < tol_sqr:
k -= 1
break

Expand All @@ -395,12 +419,13 @@ def solve(self, b, out=None):
am = am1

if verbose:
print( template.format(k, sqrt(nrmr_sqr)))
print( template.format(k, sqrt(_host(nrmr_sqr))))

if verbose:
print( "+---------+---------------------+")

# Convergence information
nrmr_sqr = _host(nrmr_sqr)
self._info = {'niter': k, 'success': nrmr_sqr < tol_sqr, 'res_norm': sqrt(nrmr_sqr) }

if recycle:
Expand Down Expand Up @@ -540,7 +565,8 @@ def solve(self, b, out=None):
# Iterate to convergence
for m in range(1, maxiter + 1):

if res_sqr < tol_sqr:
# the only copy to the host per iteration (with device vectors)
if _host(res_sqr) < tol_sqr:
m -= 1
break

Expand Down Expand Up @@ -585,12 +611,13 @@ def solve(self, b, out=None):
ps += rs

if verbose:
print( template.format(m, sqrt(res_sqr)) )
print( template.format(m, sqrt(_host(res_sqr))) )

if verbose:
print( "+---------+---------------------+")

# Convergence information
res_sqr = _host(res_sqr)
self._info = {'niter': m, 'success': res_sqr < tol_sqr, 'res_norm': sqrt(res_sqr)}

if recycle:
Expand Down Expand Up @@ -729,12 +756,12 @@ def solve(self, b, out=None):
print("+---------+---------------------+")
template = "| {:7d} | {:19.2e} |"

# Iterate to convergence
for m in range(1, maxiter + 1):

if res_sqr < tol_sqr:
m -= 1
break
# Iterate to convergence. The residual is tested once per iteration,
# after its update (the only copy to the host per iteration with device
# vectors); before the first iteration it is tested here.
n_iter = 0 if _host(res_sqr) < tol_sqr else maxiter
m = 0
for m in range(1, n_iter + 1):

# -----------------------
# MATRIX-VECTOR PRODUCTS
Expand Down Expand Up @@ -771,7 +798,7 @@ def solve(self, b, out=None):
# ||r||_2 := (r, r)
res_sqr = r.inner(r).real

if res_sqr < tol_sqr:
if _host(res_sqr) < tol_sqr:
break

# b := a / w * (r0, r)_{m+1} / (r0, r)_m
Expand All @@ -783,12 +810,13 @@ def solve(self, b, out=None):
p.mul_iadd(-b * w, v)

if verbose:
print(template.format(m, sqrt(res_sqr)))
print(template.format(m, sqrt(_host(res_sqr))))

if verbose:
print("+---------+---------------------+")

# Convergence information
res_sqr = _host(res_sqr)
self._info = {'niter': m, 'success': res_sqr < tol_sqr, 'res_norm': sqrt(res_sqr)}

if recycle:
Expand Down Expand Up @@ -963,7 +991,8 @@ def solve(self, b, out=None):
# iterate to convergence or maximum number of iterations
niter = 0

while res_sqr > tol_sqr and niter < maxiter:
# the only copy to the host per iteration (with device vectors)
while _host(res_sqr) > tol_sqr and niter < maxiter:

# v = A @ pp, vp = PC @ v, alphap = rhop/(vp.rp0)
A.dot(pp, out=v)
Expand Down Expand Up @@ -1016,12 +1045,13 @@ def solve(self, b, out=None):
niter += 1

if verbose:
print(template.format(niter, sqrt(res_sqr)))
print(template.format(niter, sqrt(_host(res_sqr))))

if verbose:
print("+---------+---------------------+")

# convergence information
res_sqr = _host(res_sqr)
self._info = {'niter': niter, 'success': res_sqr <
tol_sqr, 'res_norm': sqrt(res_sqr)}

Expand Down Expand Up @@ -1171,7 +1201,7 @@ def solve(self, b, out=None):
y *= -1.0
y.copy(out=res_old) # res = b - A*x

beta = sqrt(res_old.inner(res_old))
beta = sqrt(_host(res_old.inner(res_old)))

# Initialize other quantities
oldb = 0
Expand Down Expand Up @@ -1215,15 +1245,15 @@ def solve(self, b, out=None):
if itn >= 2:
y.mul_iadd(-(beta/oldb), res_old)

alfa = v.inner(y)
alfa = _host(v.inner(y))
y.mul_iadd(-(alfa/beta), res_new)

# We put res_new in res_old and y in res_new
res_new, res_old = res_old, res_new
y.copy(out=res_new)

oldb = beta
beta = sqrt(res_new.inner(res_new))
beta = sqrt(_host(res_new.inner(res_new)))
tnorm2 += alfa**2 + oldb**2 + beta**2

# Apply previous rotation Qk-1 to get
Expand Down Expand Up @@ -1270,7 +1300,7 @@ def solve(self, b, out=None):
# Estimate various norms and test for convergence.

Anorm = sqrt(tnorm2)
ynorm = sqrt(x.inner(x))
ynorm = sqrt(_host(x.inner(x)))

rnorm = phibar
if ynorm == 0 or Anorm == 0:test1 = inf
Expand Down Expand Up @@ -1486,16 +1516,16 @@ def solve(self, b, out=None):
btol = tol

b.copy(out=u)
normb = sqrt(b.inner(b).real)
normb = sqrt(_host(b.inner(b)).real)

A.dot(x, out=u_work)
u -= u_work
beta = sqrt(u.inner(u).real)
beta = sqrt(_host(u.inner(u)).real)

if beta > 0:
u *= (1 / beta)
At.dot(u, out=v)
alpha = sqrt(v.inner(v).real)
alpha = sqrt(_host(v.inner(v)).real)
else:
x.copy(out=v)
alpha = 0
Expand Down Expand Up @@ -1558,14 +1588,14 @@ def solve(self, b, out=None):
u *= -alpha
A.dot(v, out=u_work)
u += u_work
beta = sqrt(u.inner(u).real)
beta = sqrt(_host(u.inner(u)).real)

if beta > 0:
u *= (1 / beta)
v *= -beta
At.dot(u, out=v_work)
v += v_work
alpha = sqrt(v.inner(v).real)
alpha = sqrt(_host(v.inner(v)).real)
if alpha > 0:v *= (1 / alpha)

# At this point, beta = beta_{k+1}, alpha = alpha_{k+1}.
Expand Down Expand Up @@ -1642,7 +1672,7 @@ def solve(self, b, out=None):

# Compute norms for convergence testing.
normar = abs(zetabar)
normx = sqrt(x.inner(x).real)
normx = sqrt(_host(x.inner(x)).real)

# Now use these norms to estimate certain other quantities,
# some of which will be small near a solution.
Expand Down Expand Up @@ -1807,7 +1837,7 @@ def solve(self, b, out=None):
A.dot( x , out=r)
r -= b

am = sqrt(r.inner(r).real)
am = sqrt(_host(r.inner(r).real))
if am < tol:
self._info = {'niter': 1, 'success': am < tol, 'res_norm': am }
return x
Expand Down Expand Up @@ -1884,7 +1914,7 @@ def arnoldi(self, k, p):
h[i] = p.inner(self._Q[i])
p.mul_iadd(-h[i], self._Q[i])

h[k+1] = sqrt(p.inner(p).real)
h[k+1] = _sqrt(p.inner(p).real)
p /= h[k+1] # Normalize vector

if len(self._Q) > k + 1:
Expand Down Expand Up @@ -2007,7 +2037,7 @@ def solve(self, b, out=None):

# constraint residual: R = B1*u + B2*ue - g
R = B1.dot(u) + B2.dot(ue) - g
residual_norm = sqrt(R.inner(R).real)
residual_norm = sqrt(_host(R.inner(R)).real)

if verbose:
print(template.format(iteration, residual_norm))
Expand All @@ -2017,7 +2047,7 @@ def solve(self, b, out=None):

# pressure update: steepest descent step size
S_R = B1.dot(A11inv.dot(B1.T.dot(R))) + B2.dot(A22inv.dot(B2.T.dot(R)))
alpha = R.inner(R).real / R.inner(S_R).real
alpha = _host(R.inner(R)).real / _host(R.inner(S_R)).real
p += alpha * R

if verbose:
Expand Down
Loading
Loading