Skip to content

A heat equation with halo cells

Download notebookView on GitHubRun when the site was built.

The 2D heat equation ∂tu=α∇2u\partial_t u = \alpha \nabla^2 u, solved with a five-point stencil on a grid split over the MPI ranks. Each step updates the halo cells from the neighbours, with a zero-flux condition at the walls, and applies the stencil to every rank’s block. Run it serially, or with mpiexec -n 4 python heat_equation.py: the result is the same.

import cunumpy as xp
import matplotlib.pyplot as plt
import mpiarray as mpa
n = 128 # grid points per axis
h = 1.0 / n # grid spacing
alpha = 1.0 # diffusivity
dt = 0.2 * h**2 / alpha # stable for the explicit scheme (needs dt <= h**2 / 4)
def hot_spot(i, j):
"""Initial temperature: a Gaussian, from the global indices of each cell.
The indices come as the backend's arrays (NumPy or CuPy), hence ``xp``.
"""
x, y = (i + 0.5) * h, (j + 0.5) * h
return xp.exp(-((x - 0.3) ** 2 + (y - 0.6) ** 2) / 0.005)
u = mpa.fromfunction(hot_spot, (n, n), split=(0, 1), halo=1)
rank = u.layout.rank
print(u.layout)
Layout(shape=(128, 128), split=(0, 1), halo=(1, 1), periodic=(False, False), process_grid=(1, 1), rank=0 of 1)

fromfunction computes every rank’s block from its global indices, so the whole grid never exists on one rank. split=(0, 1) cuts both axes, and halo=1 gives every block one layer of halo cells, which is what the five-point stencil reads.

update_halos copies the neighbours’ boundary values into the halo cells. At the walls there is no neighbour; boundary="edge" copies the nearest cell instead, so no heat flows through the walls. The stencil reads the storage with its halos and writes the block.

def laplacian(v):
"""The five-point Laplacian of the cells inside a storage array ``v``."""
return (
v[2:, 1:-1] + v[:-2, 1:-1] + v[1:-1, 2:] + v[1:-1, :-2] - 4 * v[1:-1, 1:-1]
) / h**2
def step(u):
u.update_halos(boundary="edge") # zero flux through the walls
u.local[...] += dt * alpha * laplacian(u.local_with_halos)
total, peak = u.sum(), u.max()
for _ in range(400):
step(u)
new_total, new_peak = u.sum(), u.max() # collective: on every rank
if rank == 0:
print(f"total heat: {total:.6f} -> {new_total:.6f}")
print(f"peak: {peak:.4f} -> {new_peak:.4f}")
total heat: 257.359270 -> 257.359270
peak: 0.9988 -> 0.2036

The total heat stays the same, to rounding: with zero flux at the walls the stencil only moves heat around. sum and max are collective: every rank gets the same value.

to_numpy(root=0) assembles the whole grid on rank 0, as a NumPy array for plotting (gather() keeps the backend’s array type, NumPy or CuPy). It is collective, so it is called on every rank, outside the if; only rank 0 plots.

field = u.to_numpy(root=0) # NumPy, also with CuPy storage
if rank == 0:
fig, ax = plt.subplots(figsize=(4, 3.5))
image = ax.imshow(field.T, origin="lower", extent=(0, 1, 0, 1), cmap="inferno")
fig.colorbar(image, ax=ax, label="u")
ax.set(xlabel="x", ylabel="y", title="after 400 steps")
Notebook output

With wait=False, update_halos starts the messages and returns at once. The cells that do not touch a halo can be updated meanwhile; wait() then completes the halos, and the frame of the block follows. The result is the same as with wait=True, corner halo cells included. The array must not be written before wait(), so the new values go to a separate array first.

def step_overlapped(u):
pending = u.update_halos(wait=False, boundary="edge")
v = u.local_with_halos
change = xp.empty_like(u.local)
change[1:-1, 1:-1] = laplacian(v[1:-1, 1:-1]) # needs no halo cells
pending.wait()
# the frame of the block: each edge is the Laplacian of a three-cell strip
change[0, :] = laplacian(v[:3, :])[0]
change[-1, :] = laplacian(v[-3:, :])[0]
change[1:-1, 0] = laplacian(v[1:-1, :3])[:, 0]
change[1:-1, -1] = laplacian(v[1:-1, -3:])[:, 0]
u.local[...] += dt * alpha * change
a = mpa.fromfunction(hot_spot, (n, n), split=(0, 1), halo=1)
b = mpa.fromfunction(hot_spot, (n, n), split=(0, 1), halo=1)
for _ in range(50):
step(a)
step_overlapped(b)
same = bool(xp.allclose(a.gather(), b.gather()))
if rank == 0:
print("same result:", same)
same result: True