diff --git a/firedrake/variational_solver.py b/firedrake/variational_solver.py index 4031bf3c5c..21eb1d1185 100644 --- a/firedrake/variational_solver.py +++ b/firedrake/variational_solver.py @@ -96,18 +96,34 @@ def __init__(self, F, u, bcs=None, J=None, V_res = restricted_function_space(V, extract_subdomain_ids(bcs)) bcs = [bc.reconstruct(V=V_res, indices=bc._indices) for bc in bcs] self.u_restrict = Function(V_res) + v_res, u_res = TestFunction(V_res), TrialFunction(V_res) + + P = interpolate(u_res, V) + Pstar = ufl_expr.adjoint(P) + u_full = interpolate(self.u_restrict, V) + if isinstance(F, Form): F_arg, = F.arguments() self.F = replace(F, {F_arg: v_res, self.u: self.u_restrict}) else: - self.F = interpolate(v_res, replace(F, {self.u: self.u_restrict})) + F_full = replace(F, {self.u: u_full}) + self.F = ufl_expr.action(Pstar, F_full) + + if isinstance(self.J, Form): + v_arg, u_arg = self.J.arguments() + self.J = replace(self.J, {v_arg: v_res, u_arg: u_res, self.u: self.u_restrict}) + else: + J_full = replace(self.J, {self.u: u_full}) + self.J = ufl_expr.action(Pstar, ufl_expr.action(J_full, P)) - v_arg, u_arg = self.J.arguments() - self.J = replace(self.J, {v_arg: v_res, u_arg: u_res, self.u: self.u_restrict}) if self.Jp: - v_arg, u_arg = self.Jp.arguments() - self.Jp = replace(self.Jp, {v_arg: v_res, u_arg: u_res, self.u: self.u_restrict}) + if isinstance(self.Jp, Form): + v_arg, u_arg = self.Jp.arguments() + self.Jp = replace(self.Jp, {v_arg: v_res, u_arg: u_res, self.u: self.u_restrict}) + else: + Jp_full = replace(self.Jp, {self.u: u_full}) + self.Jp = ufl_expr.action(Pstar, ufl_expr.action(Jp_full, P)) self.restricted_space = V_res else: self.u_restrict = u diff --git a/tests/firedrake/regression/test_restricted_function_space.py b/tests/firedrake/regression/test_restricted_function_space.py index 727c4e5101..5c4ccd43d4 100644 --- a/tests/firedrake/regression/test_restricted_function_space.py +++ b/tests/firedrake/regression/test_restricted_function_space.py @@ -209,6 +209,39 @@ def test_poisson_inhomogeneous_bcs_high_level_interface(assembled_rhs): assert errornorm(SpatialCoordinate(mesh)[0]**2, u) < 1.e-12 +@pytest.mark.parametrize("preassembled", [False, True]) +def test_restrict_action(preassembled): + mesh = UnitSquareMesh(4, 4) + V = FunctionSpace(mesh, "CG", 1) + Q = FunctionSpace(mesh, "DG", 0) + + u, v = TrialFunction(V), TestFunction(V) + q, _ = TrialFunction(Q), TestFunction(Q) + + M = inner(q, v) * dx # Q x V -> R + I = interpolate(u, Q) # V x Q^* -> R + if preassembled: + M = assemble(M) + I = assemble(I) + A = action(M, I) # V x V^* -> R + + solution = Function(V) + bc = DirichletBC(V, 0, "on_boundary") + L = inner(1, v) * dx + + problem = LinearVariationalProblem(A, L, solution, bcs=bc, restrict=True) + + V_res = problem.restricted_space + assert problem.F.arguments()[0].function_space() == V_res + assert all(arg.function_space() == V_res for arg in problem.J.arguments()) + + matrix = assemble(A).petscmat[:, :] + matrix = np.delete(matrix, bc.nodes, axis=0) + matrix = np.delete(matrix, bc.nodes, axis=1) + restricted_matrix = assemble(problem.J).petscmat[:, :] + assert np.allclose(matrix, restricted_matrix) + + @pytest.mark.parametrize("j", [1, 2, 5]) def test_restricted_function_space_coord_change(j): mesh = UnitSquareMesh(1, 2)