Skip to content

Stencil Operators

stencil

Finite-difference stencil operators for lattice fields.

This module provides a unified implementation of spatial derivative operators using finite difference stencils. These work for any dimension (1D, 2D, 3D) with periodic boundary conditions.

For FFT-based spectral methods, see spectral.py.

All operators are JIT-compiled for performance. Any operator that takes an axis (or other structural, non-array) argument must mark it static -- @partial(jit, static_argnames="axis") -- because jnp.roll(..., axis=...) requires a concrete Python index and will raise ConcretizationTypeError on a tracer. Note that operators which merely index an array with an integer (as in operators/gauge.py) need no such treatment: dynamic indexing traces fine.

laplacian(field, dx)

Compute Laplacian \(\nabla^2 \phi\) using second-order finite differences.

Uses the standard 3-point stencil in each dimension:

\[\nabla^2 \phi = \Sigma_i [\phi(n+\hat{i}) + \phi(n-\hat{i}) - 2\phi(n)] / dx^2\]

Works for 1D, 2D, or 3D fields with periodic boundary conditions.

Parameters:

Name Type Description Default
field Array

Field array (1D, 2D, or 3D)

required
dx float

Lattice spacing (assumed uniform in all directions)

required

Returns:

Type Description
Array

Laplacian at each lattice site (same shape as input)

Source code in jaxlatt/operators/stencil.py
@jit
def laplacian(field: Array, dx: float) -> Array:
    """
    Compute Laplacian $\\nabla^2 \\phi$ using second-order finite differences.

    Uses the standard 3-point stencil in each dimension:

    $$\\nabla^2 \\phi = \\Sigma_i [\\phi(n+\\hat{i}) + \\phi(n-\\hat{i}) - 2\\phi(n)] / dx^2$$

    Works for 1D, 2D, or 3D fields with periodic boundary conditions.

    Args:
        field: Field array (1D, 2D, or 3D)
        dx: Lattice spacing (assumed uniform in all directions)

    Returns:
        Laplacian at each lattice site (same shape as input)
    """
    result = jnp.zeros_like(field)
    ndim = field.ndim

    for axis in range(ndim):
        result = result + (
            jnp.roll(field, -1, axis=axis) + jnp.roll(field, 1, axis=axis) - 2.0 * field
        )

    return result / (dx**2)

gradient_squared(field, dx)

Compute \(|\nabla\phi|^2\) using centered finite differences.

\[|\nabla\phi|^2 = \Sigma_i |\partial_i\phi|^2 \quad \text{where} \quad \partial_i\phi = [\phi(n+\hat{i}) - \phi(n-\hat{i})] / (2 \, dx)\]

Works for 1D, 2D, or 3D fields. For complex fields, returns \(|\nabla\phi|^2\).

Parameters:

Name Type Description Default
field Array

Field array (1D, 2D, or 3D), can be real or complex

required
dx float

Lattice spacing

required

Returns:

Type Description
Array

Gradient squared at each site (real-valued, same shape as input)

Source code in jaxlatt/operators/stencil.py
@jit
def gradient_squared(field: Array, dx: float) -> Array:
    """
    Compute $|\\nabla\\phi|^2$ using centered finite differences.

    $$|\\nabla\\phi|^2 = \\Sigma_i |\\partial_i\\phi|^2 \\quad \\text{where} \\quad \\partial_i\\phi = [\\phi(n+\\hat{i}) - \\phi(n-\\hat{i})] / (2 \\, dx)$$

    Works for 1D, 2D, or 3D fields. For complex fields, returns $|\\nabla\\phi|^2$.

    Args:
        field: Field array (1D, 2D, or 3D), can be real or complex
        dx: Lattice spacing

    Returns:
        Gradient squared at each site (real-valued, same shape as input)
    """
    grad_sq = jnp.zeros_like(field.real)
    ndim = field.ndim

    for axis in range(ndim):
        derivative = (jnp.roll(field, -1, axis=axis) - jnp.roll(field, 1, axis=axis)) / (2.0 * dx)
        grad_sq = grad_sq + jnp.abs(derivative) ** 2

    return grad_sq

