Skip to content
Draft
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
12 changes: 12 additions & 0 deletions firedrake/adjoint_utils/variational_solver.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import copy
from functools import wraps
import warnings
from pyadjoint.tape import get_working_tape, stop_annotating, annotate_tape, no_annotations
from firedrake.adjoint_utils.blocks import NonlinearVariationalSolveBlock
from firedrake.ufl_expr import derivative, adjoint
Expand Down Expand Up @@ -73,6 +74,17 @@ def wrapper(self, **kwargs):
tape = get_working_tape()
problem = self._ad_problem
sb_kwargs = NonlinearVariationalSolveBlock.pop_kwargs(kwargs)
if (
"adj_kwargs" in self._ad_kwargs
and "adj_kwargs" in sb_kwargs
):
raise TypeError("Cannot pass adj_kwargs to both "
"solver constructor and solve().")
if "adj_kwargs" in sb_kwargs:
warnings.warn(
"Passing adj_kwargs to the solve method has been"
" deprecated, pass it to the"
" NonlinearVariationalProblem instead.", FutureWarning)
sb_kwargs.update(kwargs)

block = NonlinearVariationalSolveBlock(problem._ad_F == 0,
Expand Down
73 changes: 40 additions & 33 deletions firedrake/variational_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,49 +194,53 @@ def __init__(self, problem, *, solver_parameters=None,
post_jacobian_callback=None,
pre_function_callback=None,
post_function_callback=None,
pre_apply_bcs=True):
pre_apply_bcs=True,
adj_kwargs=None):
r"""
:arg problem: A :class:`NonlinearVariationalProblem` to solve.
:kwarg nullspace: an optional :class:`.VectorSpaceBasis` (or
:class:`.MixedVectorSpaceBasis`) spanning the null
space of the operator.
:kwarg transpose_nullspace: as for the nullspace, but used to
make the right hand side consistent.
:kwarg near_nullspace: as for the nullspace, but used to
specify the near nullspace (for multigrid solvers).
:kwarg solver_parameters: Solver parameters to pass to PETSc.
This should be a dict mapping PETSc options to values.
:kwarg appctx: A dictionary containing application context that
is passed to the preconditioner if matrix-free.
:kwarg options_prefix: an optional prefix used to distinguish
PETSc options. If not provided a unique prefix will be
created. Use this option if you want to pass options
to the solver from the command line in addition to
through the ``solver_parameters`` dict.
:kwarg pre_jacobian_callback: A user-defined function that will
be called immediately before Jacobian assembly. This can
be used, for example, to update a coefficient function
that has a complicated dependence on the unknown solution.
:kwarg post_jacobian_callback: As above, but called after the
Jacobian has been assembled.
:kwarg pre_function_callback: As above, but called immediately
before residual assembly.
:kwarg post_function_callback: As above, but called immediately
after residual assembly.
:class:`.MixedVectorSpaceBasis`) spanning the null space of the
operator.
:kwarg transpose_nullspace: as for the nullspace, but used to make the
right hand side consistent.
:kwarg near_nullspace: as for the nullspace, but used to specify the
near nullspace (for multigrid solvers).
:kwarg solver_parameters: Solver parameters to pass to PETSc. This
should be a dict mapping PETSc options to values.
:kwarg appctx: A dictionary containing application context that is
passed to the preconditioner if matrix-free.
:kwarg options_prefix: an optional prefix used to distinguish PETSc
options. If not provided a unique prefix will be created. Use
this option if you want to pass options to the solver from the
command line in addition to through the ``solver_parameters``
dict.
:kwarg pre_jacobian_callback: A user-defined function that will be
called immediately before Jacobian assembly. This can be used,
for example, to update a coefficient function that has a
complicated dependence on the unknown solution.
:kwarg post_jacobian_callback: As above, but called after the Jacobian
has been assembled.
:kwarg pre_function_callback: As above, but called immediately before
residual assembly.
:kwarg post_function_callback: As above, but called immediately after
residual assembly.
:kwarg pre_apply_bcs: If `True`, the bcs are applied before the solve.
Otherwise, the problem is linearised around the initial guess
before imposing bcs, and the bcs are appended to the nonlinear system.
before imposing bcs, and the bcs are appended to the nonlinear
system.
:kwarg adj_kwargs: A dictionary of keyword arguments for the adjoint
LinearVariationalSolver. Arguments not set here default to the
values given for the forward solve.

Example usage of the ``solver_parameters`` option: to set the
nonlinear solver type to just use a linear solver, use
Example usage of the ``solver_parameters`` option: to set the nonlinear
solver type to just use a linear solver, use

.. code-block:: python3

{'snes_type': 'ksponly'}

PETSc flag options (where the presence of the option means something) should
be specified with ``None``.
For example:
PETSc flag options (where the presence of the option means something)
should be specified with ``None``. For example:

.. code-block:: python3

Expand All @@ -251,12 +255,15 @@ def __init__(self, problem, *, solver_parameters=None,
def update_diffusivity(current_solution):
with cursol.dat.vec_wo as v:
current_solution.copy(v)
solve(trial*test*dx == dot(grad(cursol), grad(test))*dx, diffusivity)
solve(trial*test*dx == dot(grad(cursol), grad(test))*dx,
diffusivity)

solver = NonlinearVariationalSolver(problem,
pre_jacobian_callback=update_diffusivity)

"""
# Note adj_kwargs is picked up by the wrapper it is only included here
# for documentation purposes.
assert isinstance(problem, NonlinearVariationalProblem)

