diff --git a/src/dolfinx_adjoint/__init__.py b/src/dolfinx_adjoint/__init__.py index 7a0a688..3ed5e7c 100644 --- a/src/dolfinx_adjoint/__init__.py +++ b/src/dolfinx_adjoint/__init__.py @@ -6,6 +6,7 @@ import pyadjoint as _pyad from .assembly import assemble_scalar, error_norm +from .interpolation import interpolate from .solvers import LinearProblem, NonlinearProblem from .types import Constant, Function, dirichletbc from .types.function import assign @@ -35,4 +36,5 @@ "__license__", "__email__", "__program_name__", + "interpolate", ] diff --git a/src/dolfinx_adjoint/blocks/interpolation.py b/src/dolfinx_adjoint/blocks/interpolation.py new file mode 100644 index 0000000..8cf8b58 --- /dev/null +++ b/src/dolfinx_adjoint/blocks/interpolation.py @@ -0,0 +1,202 @@ +from __future__ import annotations + +import typing +from typing import Callable + +import dolfinx +from pyadjoint import Block +from pyadjoint.overloaded_type import create_overloaded_object + +if typing.TYPE_CHECKING: + from petsc4py import PETSc + +# Global cache to prevent redundant matrix assembly across multiple blocks/iterations +_INTERPOLATION_MATRIX_CACHE: dict[tuple[int, int], dolfinx.la.MatrixCSR] = {} + + +def attach_working_array(mat: dolfinx.la.MatrixCSR): + """Attach working arrays to a dolfinx.la.MatrixCSR for efficient matrix-vector multiplication.""" + if not hasattr(mat, "_row_vec"): + mat._row_vec = dolfinx.la.vector(mat.index_map(0), mat.block_size[0], dtype=mat.data.dtype) + if not hasattr(mat, "_col_vec"): + mat._col_vec = dolfinx.la.vector(mat.index_map(1), mat.block_size[1], dtype=mat.data.dtype) + mat._row_vec.array[:] = 0.0 + mat._col_vec.array[:] = 0.0 + + +def get_mult( + mat: "PETSc.Mat" | dolfinx.la.MatrixCSR, + transpose: bool = False, # type: ignore +) -> Callable[[dolfinx.la.Vector, dolfinx.la.Vector], None]: + """Return a function that performs matrix-vector multiplication with the given matrix.""" + if isinstance(mat, dolfinx.la.MatrixCSR): + + def mult(v_in: dolfinx.la.Vector, v_out: dolfinx.la.Vector): + # Need to use vectors from + in_size_local = v_in.index_map.size_local * v_in.block_size + out_size_local = v_out.index_map.size_local * v_out.block_size + attach_working_array(mat) # Ensure working arrays are attached + if transpose: + # Prevent double-counting in parallel by zeroing ghosts of input vector + # Calculate the exact number of local degrees of freedom + mat._row_vec.array[:in_size_local] = v_in.array[:in_size_local] + mat._row_vec.scatter_forward() # Ensure ghost values are updated before multiplication + mat._col_vec.array[:out_size_local] = 0.0 + mat.mult(mat._row_vec, mat._col_vec, transpose=True) + v_out.array[:out_size_local] = mat._col_vec.array[:out_size_local] + else: + # Prevent double-counting in parallel by zeroing ghosts of input vector + # Calculate the exact number of local degrees of freedom + mat._row_vec.array[:out_size_local] = 0 + mat._col_vec.array[:in_size_local] = v_in.array[:in_size_local] + mat._col_vec.scatter_forward() # Ensure ghost values are updated before multiplication + mat.mult(mat._col_vec, mat._row_vec) + v_out.array[:out_size_local] = mat._row_vec.array[:out_size_local] + v_out.scatter_forward() + + return mult + elif dolfinx.has_petsc4py and dolfinx.has_petsc: + from petsc4py import PETSc + + if isinstance(mat, PETSc.Mat): + + def mult(v_in: dolfinx.la.Vector, v_out: dolfinx.la.Vector): + if transpose: + mat.multTranspose(v_in.petsc_vec, v_out.petsc_vec) + else: + mat.mult(v_in.petsc_vec, v_out.petsc_vec) + v_out.scatter_forward() + + return mult + else: + raise TypeError("Expected a PETSc.Mat when PETSc is available.") + else: + raise TypeError("Matrix type not supported. Expected dolfinx.la.MatrixCSR or PETSc.Mat, got {type(mat)=}.") + + +def _get_interpolation_matrix( + space_from: dolfinx.fem.FunctionSpace, space_to: dolfinx.fem.FunctionSpace, use_petsc: bool = False +) -> dolfinx.la.MatrixCSR | "PETSc.Mat": + """Retrieve or compute the interpolation matrix for a pair of spaces.""" + key = (id(space_from), id(space_to)) + + if key not in _INTERPOLATION_MATRIX_CACHE: + if use_petsc: + mat = dolfinx.fem.petsc.interpolation_matrix(space_from, space_to) + mat.assemble() + else: + mat = dolfinx.fem.interpolation_matrix(space_from, space_to) + mat.scatter_reverse() + # The built in interpolation matrix requires two working arrays + attach_working_array(mat) + + _INTERPOLATION_MATRIX_CACHE[key] = mat + + return _INTERPOLATION_MATRIX_CACHE[key] + + +class InterpolationBlock(Block): + """Block for interpolating a dolfinx.fem.Function into another space.""" + + def __init__( + self, + func_from: dolfinx.fem.Function, + func_to: dolfinx.fem.Function, + ad_block_tag: str | None = None, + petsc_mat: bool = False, + ): + super().__init__(ad_block_tag=ad_block_tag) + self.space_from = func_from.function_space + self.space_to = func_to.function_space + self._use_petsc = petsc_mat + + self.add_dependency(func_from) + + # Initialize internal caches for outputs to avoid MPI communicator exhaustion + self._adj_output: dolfinx.fem.Function | None = None + self._tlm_output: dolfinx.fem.Function | None = None + self._hessian_output: dolfinx.fem.Function | None = None + self._recompute_output: dolfinx.fem.Function | None = None + + def __str__(self): + return "interpolate_function" + + # --- Adjoint --- + + def prepare_evaluate_adj(self, inputs, adj_inputs, relevant_dependencies): + return _get_interpolation_matrix(self.space_from, self.space_to, use_petsc=self._use_petsc) + + def evaluate_adj_component(self, inputs, adj_inputs, block_variable, idx, prepared=None): + adj_input = adj_inputs[0] + mat = prepared + + if self._adj_output is None: + self._adj_output = dolfinx.fem.Function(self.space_from) + + # Action of the adjoint: A^T * adj_input + self._adj_output.x.array[:] = 0.0 # Reset the output vector before accumulation + adj_input.x.scatter_forward() + mult = get_mult(mat, transpose=True) + mult(adj_input.x, self._adj_output.x) + return self._adj_output + + # --- Tangent Linear Model (TLM) --- + + def prepare_evaluate_tlm(self, inputs, tlm_inputs, relevant_outputs): + return _get_interpolation_matrix(self.space_from, self.space_to, use_petsc=self._use_petsc) + + def evaluate_tlm_component(self, inputs, tlm_inputs, block_variable, idx, prepared=None): + tlm_input = tlm_inputs[0] + if tlm_input is None: + return None + + mat = prepared + + if self._tlm_output is None: + self._tlm_output = dolfinx.fem.Function(self.space_to) + + # Forward Jacobian action: A * tlm_input + tlm_input.x.scatter_forward() + self._tlm_output.x.array[:] = 0.0 # Reset the output vector before accumulation + mult = get_mult(mat, transpose=False) + mult(tlm_input.x, self._tlm_output.x) + return self._tlm_output + + # --- Hessian --- + + def prepare_evaluate_hessian(self, inputs, hessian_inputs, adj_inputs, relevant_dependencies): + return _get_interpolation_matrix(self.space_from, self.space_to, use_petsc=self._use_petsc) + + def evaluate_hessian_component( + self, inputs, hessian_inputs, adj_inputs, block_variable, idx, relevant_dependencies, prepared=None + ): + hessian_input = hessian_inputs[0] + mat = prepared + + if self._hessian_output is None: + self._hessian_output = dolfinx.fem.Function(self.space_from) + + # Action of the adjoint on the incoming Hessian sensitivity + hessian_input.x.scatter_forward() + self._hessian_output.x.array[:] = 0.0 # Reset the output vector before accumulation + mult = get_mult(mat, transpose=True) + mult(hessian_input.x, self._hessian_output.x) + return self._hessian_output + + # --- Recompute (Forward Pass) --- + + def prepare_recompute_component(self, inputs, relevant_outputs): + return None + + def recompute_component(self, inputs, block_variable, idx, prepared): + func_from = inputs[0] + + if self._recompute_output is None: + self._recompute_output = dolfinx.fem.Function(self.space_to) + + self._recompute_output.interpolate(func_from) + self._recompute_output.x.scatter_forward() + + # Overload the object to ensure PyAdjoint tracks it properly + output = create_overloaded_object(self._recompute_output) + return output diff --git a/src/dolfinx_adjoint/interpolation.py b/src/dolfinx_adjoint/interpolation.py new file mode 100644 index 0000000..b593897 --- /dev/null +++ b/src/dolfinx_adjoint/interpolation.py @@ -0,0 +1,35 @@ +import dolfinx +from pyadjoint.overloaded_type import create_overloaded_object +from pyadjoint.tape import annotate_tape, get_working_tape, stop_annotating + +from .blocks.interpolation import InterpolationBlock + + +def interpolate(u: dolfinx.fem.Function, V: dolfinx.fem.FunctionSpace, **kwargs): + """Interpolate a function to a different function space. + + Args: + u: The function to interpolate. + V: The function space to interpolate to. + kwargs: Keyword arguments to pass to the interpolation routine. + Includes ``"ad_block_tag"`` to tag the block in the adjoint tape, + ``"annotate"`` to control whether the assembly is annotated in the adjoint tape. + If you want to use PETSc based interpolation matrices, you can passe `petsc_mat=True` in the kwargs. + """ + ad_block_tag = kwargs.pop("ad_block_tag", None) + petsc_mat = kwargs.pop("petsc_mat", False) + annotate = annotate_tape(kwargs) + with stop_annotating(): + v = dolfinx.fem.Function(V) + v.interpolate(u) + output = create_overloaded_object(v) + + if annotate: + block = InterpolationBlock(u, output, ad_block_tag=ad_block_tag, petsc_mat=petsc_mat) + + tape = get_working_tape() + tape.add_block(block) + + block.add_output(output.block_variable) + + return output diff --git a/tests/test_interpolate.py b/tests/test_interpolate.py new file mode 100644 index 0000000..1328412 --- /dev/null +++ b/tests/test_interpolate.py @@ -0,0 +1,150 @@ +from mpi4py import MPI + +import dolfinx +import numpy as np +import pyadjoint +import pytest +import ufl + +from dolfinx_adjoint import Function, assemble_scalar, interpolate +from dolfinx_adjoint.blocks.interpolation import InterpolationBlock + +# Dynamically determine available matrix backends +petsc_options = [False] +if getattr(dolfinx, "has_petsc", False) and getattr(dolfinx, "has_petsc4py", False): + petsc_options.append(True) + + +@pytest.fixture(scope="module") +def mesh_1D(): + return dolfinx.mesh.create_unit_interval(MPI.COMM_WORLD, 10) + + +@pytest.fixture(scope="module") +def mesh_2D(): + return dolfinx.mesh.create_unit_square(MPI.COMM_WORLD, 7, 7) + + +@pytest.fixture(scope="module") +def mesh_3D(): + return dolfinx.mesh.create_unit_cube(MPI.COMM_WORLD, 5, 5, 5) + + +# ============================================================================== +# Test 1: Algebraic Adjoint Property ( == ) +# ============================================================================== + + +@pytest.mark.parametrize( + "family, degree_from, degree_to", + [ + ("Lagrange", 1, 2), + ("DG", 0, 1), + ("N1curl", 1, 2), + ], +) +@pytest.mark.parametrize("use_petsc", petsc_options) +def test_interpolation_block_adjoint_property(mesh_3D, family, degree_from, degree_to, use_petsc): + """Verifies that the TLM and Adjoint exactly satisfy the linear adjoint identity.""" + mesh = mesh_3D + + V_from = dolfinx.fem.functionspace(mesh, (family, degree_from)) + V_to = dolfinx.fem.functionspace(mesh, (family, degree_to)) + + # Use modern NumPy random generator with a fixed seed for reproducibility + rng = np.random.default_rng(seed=42) + + u = Function(V_from) + u.x.array[:] = rng.random(len(u.x.array)) + + v = Function(V_to) + v.x.array[:] = rng.random(len(v.x.array)) + + # Initialize the block directly, passing the PETSc backend flag + block = InterpolationBlock(u, v, petsc_mat=use_petsc) + + mat_tlm = block.prepare_evaluate_tlm([u], [u], None) + mat_adj = block.prepare_evaluate_adj([u], [v], None) + + tlm_output = block.evaluate_tlm_component(inputs=[u], tlm_inputs=[u], block_variable=None, idx=0, prepared=mat_tlm) + adj_output = block.evaluate_adj_component(inputs=[u], adj_inputs=[v], block_variable=None, idx=0, prepared=mat_adj) + + inner_forward = dolfinx.cpp.la.inner_product(tlm_output.x._cpp_object, v.x._cpp_object) + inner_adjoint = dolfinx.cpp.la.inner_product(u.x._cpp_object, adj_output.x._cpp_object) + + comm = mesh.comm + global_inner_forward = comm.allreduce(inner_forward, op=MPI.SUM) + global_inner_adjoint = comm.allreduce(inner_adjoint, op=MPI.SUM) + + np.testing.assert_allclose( + global_inner_forward, + global_inner_adjoint, + rtol=1e-12, + atol=1e-12, + err_msg=f"Adjoint property failed (PETSc={use_petsc}): != ", + ) + + +# ============================================================================== +# Test 2: Taylor Remainder Convergence (Graph & Optimization Integration) +# ============================================================================== + + +@pytest.mark.parametrize("mesh_var_name", ["mesh_1D", "mesh_2D", "mesh_3D"]) +@pytest.mark.parametrize("use_petsc", petsc_options) +def test_interpolation_taylor_test(mesh_var_name: str, request, use_petsc): + """Verifies that the exposed interpolate function works with ReducedFunctional.""" + pyadjoint.get_working_tape().clear_tape() + mesh = request.getfixturevalue(mesh_var_name) + + V_from = dolfinx.fem.functionspace(mesh, ("Lagrange", 1)) + V_to = dolfinx.fem.functionspace(mesh, ("Lagrange", 2)) + + u = Function(V_from) + u.name = "u_control" + u.x.array[:] = 0.2 + + # Forward the backend flag to the exposed wrapper + v = interpolate(u, V_to, petsc_mat=use_petsc) + + def u_ex(mod, x_coords): + return x_coords[0] + + x = ufl.SpatialCoordinate(mesh) + c = u_ex(ufl, x) + error = ufl.inner(v - c, v - c) * ufl.inner(v - c, v - c) * ufl.dx(domain=mesh) + + J = assemble_scalar(error) + + derivative_options = { + "riesz_representation": "L2", + "petsc_options": {"ksp_type": "preonly", "pc_type": "lu", "pc_factor_mat_solver_type": "mumps"}, + } + + control = pyadjoint.Control(u, riesz_map=derivative_options) + Jh = pyadjoint.ReducedFunctional(J, control) + assert Jh(u) > 0 + + du = Function(V_from) + du.interpolate(lambda x_coords: np.sin(x_coords[0])) + + # --- 1. Zero-order Taylor test --- + Jh(u) + min_rate = pyadjoint.taylor_test(Jh, u, du, dJdm=0) + assert np.isclose(min_rate, 1.0, rtol=1e-2, atol=1e-2) + + # --- 2. First-order Taylor test --- + Jh(u) + min_rate = pyadjoint.taylor_test(Jh, u, du) + assert np.isclose(min_rate, 2.0, rtol=1e-2, atol=1e-2) + + # --- 3. Second-order Taylor test --- + Jh(u) + dJdm = Jh.derivative()._ad_dot(du) + hessian = Jh.hessian(du) + dHddu = hessian._ad_dot(du) + + min_rate = pyadjoint.taylor_test(Jh, u, du, dJdm=dJdm, Hm=dHddu) + assert np.isclose(min_rate, 3.0, rtol=1e-3, atol=1e-3) + + pyadjoint.get_working_tape().clear_tape()