gradient(field, dx)

Compute gradient \(\nabla\phi\) using centered finite differences.

Returns a tuple of arrays, one for each spatial direction.

Parameters:

Name Type Description Default
field Array

Field array (1D, 2D, or 3D)

required
dx float

Lattice spacing

required

Returns:

Type Description
tuple[Array, ...]

Tuple of gradient components \((\partial_x \phi, \partial_y \phi, \ldots)\) for each dimension

Source code in jaxlatt/operators/stencil.py
@jit
def gradient(field: Array, dx: float) -> tuple[Array, ...]:
    """
    Compute gradient $\\nabla\\phi$ using centered finite differences.

    Returns a tuple of arrays, one for each spatial direction.

    Args:
        field: Field array (1D, 2D, or 3D)
        dx: Lattice spacing

    Returns:
        Tuple of gradient components $(\\partial_x \\phi, \\partial_y \\phi, \\ldots)$ for each dimension
    """
    ndim = field.ndim
    grads = []

    for axis in range(ndim):
        grad_i = (jnp.roll(field, -1, axis=axis) - jnp.roll(field, 1, axis=axis)) / (2.0 * dx)
        grads.append(grad_i)

    return tuple(grads)

divergence(vector_field, dx)

Compute divergence \(\nabla \cdot F\) for a vector field.

Parameters:

Name Type Description Default
vector_field Array

Array of shape (ndim, *spatial_shape) First axis is the vector component index.

required
dx float

Lattice spacing

required

Returns:

Type Description
Array

Divergence at each site (shape = spatial_shape)

Source code in jaxlatt/operators/stencil.py
@jit
def divergence(vector_field: Array, dx: float) -> Array:
    """
    Compute divergence $\\nabla \\cdot F$ for a vector field.

    Args:
        vector_field: Array of shape (ndim, *spatial_shape)
            First axis is the vector component index.
        dx: Lattice spacing

    Returns:
        Divergence at each site (shape = spatial_shape)
    """
    ndim = vector_field.shape[0]
    div = jnp.zeros(vector_field.shape[1:])

    for i in range(ndim):
        # Central difference for $\\partial_i F_i$
        div = div + (
            jnp.roll(vector_field[i], -1, axis=i) - jnp.roll(vector_field[i], 1, axis=i)
        ) / (2.0 * dx)

    return div

forward_gradient(field, dx, axis)

Compute forward finite difference \(\partial_i \phi = [\phi(n+\hat{i}) - \phi(n)] / dx\).

Useful for gauge-covariant derivatives and link-based calculations.

Parameters:

Name Type Description Default
field Array

Field array

required
dx float

Lattice spacing

required
axis int

Direction of derivative (0, 1, or 2). Static: it selects the jnp.roll axis, so it must be a concrete Python int, not a traced value.

required

Returns:

Type Description
Array

Forward derivative along specified axis

Source code in jaxlatt/operators/stencil.py
@partial(jit, static_argnames="axis")
def forward_gradient(field: Array, dx: float, axis: int) -> Array:
    """
    Compute forward finite difference $\\partial_i \\phi = [\\phi(n+\\hat{i}) - \\phi(n)] / dx$.

    Useful for gauge-covariant derivatives and link-based calculations.

    Args:
        field: Field array
        dx: Lattice spacing
        axis: Direction of derivative (0, 1, or 2). Static: it selects the
            ``jnp.roll`` axis, so it must be a concrete Python int, not a
            traced value.

    Returns:
        Forward derivative along specified axis
    """
    return (jnp.roll(field, -1, axis=axis) - field) / dx

backward_gradient(field, dx, axis)

Compute backward finite difference \(\partial_i \phi = [\phi(n) - \phi(n-\hat{i})] / dx\).

Parameters:

Name Type Description Default
field Array

Field array

required
dx float

Lattice spacing

required
axis int