solver_parameters = flatten_parameters(solver_parameters or {})
Expand Down
118 changes: 118 additions & 0 deletions tests/firedrake/adjoint/test_adj_kwargs.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
import pytest
import warnings

from firedrake import *
from firedrake.adjoint import *


@pytest.fixture(autouse=True)
def autouse_set_test_tape(set_test_tape):
pass


@pytest.fixture()
def forcing():
mesh = UnitSquareMesh(20,20)

Check failure on line 15 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E231

tests/firedrake/adjoint/test_adj_kwargs.py:15:29: E231 missing whitespace after ','
U = FunctionSpace(mesh, "CG", 1)
return Function(U)


@pytest.fixture()
def nvp(forcing):
U = forcing.function_space()
u = Function(U)
v = TestFunction(U)

F = inner(grad(u), grad(v))*dx - inner(forcing, v)*dx

bcs = [DirichletBC(U, Constant(1), (4,)),
DirichletBC(U, Constant(0), (1, 2, 3))]

Check failure on line 29 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E128

tests/firedrake/adjoint/test_adj_kwargs.py:29:9: E128 continuation line under-indented for visual indent

return NonlinearVariationalProblem(F, u, bcs=bcs)


# Spike the adjoint solve so we can check that the options are actually
# set.
sp_lu = {
"pc_type": "none",

Check failure on line 37 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E126

tests/firedrake/adjoint/test_adj_kwargs.py:37:9: E126 continuation line over-indented for hanging indent
"ksp_type": "cg",
"ksp_max_it": 1,
"ksp_view": None
}

Check failure on line 41 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E123

tests/firedrake/adjoint/test_adj_kwargs.py:41:9: E123 closing bracket does not match indentation of opening bracket's line

#@pytest.mark.parametrize("mode", ("solver", "solve", "both", "none", "free"))

Check failure on line 43 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E265

tests/firedrake/adjoint/test_adj_kwargs.py:43:1: E265 block comment should start with '# '


@pytest.mark.skipcomplex
def test_adj_kwargs_solver(nvp, forcing):
nvs = NonlinearVariationalSolver(nvp,
adj_kwargs={"solver_parameters": sp_lu})
nvs.solve()
u = nvp.u

J = assemble(inner(u, u)*dx)
Jhat = ReducedFunctional(J, Control(forcing))

with pytest.raises(ConvergenceError):
Jhat.derivative()

Check failure on line 58 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

W293

tests/firedrake/adjoint/test_adj_kwargs.py:58:1: W293 blank line contains whitespace

@pytest.mark.skipcomplex
def test_adj_kwargs_none(nvp, forcing):
# Don't pass adj_kwargs anywhere. Should succeed.

Check failure on line 62 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E265

tests/firedrake/adjoint/test_adj_kwargs.py:62:5: E265 block comment should start with '# '
nvs = NonlinearVariationalSolver(nvp)
nvs.solve()
u=nvp.u

Check failure on line 65 in tests/firedrake/adjoint/test_adj_kwargs.py

View workflow job for this annotation

GitHub Actions / test / Lint codebase

E225

tests/firedrake/adjoint/test_adj_kwargs.py:65:6: E225 missing whitespace around operator
J = assemble(inner(u, u)*dx)
Jhat = ReducedFunctional(J, Control(forcing))
Jhat.derivative()


@pytest.mark.skipcomplex
def test_adj_kwargs_both(nvp):
nvs = NonlinearVariationalSolver(nvp,
adj_kwargs={"solver_parameters": sp_lu})
# Passing adj_kwargs to both the solver and solve() is an error.
with pytest.raises(TypeError):
nvs.solve(adj_kwargs={"solver_parameters": sp_lu})


@pytest.mark.skipcomplex
def test_adj_kwargs_solve(nvp, forcing):
nvs = NonlinearVariationalSolver(nvp)
with warnings.catch_warnings(record=True) as w:
# Cause all warnings to always be triggered.
warnings.simplefilter("always")
# Trigger a warning.
nvs.solve(adj_kwargs={"solver_parameters": sp_lu})
# Verify some things
assert len(w) == 1
assert issubclass(w[-1].category, FutureWarning)

u = nvp.u
J = assemble(inner(u, u)*dx)
Jhat = ReducedFunctional(J, Control(forcing))

with pytest.raises(ConvergenceError):
Jhat.derivative()


# Unclear why this doesn't work.
@pytest.mark.xfail
@pytest.mark.skipcomplex
def test_adj_kwargs_solve_free_function(nvp, forcing):
u = nvp.u

with warnings.catch_warnings(record=True) as w:
# Cause all warnings to always be triggered.
warnings.simplefilter("always")
# Should not trigger warnings.
solve(nvp.F == 0, u, adj_kwargs={"solver_parameters": sp_lu})
# Verify no warning raised.
assert not w

J = assemble(inner(u, u)*dx)
Jhat = ReducedFunctional(J, Control(forcing))

with pytest.raises(ConvergenceError):
Jhat.derivative()
Loading