diff --git a/firedrake/adjoint_utils/blocks/__init__.py b/firedrake/adjoint_utils/blocks/__init__.py index fa10ff29fa..822aece800 100644 --- a/firedrake/adjoint_utils/blocks/__init__.py +++ b/firedrake/adjoint_utils/blocks/__init__.py @@ -1,8 +1,6 @@ from firedrake.adjoint_utils.blocks.assembly import AssembleBlock # noqa F401 from firedrake.adjoint_utils.blocks.solving import ( # noqa F401 - CachedSolverBlock, GenericSolveBlock, - ProjectBlock, SupermeshProjectBlock, SolveVarFormBlock, - NonlinearVariationalSolveBlock + CachedSolverBlock, SupermeshProjectBlock ) from firedrake.adjoint_utils.blocks.function import ( # noqa F401 FunctionAssignBlock, FunctionMergeBlock, SubfunctionBlock diff --git a/firedrake/adjoint_utils/blocks/solving.py b/firedrake/adjoint_utils/blocks/solving.py index a2d8582aa6..12c176a8e6 100644 --- a/firedrake/adjoint_utils/blocks/solving.py +++ b/firedrake/adjoint_utils/blocks/solving.py @@ -1,12 +1,6 @@ -import numpy -import ufl -from ufl.domain import extract_domains, extract_unique_domain -from ufl import replace -from ufl.formatting.ufl2unicode import ufl2unicode from enum import Enum -from pyadjoint import Block, stop_annotating, get_working_tape -from pyadjoint.enlisting import Enlist +from pyadjoint import Block, stop_annotating import firedrake from firedrake.adjoint_utils.checkpointing import maybe_disk_checkpoint @@ -311,529 +305,6 @@ def evaluate_hessian_component(self, inputs, hessian_inputs, adj_inputs, block_v return hessian_output -class GenericSolveBlock(Block): - pop_kwargs_keys = ["adj_cb", "adj_bdy_cb", "adj2_cb", "adj2_bdy_cb", - "forward_args", "forward_kwargs", "adj_args", - "adj_kwargs"] - - def __init__(self, lhs, rhs, func, bcs, *args, **kwargs): - super().__init__(ad_block_tag=kwargs.pop('ad_block_tag', None)) - self.adj_cb = kwargs.pop("adj_cb", None) - self.adj_bdy_cb = kwargs.pop("adj_bdy_cb", None) - self.adj2_cb = kwargs.pop("adj2_cb", None) - self.adj2_bdy_cb = kwargs.pop("adj2_bdy_cb", None) - - self.forward_args = [] - self.forward_kwargs = {} - self.adj_args = [] - self.adj_kwargs = {} - self.assemble_kwargs = {} - - # Equation LHS - self.lhs = lhs - # Equation RHS - self.rhs = rhs - # Solution function - self.func = func - self.function_space = self.func.function_space() - # Storage for adjoint solution of this block - self.adj_state_buf = func.copy(deepcopy=True) - # Boundary conditions - self.bcs = [] - if bcs is not None: - self.bcs = Enlist(bcs) - - if isinstance(self.lhs, ufl.Form) and isinstance(self.rhs, (ufl.Form, ufl.Cofunction)): - self.linear = True - for c in self.rhs.coefficients(): - self.add_dependency(c, no_duplicates=True) - else: - self.linear = False - - for c in self.lhs.coefficients(): - self.add_dependency(c, no_duplicates=True) - - for bc in self.bcs: - self.add_dependency(bc, no_duplicates=True) - - try: # add all meshes as dependency - for mesh in extract_domains(self.lhs): - self.add_dependency(mesh, no_duplicates=True) - except AttributeError: - pass - - if isinstance(self.rhs, ufl.BaseForm): - # add all meshes as dependency - for mesh in extract_domains(self.rhs): - self.add_dependency(mesh, no_duplicates=True) - - self._init_solver_parameters(args, kwargs) - - def _init_solver_parameters(self, args, kwargs): - self.forward_args = kwargs.pop("forward_args", []) - self.forward_kwargs = kwargs.pop("forward_kwargs", {}) - self.adj_args = kwargs.pop("adj_args", []) - self.adj_kwargs = kwargs.pop("adj_kwargs", {}) - self.assemble_kwargs = {} - - def __str__(self): - try: - lhs_string = ufl2unicode(self.lhs) - except AttributeError: - lhs_string = str(self.lhs) - try: - rhs_string = ufl2unicode(self.rhs) - except AttributeError: - rhs_string = str(self.rhs) - return "solve({} = {})".format(lhs_string, rhs_string) - - def _create_F_form(self): - # Process the equation forms, replacing values with checkpoints, - # and gathering lhs and rhs in one single form. - if self.linear: - tmp_u = firedrake.Function(self.function_space) - F_form = firedrake.action(self.lhs, tmp_u) - self.rhs - else: - tmp_u = self.func - F_form = self.lhs - - replace_map = self._replace_map(F_form) - replace_map[tmp_u] = self.get_outputs()[0].saved_output - return ufl.replace(F_form, replace_map) - - def _homogenize_bcs(self): - bcs = [] - for bc in self.bcs: - if isinstance(bc, firedrake.DirichletBC): - bc = bc.reconstruct(g=0) - bcs.append(bc) - return bcs - - def _create_initial_guess(self): - return firedrake.Function(self.function_space) - - def _recover_bcs(self): - bcs = [] - for block_variable in self.get_dependencies(): - c = block_variable.output - c_rep = block_variable.saved_output - - if isinstance(c, firedrake.DirichletBC): - bcs.append(c_rep) - return bcs - - def _replace_map(self, form): - replace_coeffs = {} - for block_variable in self.get_dependencies(): - coeff = block_variable.output - if coeff in form.coefficients(): - replace_coeffs[coeff] = block_variable.saved_output - return replace_coeffs - - def _replace_form(self, form, func=None): - """Replace the form coefficients with checkpointed values - - func represents the initial guess if relevant. - """ - replace_map = self._replace_map(form) - if func is not None and self.func in replace_map: - firedrake.Function.assign(func, replace_map[self.func]) - replace_map[self.func] = func - return ufl.replace(form, replace_map) - - def _should_compute_boundary_adjoint(self, relevant_dependencies): - # Check if DirichletBC derivative is relevant - bdy = False - for _, dep in relevant_dependencies: - if isinstance(dep.output, firedrake.DirichletBC): - bdy = True - break - return bdy - - @property - def adj_sol(self): - return self.adj_state - - @adj_sol.setter - def adj_sol(self, value): - if self.adj_state is None: - self.adj_state = value.copy(deepcopy=True) - else: - self.adj_state.assign(value) - - def prepare_evaluate_adj(self, inputs, adj_inputs, relevant_dependencies): - fwd_block_variable = self.get_outputs()[0] - u = fwd_block_variable.output - - dJdu = adj_inputs[0] - - F_form = self._create_F_form() - - dFdu = firedrake.derivative( - F_form, - fwd_block_variable.saved_output, - firedrake.TrialFunction( - u.function_space() - ) - ) - dFdu_form = firedrake.adjoint(dFdu) - dJdu = dJdu.copy() - - compute_bdy = self._should_compute_boundary_adjoint( - relevant_dependencies - ) - adj_sol, adj_sol_bdy = self._assemble_and_solve_adj_eq( - dFdu_form, dJdu, compute_bdy - ) - self.adj_sol = adj_sol - self.adj_state = adj_sol - self.adj_state_buf.assign(adj_sol) - if self.adj_cb is not None: - self.adj_cb(adj_sol) - if self.adj_bdy_cb is not None and compute_bdy: - self.adj_bdy_cb(adj_sol_bdy) - - r = {} - r["form"] = F_form - r["adj_sol"] = adj_sol - r["adj_sol_bdy"] = adj_sol_bdy - return r - - def evaluate_adj_component(self, inputs, adj_inputs, block_variable, idx, - prepared=None): - if not self.linear and self.func == block_variable.output: - # We are not able to calculate derivatives wrt initial guess. - return None - F_form = prepared["form"] - adj_sol = prepared["adj_sol"] - adj_sol_bdy = prepared["adj_sol_bdy"] - c = block_variable.output - c_rep = block_variable.saved_output - - if isinstance(c, (firedrake.Function, firedrake.Cofunction)): - trial_function = firedrake.TrialFunction(c.function_space()) - elif isinstance(c, firedrake.DirichletBC): - tmp_bc = c.reconstruct( - g=extract_subfunction(adj_sol_bdy, c.function_space()) - ) - return [tmp_bc] - elif isinstance(c, firedrake.MeshGeometry): - # Using CoordinateDerivative requires us to do action before - # differentiating, might change in the future. - F_form_tmp = firedrake.action(F_form, adj_sol) - X = firedrake.SpatialCoordinate(c_rep) - dFdm = firedrake.derivative( - -F_form_tmp, X, - firedrake.TestFunction(c._ad_function_space()) - ) - - if dFdm == 0: - return firedrake.Function(c._ad_function_space().dual()) - - dFdm = firedrake.assemble(dFdm, **self.assemble_kwargs) - return dFdm - - dFdm = -firedrake.derivative(F_form, c_rep, trial_function) - if isinstance(dFdm, ufl.Form): - dFdm = firedrake.adjoint(dFdm) - dFdm = firedrake.action(dFdm, adj_sol) - else: - dFdm = dFdm(adj_sol) - dFdm = firedrake.assemble(dFdm, **self.assemble_kwargs) - return dFdm - - def _assemble_dFdu_adj(self, dFdu_adj_form, **kwargs): - return firedrake.assemble(dFdu_adj_form, **kwargs) - - def _assemble_and_solve_adj_eq(self, dFdu_adj_form, dJdu, compute_bdy): - dJdu_copy = dJdu.copy() - # Homogenize and apply boundary conditions on adj_dFdu. - bcs = self._homogenize_bcs() - dFdu = firedrake.assemble(dFdu_adj_form, bcs=bcs, **self.assemble_kwargs) - - adj_sol = firedrake.Function(self.function_space) - firedrake.solve( - dFdu, adj_sol, dJdu, *self.adj_args, **self.adj_kwargs - ) - - adj_sol_bdy = None - if compute_bdy: - adj_sol_bdy = self._compute_adj_bdy( - adj_sol, adj_sol_bdy, dFdu_adj_form, dJdu_copy) - - return adj_sol, adj_sol_bdy - - def _compute_adj_bdy(self, adj_sol, adj_sol_bdy, dFdu_adj_form, dJdu): - adj_sol_bdy = firedrake.assemble(dJdu - firedrake.action(dFdu_adj_form, adj_sol)) - return adj_sol_bdy.riesz_representation("l2") - - def prepare_evaluate_tlm(self, inputs, tlm_inputs, relevant_outputs): - pass - - def evaluate_tlm_component(self, inputs, tlm_inputs, block_variable, idx, - prepared=None): - fwd_block_variable = self.get_outputs()[0] - u = fwd_block_variable.output - - F_form = self._create_F_form() - - # Obtain dFdu. - dFdu = firedrake.derivative( - F_form, - fwd_block_variable.saved_output, - firedrake.TrialFunction(u.function_space()) - ) - V = self.get_outputs()[idx].output.function_space() - - bcs = [] - dFdm = 0. - for block_variable in self.get_dependencies(): - tlm_value = block_variable.tlm_value - c = block_variable.output - c_rep = block_variable.saved_output - - if isinstance(c, firedrake.DirichletBC): - if tlm_value is None: - bcs.append(c.reconstruct(g=0)) - else: - bcs.append(tlm_value) - continue - elif isinstance(c, firedrake.MeshGeometry): - X = firedrake.SpatialCoordinate(c) - c_rep = X - - if tlm_value is None: - continue - - if c == self.func and not self.linear: - continue - - dFdm += firedrake.derivative(-F_form, c_rep, tlm_value) - - if isinstance(dFdm, float): - v = dFdu.arguments()[0] - dFdm = firedrake.inner( - firedrake.Constant(numpy.zeros(v.ufl_shape)), v - ) * firedrake.dx - - dFdm = ufl.algorithms.expand_derivatives(dFdm) - dFdm = firedrake.assemble(dFdm) - dudm = firedrake.Function(V) - result = self._assemble_and_solve_tlm_eq( - firedrake.assemble(dFdu, bcs=bcs, **self.assemble_kwargs), - dFdm, dudm, bcs - ) - return result - - def _assemble_and_solve_tlm_eq(self, dFdu, dFdm, dudm, bcs): - return self._assembled_solve(dFdu, dFdm, dudm, bcs) - - def _assemble_soa_eq_rhs(self, dFdu_form, adj_sol, hessian_input, d2Fdu2): - # Start piecing together the rhs of the soa equation - b = hessian_input.copy() - if len(d2Fdu2.integrals()) > 0: - b_form = firedrake.action(firedrake.adjoint(d2Fdu2), adj_sol) - else: - b_form = d2Fdu2 - - for bo in self.get_dependencies(): - c = bo.output - c_rep = bo.saved_output - tlm_input = bo.tlm_value - - if (c == self.func and not self.linear) or tlm_input is None: - continue - - if isinstance(c, firedrake.MeshGeometry): - X = firedrake.SpatialCoordinate(c) - dFdu_adj = firedrake.action(firedrake.adjoint(dFdu_form), - adj_sol) - d2Fdudm = ufl.algorithms.expand_derivatives( - firedrake.derivative(dFdu_adj, X, tlm_input)) - if len(d2Fdudm.integrals()) > 0: - b_form += d2Fdudm - elif not isinstance(c, firedrake.DirichletBC): - dFdu_adj = firedrake.action(firedrake.adjoint(dFdu_form), - adj_sol) - # b_form += firedrake.derivative(dFdu_adj, c_rep, tlm_input) - bo_form = ufl.algorithms.expand_derivatives( - firedrake.derivative(dFdu_adj, c_rep, tlm_input)) - b_form += bo_form - - b_form = ufl.algorithms.expand_derivatives(b_form) - if len(b_form.integrals()) > 0: - b -= firedrake.assemble(b_form) - - return b - - def _assemble_and_solve_soa_eq(self, dFdu_form, adj_sol, hessian_input, - d2Fdu2, compute_bdy): - b = self._assemble_soa_eq_rhs(dFdu_form, adj_sol, hessian_input, - d2Fdu2) - dFdu_form = firedrake.adjoint(dFdu_form) - adj_sol2, adj_sol2_bdy = self._assemble_and_solve_adj_eq(dFdu_form, b, - compute_bdy) - if self.adj2_cb is not None: - self.adj2_cb(adj_sol2) - if self.adj2_bdy_cb is not None and compute_bdy: - self.adj2_bdy_cb(adj_sol2_bdy) - return adj_sol2, adj_sol2_bdy - - def prepare_evaluate_hessian(self, inputs, hessian_inputs, adj_inputs, - relevant_dependencies): - # First fetch all relevant values - fwd_block_variable = self.get_outputs()[0] - hessian_input = hessian_inputs[0] - tlm_output = fwd_block_variable.tlm_value - - self.adj_state = self.adj_state_buf.copy(deepcopy=True) - - if hessian_input is None: - return - - if tlm_output is None: - return - - F_form = self._create_F_form() - - # Using the equation Form derive dF/du, d^2F/du^2 * du/dm * direction. - dFdu_form = firedrake.derivative(F_form, - fwd_block_variable.saved_output) - d2Fdu2 = ufl.algorithms.expand_derivatives( - firedrake.derivative(dFdu_form, fwd_block_variable.saved_output, - tlm_output)) - - adj_sol = self.adj_sol - if adj_sol is None: - raise RuntimeError("Hessian computation was run before adjoint.") - bdy = self._should_compute_boundary_adjoint(relevant_dependencies) - adj_sol2, adj_sol2_bdy = self._assemble_and_solve_soa_eq( - dFdu_form, adj_sol, hessian_input, d2Fdu2, bdy - ) - - r = {} - r["adj_sol2"] = adj_sol2 - r["adj_sol2_bdy"] = adj_sol2_bdy - r["form"] = F_form - r["adj_sol"] = adj_sol - - return r - - def evaluate_hessian_component(self, inputs, hessian_inputs, adj_inputs, - block_variable, idx, relevant_dependencies, - prepared=None): - c = block_variable.output - if c == self.func and not self.linear: - return None - - adj_sol2 = prepared["adj_sol2"] - adj_sol2_bdy = prepared["adj_sol2_bdy"] - F_form = prepared["form"] - adj_sol = prepared["adj_sol"] - fwd_block_variable = self.get_outputs()[0] - tlm_output = fwd_block_variable.tlm_value - - c_rep = block_variable.saved_output - - # If m = DirichletBC then d^2F(u,m)/dm^2 = 0 and d^2F(u,m)/dudm = 0, - # so we only have the term dF(u,m)/dm * adj_sol2 - if isinstance(c, firedrake.DirichletBC): - tmp_bc = c.reconstruct( - g=extract_subfunction(adj_sol2_bdy, c.function_space()) - ) - return [tmp_bc] - - if isinstance(c, firedrake.MeshGeometry): - X = firedrake.SpatialCoordinate(c) - W = c._ad_function_space() - else: - W = c.function_space() - - dc = firedrake.TestFunction(W) - form_adj = firedrake.action(F_form, adj_sol) - form_adj2 = firedrake.action(F_form, adj_sol2) - if isinstance(c, firedrake.MeshGeometry): - dFdm_adj = firedrake.derivative(form_adj, X, dc) - dFdm_adj2 = firedrake.derivative(form_adj2, X, dc) - else: - dFdm_adj = firedrake.derivative(form_adj, c_rep, dc) - dFdm_adj2 = firedrake.derivative(form_adj2, c_rep, dc) - - # TODO: Old comment claims this might break on split. Confirm if true - # or not. - d2Fdudm = ufl.algorithms.expand_derivatives( - firedrake.derivative(dFdm_adj, fwd_block_variable.saved_output, - tlm_output)) - - d2Fdm2 = 0 - # We need to add terms from every other dependency - # i.e. the terms d^2F/dm_1dm_2 - for _, bv in relevant_dependencies: - c2 = bv.output - c2_rep = bv.saved_output - - if isinstance(c2, firedrake.DirichletBC): - continue - - tlm_input = bv.tlm_value - if tlm_input is None: - continue - - if c2 == self.func and not self.linear: - continue - - # TODO: If tlm_input is a Sum, this crashes in some instances? - if isinstance(c2_rep, firedrake.MeshGeometry): - X = firedrake.SpatialCoordinate(c2_rep) - d2Fdm2 += ufl.algorithms.expand_derivatives( - firedrake.derivative(dFdm_adj, X, tlm_input) - ) - else: - d2Fdm2 += ufl.algorithms.expand_derivatives( - firedrake.derivative(dFdm_adj, c2_rep, tlm_input) - ) - - hessian_form = ufl.algorithms.expand_derivatives( - d2Fdm2 + dFdm_adj2 + d2Fdudm - ) - hessian_output = 0 - if not hessian_form.empty(): - hessian_output = firedrake.assemble(hessian_form) - hessian_output *= -1. - - return hessian_output - - def prepare_recompute_component(self, inputs, relevant_outputs): - return self._replace_recompute_form() - - def _replace_recompute_form(self): - func = self._create_initial_guess() - - bcs = self._recover_bcs() - lhs = self._replace_form(self.lhs, func=func) - rhs = 0 - if self.linear: - rhs = self._replace_form(self.rhs) - - return lhs, rhs, func, bcs - - def _forward_solve(self, lhs, rhs, func, bcs): - firedrake.solve(lhs == rhs, func, bcs, *self.forward_args, - **self.forward_kwargs) - return func - - def _assembled_solve(self, lhs, rhs, func, bcs, **kwargs): - firedrake.solve(lhs, func, rhs, **kwargs) - return func - - def recompute_component(self, inputs, block_variable, idx, prepared): - lhs, rhs, func, bcs = prepared - result = self._forward_solve(lhs, rhs, func, bcs) - if isinstance(block_variable.checkpoint, firedrake.Function): - result = block_variable.checkpoint.assign(result) - return maybe_disk_checkpoint(result) - - def solve_init_params(self, args, kwargs, varform): if len(self.forward_args) <= 0: self.forward_args = args @@ -877,230 +348,6 @@ def solve_init_params(self, args, kwargs, varform): self.assemble_kwargs["appctx"] = kwargs["appctx"] -class SolveVarFormBlock(GenericSolveBlock): - def __init__(self, equation, func, bcs=[], *args, **kwargs): - lhs = equation.lhs - rhs = equation.rhs - super().__init__(lhs, rhs, func, bcs, *args, **kwargs) - - def _init_solver_parameters(self, args, kwargs): - super()._init_solver_parameters(args, kwargs) - solve_init_params(self, args, kwargs, varform=True) - - -class NonlinearVariationalSolveBlock(GenericSolveBlock): - def __init__(self, equation, func, bcs, adj_cache, problem_J, - solver_kwargs, **kwargs): - lhs = equation.lhs - rhs = equation.rhs - - self._adj_cache = adj_cache - self._dFdm_cache = adj_cache.setdefault("dFdm_cache", {}) - self.problem_J = problem_J - self.solver_kwargs = solver_kwargs - - self.adj_state_buf = func.copy(deepcopy=True) - - super().__init__(lhs, rhs, func, bcs, **{**solver_kwargs, **kwargs}) - - if self.problem_J is not None: - for coeff in self.problem_J.coefficients(): - self.add_dependency(coeff, no_duplicates=True) - - def _init_solver_parameters(self, args, kwargs): - super()._init_solver_parameters(args, kwargs) - solve_init_params(self, args, kwargs, varform=True) - - def recompute_component(self, inputs, block_variable, idx, prepared): - tape = get_working_tape() - if self._ad_solvers["recompute_count"] == tape.recompute_count - 1: - # Update how many times the block has been recomputed. - self._ad_solvers["recompute_count"] = tape.recompute_count - if self._ad_solvers["forward_nlvs"]._problem._constant_jacobian: - self._ad_solvers["forward_nlvs"].invalidate_jacobian() - self._ad_solvers["update_adjoint"] = True - return super().recompute_component(inputs, block_variable, idx, prepared) - - def _forward_solve(self, lhs, rhs, func, bcs, **kwargs): - self._ad_solver_replace_forms() - self._ad_solvers["forward_nlvs"].solve() - func.assign(self._ad_solvers["forward_nlvs"]._problem.u) - return func - - def _adjoint_solve(self, dJdu, compute_bdy): - dJdu_copy = dJdu.copy() - # Homogenize and apply boundary conditions on adj_dFdu and dJdu. - for bc in self.bcs: - bc.zero(dJdu) - - if ( - self._ad_solvers["forward_nlvs"]._problem._constant_jacobian - and self._ad_solvers["update_adjoint"] - ): - # Update left hand side of the adjoint equation. - self._ad_solver_replace_forms(SolverType.ADJOINT) - self._ad_solvers["adjoint_lvs"].invalidate_jacobian() - self._ad_solvers["update_adjoint"] = False - elif not self._ad_solvers["forward_nlvs"]._problem._constant_jacobian: - # Update left hand side of the adjoint equation. - self._ad_solver_replace_forms(SolverType.ADJOINT) - - # Update the right hand side of the adjoint equation. - # problem.F._component[1] is the right hand side of the adjoint. - self._ad_solvers["adjoint_lvs"]._problem.F._components[1].assign(dJdu) - - # Solve the adjoint linear variational solver. - self._ad_solvers["adjoint_lvs"].solve() - u_sol = self._ad_solvers["adjoint_lvs"]._problem.u - - adj_sol_bdy = None - if compute_bdy: - jac_adj = self._ad_solvers["adjoint_lvs"]._problem.J - adj_sol_bdy = self._compute_adj_bdy( - u_sol, adj_sol_bdy, jac_adj, dJdu_copy) - return u_sol, adj_sol_bdy - - def _ad_assign_map(self, form, solver): - if solver == SolverType.FORWARD: - count_map = self._ad_solvers["forward_nlvs"]._problem._ad_count_map - else: - count_map = self._ad_solvers["adjoint_lvs"]._problem._ad_count_map - assign_map = {} - form_ad_count_map = dict((count_map[coeff], coeff) - for coeff in form.coefficients()) - for block_variable in self.get_dependencies(): - coeff = block_variable.output - if isinstance(coeff, - (firedrake.Coefficient, firedrake.Constant, - firedrake.Cofunction)): - coeff_count = coeff.count() - if coeff_count in form_ad_count_map: - assign_map[form_ad_count_map[coeff_count]] = \ - block_variable.saved_output - - if ( - solver == SolverType.ADJOINT - and not self._ad_solvers["forward_nlvs"]._problem._constant_jacobian - ): - block_variable = self.get_outputs()[0] - coeff_count = block_variable.output.count() - if coeff_count in form_ad_count_map: - assign_map[form_ad_count_map[coeff_count]] = \ - block_variable.saved_output - return assign_map - - def _ad_assign_coefficients(self, form, solver): - assign_map = self._ad_assign_map(form, solver) - for coeff, value in assign_map.items(): - coeff.assign(value) - - def _ad_solver_replace_forms(self, solver=SolverType.FORWARD): - if solver == SolverType.FORWARD: - problem = self._ad_solvers["forward_nlvs"]._problem - self._ad_assign_coefficients(problem.F, solver) - self._ad_assign_coefficients(problem.J, solver) - else: - self._ad_assign_coefficients( - self._ad_solvers["adjoint_lvs"]._problem.J, solver) - - def prepare_evaluate_adj(self, inputs, adj_inputs, relevant_dependencies): - compute_bdy = self._should_compute_boundary_adjoint( - relevant_dependencies - ) - adj_sol, adj_sol_bdy = self._adjoint_solve(adj_inputs[0], compute_bdy) - self.adj_state = adj_sol - self.adj_state_buf.assign(adj_sol) - if self.adj_cb is not None: - self.adj_cb(adj_sol) - if self.adj_bdy_cb is not None and compute_bdy: - self.adj_bdy_cb(adj_sol_bdy) - - r = {} - r["form"] = self._create_F_form() - r["adj_sol"] = self.adj_state - r["adj_sol_bdy"] = adj_sol_bdy - return r - - def evaluate_adj_component(self, inputs, adj_inputs, block_variable, idx, - prepared=None): - if not self.linear and self.func == block_variable.output: - # We are not able to calculate derivatives wrt initial guess. - return None - F_form = prepared["form"] - adj_sol = prepared["adj_sol"] - adj_sol_bdy = prepared["adj_sol_bdy"] - c = block_variable.output - c_rep = block_variable.saved_output - - if isinstance(c, (firedrake.Function, firedrake.Cofunction)): - trial_function = firedrake.TrialFunction(c.function_space()) - elif isinstance(c, firedrake.Constant): - try: - mesh = extract_unique_domain(F_form) - except ValueError: - raise ValueError("Expecting a single mesh") - trial_function = firedrake.TrialFunction( - c._ad_function_space(mesh) - ) - elif isinstance(c, firedrake.DirichletBC): - tmp_bc = c.reconstruct( - g=extract_subfunction(adj_sol_bdy, c.function_space()) - ) - return [tmp_bc] - elif isinstance(c, firedrake.MeshGeometry): - # Using CoordianteDerivative requires us to do action before - # differentiating, might change in the future. - F_form_tmp = firedrake.action(F_form, adj_sol) - X = firedrake.SpatialCoordinate(c_rep) - dFdm = firedrake.derivative(-F_form_tmp, X, firedrake.TestFunction( - c._ad_function_space()) - ) - - dFdm = firedrake.assemble(dFdm, **self.assemble_kwargs) - return dFdm - - # dFdm_cache works with original variables, not block saved outputs. - if c in self._dFdm_cache: - dFdm = self._dFdm_cache[c] - else: - dFdm = -firedrake.derivative(self.lhs, c, trial_function) - dFdm = firedrake.adjoint(dFdm) - self._dFdm_cache[c] = dFdm - - # Replace the form coefficients with checkpointed values. - replace_map = self._replace_map(dFdm) - replace_map[self.func] = self.get_outputs()[0].saved_output - dFdm = replace(dFdm, replace_map) - - if isinstance(dFdm, firedrake.Argument): - # Corner case. Should be fixed more permanently upstream in UFL. - # See: https://github.com/FEniCS/ufl/issues/395 - dFdm = ufl.Action(dFdm, adj_sol) - else: - dFdm = dFdm * adj_sol - dFdm = firedrake.assemble(dFdm, **self.assemble_kwargs) - - return dFdm - - -class ProjectBlock(SolveVarFormBlock): - def __init__(self, v, V, output, bcs=[], *args, **kwargs): - mesh = kwargs.pop("mesh", None) - if mesh is None: - mesh = V.mesh() - dx = firedrake.dx(mesh) - w = firedrake.TestFunction(V) - Pv = firedrake.TrialFunction(V) - a = firedrake.inner(Pv, w) * dx - L = firedrake.inner(v, w) * dx - - super().__init__(a == L, output, bcs, *args, **kwargs) - - def _init_solver_parameters(self, args, kwargs): - super()._init_solver_parameters(args, kwargs) - solve_init_params(self, args, kwargs, varform=True) - - class SupermeshProjectBlock(Block): r""" Annotates supermesh projection. @@ -1124,6 +371,11 @@ class SupermeshProjectBlock(Block): Step 2. solve linear system. """ + + pop_kwargs_keys = ["adj_cb", "adj_bdy_cb", "adj2_cb", "adj2_bdy_cb", + "forward_args", "forward_kwargs", "adj_args", + "adj_kwargs"] + def __init__(self, source, target_space, target, bcs=[], **kwargs): super(SupermeshProjectBlock, self).__init__( ad_block_tag=kwargs.pop("ad_block_tag", None) diff --git a/firedrake/adjoint_utils/projection.py b/firedrake/adjoint_utils/projection.py index 75ef9a7b5c..7133ddc0a6 100644 --- a/firedrake/adjoint_utils/projection.py +++ b/firedrake/adjoint_utils/projection.py @@ -1,6 +1,6 @@ from functools import wraps from pyadjoint.tape import annotate_tape, stop_annotating, get_working_tape -from firedrake.adjoint_utils.blocks import ProjectBlock, SupermeshProjectBlock +from firedrake.adjoint_utils.blocks import SupermeshProjectBlock def annotate_project(project): @@ -35,7 +35,7 @@ def wrapper(self, **kwargs): V = self.target.function_space() if annotate: bcs = kwargs.get("bcs", []) - sb_kwargs = ProjectBlock.pop_kwargs(kwargs) + sb_kwargs = SupermeshProjectBlock.pop_kwargs(kwargs) if self._target_is_function: # block should be created before project because output might also be an input that needs checkpointing block = SupermeshProjectBlock(self.source, V, self.target, bcs, ad_block_tag=self.ad_block_tag, **sb_kwargs) diff --git a/firedrake/adjoint_utils/solving.py b/firedrake/adjoint_utils/solving.py index c7301288f4..767294684a 100644 --- a/firedrake/adjoint_utils/solving.py +++ b/firedrake/adjoint_utils/solving.py @@ -1,5 +1,5 @@ from pyadjoint.tape import get_working_tape -from firedrake.adjoint_utils.blocks import CachedSolverBlock, GenericSolveBlock, ProjectBlock +from firedrake.adjoint_utils.blocks import CachedSolverBlock def get_solve_blocks(): @@ -11,6 +11,6 @@ def get_solve_blocks(): return [ block for block in get_working_tape().get_blocks() - if issubclass(type(block), (CachedSolverBlock, GenericSolveBlock)) - and not issubclass(type(block), ProjectBlock) + if issubclass(type(block), CachedSolverBlock) + and not getattr(block, "_is_project", False) ] diff --git a/firedrake/adjoint_utils/variational_solver.py b/firedrake/adjoint_utils/variational_solver.py index 3c8f9ce49a..3cf15840fa 100644 --- a/firedrake/adjoint_utils/variational_solver.py +++ b/firedrake/adjoint_utils/variational_solver.py @@ -1,8 +1,7 @@ import copy from functools import wraps, cached_property from pyadjoint.tape import get_working_tape, stop_annotating, annotate_tape, no_annotations -from firedrake.adjoint_utils.blocks import ( - NonlinearVariationalSolveBlock, CachedSolverBlock) +from firedrake.adjoint_utils.blocks import CachedSolverBlock from firedrake.adjoint_utils.blocks.solving import solve_init_params from firedrake.ufl_expr import derivative, adjoint, action from ufl import replace, Action @@ -503,6 +502,9 @@ def wrapper(self, **kwargs): self._ad_hessian_cache, ad_block_tag=self.ad_block_tag) + if hasattr(self, "_is_project"): + block._is_project = self._is_project + for dep in self._ad_dependencies_to_add: block.add_dependency(dep, no_duplicates=True) # mesh = self._ad_problem.u.function_space().mesh() @@ -520,66 +522,6 @@ def wrapper(self, **kwargs): return wrapper - @staticmethod - def _ad_annotate_solve_old(solve): - @wraps(solve) - def wrapper(self, **kwargs): - """To disable the annotation, just pass :py:data:`annotate=False` to this routine, and it acts exactly like the - Firedrake solve call. This is useful in cases where the solve is known to be irrelevant or diagnostic - for the purposes of the adjoint computation (such as projecting fields to other function spaces - for the purposes of visualisation).""" - from firedrake import LinearVariationalSolver - annotate = annotate_tape(kwargs) - if annotate: - bounds = kwargs.pop("bounds", None) - if bounds is not None: - raise ValueError( - "MissingMathsError: we do not know how to differentiate through a variational inequality") - - tape = get_working_tape() - problem = self._ad_problem - sb_kwargs = NonlinearVariationalSolveBlock.pop_kwargs(kwargs) - sb_kwargs.update(kwargs) - - block = NonlinearVariationalSolveBlock(problem._ad_F == 0, - problem._ad_u, - problem._ad_bcs, - adj_cache=self._ad_adj_cache, - problem_J=problem._ad_J, - solver_kwargs=self._ad_kwargs, - ad_block_tag=self.ad_block_tag, - **sb_kwargs) - - # Forward variational solver. - if not self._ad_solvers["forward_nlvs"]: - self._ad_solvers["forward_nlvs"] = type(self)( - self._ad_problem_clone(self._ad_problem, block.get_dependencies()), - **self._ad_kwargs - ) - - # Adjoint variational solver. - if not self._ad_solvers["adjoint_lvs"]: - with stop_annotating(): - self._ad_solvers["adjoint_lvs"] = LinearVariationalSolver( - self._ad_adj_lvs_problem(block, problem._ad_adj_F), - *block.adj_args, **block.adj_kwargs) - if self._ad_problem._constant_jacobian: - self._ad_solvers["update_adjoint"] = False - - block._ad_solvers = self._ad_solvers - - tape.add_block(block) - - with stop_annotating(): - out = solve(self, **kwargs) - - if annotate: - block.add_output(self._ad_problem._ad_u.create_block_variable()) - - return out - - return wrapper - @no_annotations def _ad_problem_clone(self, problem, dependencies): """Replaces every coefficient in the residual and jacobian with a deepcopy to return diff --git a/firedrake/projection.py b/firedrake/projection.py index b40cec08d7..84e2101f86 100644 --- a/firedrake/projection.py +++ b/firedrake/projection.py @@ -278,6 +278,7 @@ def __init__(self, *args, **kwargs): @NonlinearVariationalSolverMixin._ad_annotate_init def _init_as_solver(self, problem, **kwargs): + self._is_project = True return @NonlinearVariationalSolverMixin._ad_annotate_solve