Direction of derivative (0, 1, or 2). Static: it selects the jnp.roll axis, so it must be a concrete Python int, not a traced value.

required

Returns:

Type Description
Array

Backward derivative along specified axis

Source code in jaxlatt/operators/stencil.py
@partial(jit, static_argnames="axis")
def backward_gradient(field: Array, dx: float, axis: int) -> Array:
    """
    Compute backward finite difference $\\partial_i \\phi = [\\phi(n) - \\phi(n-\\hat{i})] / dx$.

    Args:
        field: Field array
        dx: Lattice spacing
        axis: Direction of derivative (0, 1, or 2). Static: it selects the
            ``jnp.roll`` axis, so it must be a concrete Python int, not a
            traced value.

    Returns:
        Backward derivative along specified axis
    """
    return (field - jnp.roll(field, 1, axis=axis)) / dx

hessian_squared(field, dx)

Pointwise squared Frobenius norm of the Hessian matrix.

Computes

\[\sum_{i,j} (\partial_i \partial_j \phi)^2\]

using second-order centred finite differences for diagonal terms and forward-then-backward (mixed) differences for off-diagonal terms. Exploits Hessian symmetry — each off-diagonal pair is computed once and weighted by 2. Periodic boundary conditions throughout.

Arises in higher-derivative scalar field equations, where a combination such as :math:(\Box\phi)^2 - (\nabla_\mu\nabla_\nu\phi)^2 cancels the second-time-derivative piece and leaves this purely spatial contraction behind.

Parameters:

Name Type Description Default
field Array

Real scalar field array (1-D, 2-D, or 3-D).

required
dx float

Lattice spacing (uniform in all directions).

required

Returns:

Type Description
Array

Array of the same shape as field with the squared Hessian norm

Array

at each site.

Source code in jaxlatt/operators/stencil.py
@jit
def hessian_squared(field: Array, dx: float) -> Array:
    r"""Pointwise squared Frobenius norm of the Hessian matrix.

    Computes

    $$\sum_{i,j} (\partial_i \partial_j \phi)^2$$

    using second-order centred finite differences for diagonal terms and
    forward-then-backward (mixed) differences for off-diagonal terms.
    Exploits Hessian symmetry — each off-diagonal pair is computed once and
    weighted by 2. Periodic boundary conditions throughout.

    Arises in higher-derivative scalar field equations, where a combination
    such as :math:`(\Box\phi)^2 - (\nabla_\mu\nabla_\nu\phi)^2` cancels the
    second-time-derivative piece and leaves this purely spatial contraction
    behind.

    Args:
        field: Real scalar field array (1-D, 2-D, or 3-D).
        dx: Lattice spacing (uniform in all directions).

    Returns:
        Array of the same shape as ``field`` with the squared Hessian norm
        at each site.
    """
    ndim = field.ndim
    result = jnp.zeros_like(field)

    # Diagonal terms: (∂_i² φ)²  using 3-point centred stencil.
    for i in range(ndim):
        d2_ii = (jnp.roll(field, -1, axis=i) + jnp.roll(field, 1, axis=i) - 2.0 * field) / (dx**2)
        result = result + d2_ii**2

    # Off-diagonal terms: 2 × (∂_i ∂_j φ)²  for j > i.
    # Mixed derivative via forward-then-backward difference:
    #   ∂_i ∂_j φ ≈ [φ(n+î+ĵ) - φ(n+î-ĵ) - φ(n-î+ĵ) + φ(n-î-ĵ)] / (4 dx²)
    for i in range(ndim):
        for j in range(i + 1, ndim):
            d2_ij = (
                jnp.roll(jnp.roll(field, -1, axis=i), -1, axis=j)
                - jnp.roll(jnp.roll(field, -1, axis=i), 1, axis=j)
                - jnp.roll(jnp.roll(field, 1, axis=i), -1, axis=j)
                + jnp.roll(jnp.roll(field, 1, axis=i), 1, axis=j)
            ) / (4.0 * dx**2)
            result = result + 2.0 * d2_ij**2

    return result