diff --git a/.github/actions/install/action.yml b/.github/actions/install/action.yml index 54c1411d7f..e55f629855 100644 --- a/.github/actions/install/action.yml +++ b/.github/actions/install/action.yml @@ -150,11 +150,14 @@ runs: EXTRA_PIP_FLAGS='--no-build-isolation' elif [ ${{ inputs.base_ref }} = 'release' ]; then EXTRA_PIP_FLAGS='' + export PIP_BUILD_CONSTRAINT=constraints.txt else echo "Unrecognised 'base_ref' input: '${{ inputs.base_ref }}" exit 1 fi + : # DROP BEFORE MERGE + pip install -v --no-deps --ignore-installed git+https://github.com/firedrakeproject/fiat.git@pbrubeck/fix/dual-enriched-again pip install --verbose $EXTRA_PIP_FLAGS \ --no-binary h5py \ --extra-index-url https://download.pytorch.org/whl/cpu \ diff --git a/.github/workflows/core.yml b/.github/workflows/core.yml index 44b436660f..89e093b0bd 100644 --- a/.github/workflows/core.yml +++ b/.github/workflows/core.yml @@ -582,6 +582,9 @@ jobs: - name: Install Firedrake id: install run: | + : # DROP BEFORE MERGE + pip install -v --no-deps --ignore-installed git+https://github.com/firedrakeproject/fiat.git@pbrubeck/fix/dual-enriched-again + pip install -v --no-deps --ignore-installed git+https://github.com/firedrakeproject/ufl.git@pbrubeck/interpolate-reference-lowering pip install --verbose -r ./firedrake-repo/requirements-build.txt CC=mpicc CXX=mpicxx \ pip install --verbose --no-build-isolation './firedrake-repo[docs]' diff --git a/docs/source/conf.py b/docs/source/conf.py index b13e8fc83c..609ed119eb 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -167,6 +167,7 @@ ('py:class', 'pyop2.caching.Cached'), ('py:class', 'pyop2.op2.Kernel'), ('py:class', 'pyop2.types.mat.Mat'), + ('py:class', 'pyop2.types.access.Access'), # Ignore mission docs from Firedrake internal "private" code # Any "Base" class eg: # firedrake.adjoint.checkpointing.CheckpointBase diff --git a/docs/source/interpolation.rst b/docs/source/interpolation.rst index 1d0a9c233d..8c158987aa 100644 --- a/docs/source/interpolation.rst +++ b/docs/source/interpolation.rst @@ -458,7 +458,8 @@ each block given by \end{pmatrix} The off-diagonal blocks are zero since the dofs are applied component-wise. Firedrake's form -compiler recognises this and avoids assembling the zero blocks. +compiler recognises this and avoids assembling the zero blocks, which the nest +still allocates. We can assemble more general interpolation matrices between mixed function spaces by interpolating vector expressions with arguments. For example, by doing diff --git a/firedrake/__init__.py b/firedrake/__init__.py index 29e9ea9664..3bd63a02fb 100644 --- a/firedrake/__init__.py +++ b/firedrake/__init__.py @@ -70,6 +70,7 @@ def init_petsc(): FiredrakeException, ConvergenceError, MismatchingDomainError, VertexOnlyMeshMissingPointsError, DofNotDefinedError, DofTypeError, SerialExecutionOnlyError, PointNotInDomainError, + MismatchingFunctionSpaceError, ) from firedrake.function import ( # noqa: F401 Function, CoordinatelessFunction, PointEvaluator diff --git a/firedrake/assemble.py b/firedrake/assemble.py index 2a84ed28b7..96c9fd7fba 100644 --- a/firedrake/assemble.py +++ b/firedrake/assemble.py @@ -13,7 +13,7 @@ from pyadjoint.tape import annotate_tape from tsfc import kernel_args from finat.element_factory import create_element -from tsfc.ufl_utils import extract_firedrake_constants +from tsfc.ufl_utils import extract_firedrake_constants, RUNTIME_POINT_VARIABLE import ufl import finat.ufl from firedrake import (extrusion_utils as eutils, parameters, solving, @@ -21,10 +21,12 @@ from firedrake.adjoint_utils import annotate_assemble from firedrake.ufl_expr import extract_domains from firedrake.bcs import DirichletBC, EquationBC, EquationBCSplit +from firedrake.exceptions import MismatchingFunctionSpaceError from firedrake.matrix import MatrixBase, Matrix, ImplicitMatrix +from firedrake.mesh import MeshGeometry, VertexOnlyMeshTopology from firedrake.functionspaceimpl import WithGeometry, FunctionSpace, FiredrakeDualSpace from firedrake.functionspacedata import entity_dofs_key, entity_permutations_key -from firedrake.interpolation import get_interpolator +from firedrake.interpolation import get_assembly_entity_node_map, get_interpolator, SameMeshInterpolator from firedrake.petsc import PETSc from firedrake.slate import slac, slate from firedrake.slate.slac.kernel_builder import CellFacetKernelArg, LayerCountKernelArg @@ -151,6 +153,28 @@ def assemble(expr, *args, **kwargs): return get_assembler(expr, *args, **kwargs).assemble(**assemble_kwargs) +def get_form_assembler(form: ufl.form.Form | ufl.Interpolate | slate.TensorBase, + *args, **kwargs) -> "FormAssembler": + """Construct the assembler for the rank of ``form``, forwarding the options it takes.""" + diagonal = kwargs.pop("diagonal", False) + nargs = len(form.arguments()) + if nargs == 0: + return ZeroFormAssembler(form, form_compiler_parameters=kwargs.get("form_compiler_parameters")) + elif nargs == 1 or diagonal: + return OneFormAssembler(form, *args, + bcs=kwargs.get("bcs", None), + form_compiler_parameters=kwargs.get("form_compiler_parameters"), + needs_zeroing=kwargs.get("needs_zeroing", True), + zero_bc_nodes=kwargs.get("zero_bc_nodes", True), + access=kwargs.get("access", op2.INC), + diagonal=diagonal, + weight=kwargs.get("weight", 1.0)) + elif nargs == 2: + return TwoFormAssembler(form, *args, **kwargs) + else: + raise ValueError('Expecting a 0-, 1-, or 2-form: got %s' % (form)) + + def get_assembler(form, *args, **kwargs): """Create an assembler. @@ -171,31 +195,20 @@ def get_assembler(form, *args, **kwargs): # Only pre-process `form` once beforehand to avoid pre-processing for each assembly call form = BaseFormAssembler.preprocess_base_form(form, mat_type=mat_type, form_compiler_parameters=fc_params) if isinstance(form, (ufl.form.Form, slate.TensorBase)) and not BaseFormAssembler.base_form_operands(form): - diagonal = kwargs.pop('diagonal', False) - if len(form.arguments()) == 0: - return ZeroFormAssembler(form, form_compiler_parameters=fc_params) - elif len(form.arguments()) == 1 or diagonal: - return OneFormAssembler(form, *args, - bcs=kwargs.get("bcs", None), - form_compiler_parameters=fc_params, - needs_zeroing=kwargs.get("needs_zeroing", True), - zero_bc_nodes=kwargs.get("zero_bc_nodes", True), - diagonal=diagonal, - weight=kwargs.get("weight", 1.0)) - elif len(form.arguments()) == 2: - return TwoFormAssembler(form, *args, **kwargs) - else: - raise ValueError('Expecting a 0-, 1-, or 2-form: got %s' % (form)) + return get_form_assembler(form, *args, **kwargs) elif isinstance(form, ufl.core.expr.Expr) and not isinstance(form, ufl.core.base_form_operator.BaseFormOperator): # BaseForm preprocessing can turn BaseForm into an Expr (cf. case (6) in `restructure_base_form`) return ExprAssembler(form) elif isinstance(form, ufl.form.BaseForm): return BaseFormAssembler(form, *args, **kwargs) + elif isinstance(form, slate.TensorBase): + raise NotImplementedError("Assemble the interpolation in this Slate tensor first: " + "TSFC cannot fuse it into the kernels of the form that holds it.") else: raise ValueError(f'Expecting a BaseForm, slate.TensorBase, or Expr object: got {form}') -class ExprAssembler(object): +class ExprAssembler: """Expression assembler. Parameters @@ -319,6 +332,37 @@ def assemble(self, tensor=None, current_state=None): """ +def _can_fuse_operator(operator: ufl.core.base_form_operator.BaseFormOperator, + covering_domains: set) -> bool: + """Can TSFC assemble ``operator`` in the kernel of the expression that holds it? + + ``covering_domains`` are the domains that the enclosing integral's measure covers, + or the domains that the enclosing interpolation targets. + """ + if not isinstance(operator, ufl.Interpolate): + return False + # A non-terminal dual argument needs to be assembled on its own + # beforehand, so it cannot share a kernel with the expression. + dual_arg, expression = operator.argument_slots() + if not isinstance(dual_arg, (ufl.Coargument, ufl.Cofunction)): + return False + if not all(any(domain is covering_domain + or (isinstance(domain, MeshGeometry) and isinstance(covering_domain, MeshGeometry) + and domain.submesh_ancestors[-1] is covering_domain.submesh_ancestors[-1] + and domain.topological_dimension == covering_domain.topological_dimension) + for covering_domain in covering_domains) + for domain in extract_domains(operator)): + return False + # The expression that holds the interpolation iterates over its own cells, + # so it can only absorb an interpolation that maps cells to cells and that + # needs no subset. Any other interpolation keeps its own interpolator. + interpolator = get_interpolator(operator) + if not isinstance(interpolator, SameMeshInterpolator) or interpolator.subset is not None: + return False + return all(_can_fuse_operator(op, covering_domains) + for op in ufl.algorithms.extract_base_form_operators(expression)) + + class BaseFormAssembler(AbstractFormAssembler): """Base form assembler. @@ -453,19 +497,22 @@ def base_form_assembly_visitor(self, expr, tensor, bcs, *args): # Substitute the base form operators by their output expr = ufl.replace(expr, dict(zip(base_form_operators, args))) form = expr - rank = len(form.arguments()) - if rank == 0: - assembler = ZeroFormAssembler(form, form_compiler_parameters=self._form_compiler_params) - elif rank == 1 or (rank == 2 and self._diagonal): - assembler = OneFormAssembler(form, form_compiler_parameters=self._form_compiler_params, - zero_bc_nodes=self._zero_bc_nodes, diagonal=self._diagonal, weight=self._weight) - elif rank == 2: - assembler = TwoFormAssembler(form, bcs=bcs, form_compiler_parameters=self._form_compiler_params, - mat_type=self._mat_type, sub_mat_type=self._sub_mat_type, - options_prefix=self._options_prefix, appctx=self._appctx, weight=self._weight, - allocation_integral_types=self.allocation_integral_types) - else: - raise AssertionError + # The bcs of a 1-form are applied to the assembled result instead, + # so only a matrix takes them here, and only a matrix allocates a + # sparsity to match the one that the result was allocated with. + is_matrix = len(form.arguments()) == 2 and not self._diagonal + assembler = get_form_assembler( + form, + bcs=bcs if is_matrix else (), + form_compiler_parameters=self._form_compiler_params, + mat_type=self._mat_type, + sub_mat_type=self._sub_mat_type, + options_prefix=self._options_prefix, + appctx=self._appctx, + diagonal=self._diagonal, + weight=self._weight, + allocation_integral_types=self.allocation_integral_types if is_matrix else None, + ) return assembler.assemble(tensor=tensor) elif isinstance(expr, ufl.Adjoint): if len(args) != 1: @@ -691,14 +738,28 @@ def reconstruct_node_from_operands(expr, operands): def base_form_operands(expr): if isinstance(expr, (ufl.FormSum, ufl.Adjoint, ufl.Action)): return expr.ufl_operands + if isinstance(expr, slate.TensorBase): + # A zero tensor wraps no integrals for TSFC to compile. + operands = [BaseFormAssembler.base_form_operands(child) + for child in expr.operands or (expr.form,) + if not isinstance(child, ufl.ZeroBaseForm)] + return list(dict.fromkeys(itertools.chain.from_iterable(operands))) if isinstance(expr, ufl.Form): + # A fusible interpolation is not a child: descending would assemble it alone. + children = set() + for integral in expr.integrals(): + domains = {integral.ufl_domain(), *integral.extra_domain_integral_type_map()} + children.update(op for op in ufl.algorithms.extract_base_form_operators(integral.integrand()) + if not _can_fuse_operator(op, domains)) # Use reversed to treat base form operators # in the order in which they have been made. - return list(reversed(expr.base_form_operators())) + return [op for op in reversed(expr.base_form_operators()) if op in children] if isinstance(expr, ufl.core.base_form_operator.BaseFormOperator): + # An interpolation shares the kernel of the one that targets its domains. + domains = set(extract_domains(expr.argument_slots()[0])) if isinstance(expr, ufl.Interpolate) else set() # Conserve order children = dict.fromkeys(e for e in (expr.argument_slots() + expr.ufl_operands) - if isinstance(e, ufl.form.BaseForm)) + if isinstance(e, ufl.form.BaseForm) and not _can_fuse_operator(e, domains)) return list(children) return [] @@ -893,8 +954,8 @@ def preprocess_base_form(expr, mat_type=None, form_compiler_parameters=None): expr = BaseFormAssembler.restructure_base_form_postorder(expr) # Preprocessing the form makes a new object -> current form caching mechanism # will populate `expr`'s cache which is now different than `original_expr`'s cache so we need - # to transmit the cache. All of this only holds when both are `ufl.Form` objects. - if isinstance(original_expr, ufl.form.Form) and isinstance(expr, ufl.form.Form): + # to transmit the cache. Both objects must support the assembler cache. + if isinstance(original_expr, (ufl.form.Form, ufl.Interpolate)) and isinstance(expr, (ufl.form.Form, ufl.Interpolate)): expr._cache = original_expr._cache return expr @@ -949,8 +1010,8 @@ class FormAssembler(AbstractFormAssembler): def __new__(cls, *args, **kwargs): form = args[0] - if not isinstance(form, (ufl.Form, slate.TensorBase)): - raise TypeError(f"The first positional argument must be of ufl.Form or slate.TensorBase: got {type(form)} ({form})") + if not isinstance(form, (ufl.Form, ufl.Interpolate, slate.TensorBase)): + raise TypeError(f"The first positional argument must be of ufl.Form, ufl.Interpolate, or slate.TensorBase: got {type(form)} ({form})") # It is expensive to construct new assemblers because extracting the data # from the form is slow. Since all of the data structures in the assembler # are persistent apart from the output tensor, we stash the assembler on the @@ -971,9 +1032,10 @@ def __new__(cls, *args, **kwargs): self = super().__new__(cls) self._initialised = False self.__init__(*args, **kwargs) - if _FORM_CACHE_KEY not in form._cache: - form._cache[_FORM_CACHE_KEY] = {} - form._cache[_FORM_CACHE_KEY][key] = self + if key is not None: + if _FORM_CACHE_KEY not in form._cache: + form._cache[_FORM_CACHE_KEY] = {} + form._cache[_FORM_CACHE_KEY][key] = self return self @classmethod @@ -1012,11 +1074,12 @@ class ParloopFormAssembler(FormAssembler): Should ``tensor`` be zeroed before assembling? """ - def __init__(self, form, bcs=None, form_compiler_parameters=None, needs_zeroing=True): + def __init__(self, form, bcs=None, form_compiler_parameters=None, needs_zeroing=True, access=op2.INC): super().__init__(form, bcs=bcs, form_compiler_parameters=form_compiler_parameters) self._needs_zeroing = needs_zeroing + self._access = access - def assemble(self, tensor=None, current_state=None): + def assemble(self, tensor=None, current_state=None, needs_zeroing: bool | None = None): """Assemble the form. Parameters @@ -1026,6 +1089,9 @@ def assemble(self, tensor=None, current_state=None): current_state : firedrake.function.Function or None If provided, the boundary condition nodes are set to the boundary condition residual computed as ``current_state`` minus the boundary condition value. + needs_zeroing : bool or None + Override whether to zero a supplied output tensor before assembly. If omitted, use the + value provided when constructing the assembler. Returns ------- @@ -1043,7 +1109,9 @@ def assemble(self, tensor=None, current_state=None): tensor = self.allocate() else: self._check_tensor(tensor) - if self._needs_zeroing: + if needs_zeroing is None: + needs_zeroing = self._needs_zeroing + if needs_zeroing: self._as_pyop2_type(tensor).zero() self.execute_parloops(tensor) @@ -1053,6 +1121,14 @@ def assemble(self, tensor=None, current_state=None): return self.result(tensor) + def compile(self): + """Compile the local kernels now, rather than lazily inside `assemble`. + + `DirichletBC` calls this to learn whether its value can be interpolated + before it commits to interpolating rather than projecting. + """ + self.local_kernels + @abc.abstractmethod def _apply_bc(self, tensor, bc, u=None): """Apply boundary condition.""" @@ -1113,10 +1189,10 @@ def local_kernels(self): each possible combination. """ - if isinstance(self._form, ufl.Form): + if isinstance(self._form, (ufl.Form, ufl.Interpolate)): kernels = tsfc_interface.compile_form( self._form, "form", diagonal=self.diagonal, - parameters=self._form_compiler_params + parameters=self._form_compiler_params, access=self._access ) elif isinstance(self._form, slate.TensorBase): kernels = slac.compile_expression( @@ -1217,14 +1293,15 @@ class OneFormAssembler(ParloopFormAssembler): @classmethod def _cache_key(cls, form, bcs=None, form_compiler_parameters=None, needs_zeroing=True, - zero_bc_nodes=True, diagonal=False, weight=1.0): + zero_bc_nodes=True, diagonal=False, weight=1.0, access=op2.INC): bcs = solving._extract_bcs(bcs) - return tuple(bcs), tuplify(form_compiler_parameters), needs_zeroing, zero_bc_nodes, diagonal, weight + return tuple(bcs), tuplify(form_compiler_parameters), needs_zeroing, zero_bc_nodes, diagonal, weight, access @FormAssembler._skip_if_initialised def __init__(self, form, bcs=None, form_compiler_parameters=None, needs_zeroing=True, - zero_bc_nodes=True, diagonal=False, weight=1.0): - super().__init__(form, bcs=bcs, form_compiler_parameters=form_compiler_parameters, needs_zeroing=needs_zeroing) + zero_bc_nodes=True, diagonal=False, weight=1.0, access=op2.INC): + super().__init__(form, bcs=bcs, form_compiler_parameters=form_compiler_parameters, + needs_zeroing=needs_zeroing, access=access) self._weight = weight self._diagonal = diagonal self._zero_bc_nodes = zero_bc_nodes @@ -1286,6 +1363,11 @@ def _as_pyop2_type(tensor, indices=None): return tensor.dat def execute_parloops(self, tensor): + if self._access is not op2.INC: + for parloop in self.parloops(tensor): + parloop() + return + # We are repeatedly incrementing into the same Dat so intermediate halo exchanges # can be skipped. with tensor.dat.frozen_halo(op2.INC): @@ -1301,7 +1383,7 @@ def result(self, tensor): def TwoFormAssembler(form, *args, **kwargs): - assert isinstance(form, (ufl.form.Form, slate.TensorBase)) + assert isinstance(form, (ufl.form.Form, ufl.Interpolate, slate.TensorBase)) mat_type = kwargs.pop('mat_type', None) sub_mat_type = kwargs.pop('sub_mat_type', None) mat_type, sub_mat_type = _get_mat_type(mat_type, sub_mat_type, form.arguments()) @@ -1372,8 +1454,9 @@ def _cache_key(cls, *args, **kwargs): @FormAssembler._skip_if_initialised def __init__(self, form, bcs=None, form_compiler_parameters=None, needs_zeroing=True, mat_type=None, sub_mat_type=None, options_prefix=None, appctx=None, weight=1.0, - allocation_integral_types=None): - super().__init__(form, bcs=bcs, form_compiler_parameters=form_compiler_parameters, needs_zeroing=needs_zeroing) + allocation_integral_types=None, access=op2.INC): + super().__init__(form, bcs=bcs, form_compiler_parameters=form_compiler_parameters, + needs_zeroing=needs_zeroing, access=access) self._mat_type = mat_type self._sub_mat_type = sub_mat_type self._options_prefix = options_prefix @@ -1438,8 +1521,8 @@ def _make_maps_and_regions(self): # Make Sparsity independent of the subdomain of integration for better reusability; # subdomain_id is passed here only to determine the integration_type on the target domain # (see ``entity_node_map``). - rmap_ = test.function_space().topological[i].entity_node_map(mesh.topology, integral_type, subdomain_id, all_subdomain_ids) - cmap_ = trial.function_space().topological[j].entity_node_map(mesh.topology, integral_type, subdomain_id, all_subdomain_ids) + rmap_ = get_assembly_entity_node_map(test.function_space()[i], mesh, integral_type, subdomain_id, all_subdomain_ids) + cmap_ = get_assembly_entity_node_map(trial.function_space()[j], mesh, integral_type, subdomain_id, all_subdomain_ids) region = ExplicitMatrixAssembler._integral_type_region_map[integral_type] maps_and_regions[(i, j)][(rmap_, cmap_)].add(region) return {block_indices: [map_pair + (tuple(region_set), ) for map_pair, region_set in map_pair_to_region_set.items()] @@ -1458,8 +1541,8 @@ def _make_maps_and_regions_default(test, trial, allocation_integral_types): for i, Vrow in enumerate(test.function_space()): for j, Vcol in enumerate(trial.function_space()): mesh = Vrow.mesh() - rmap_ = Vrow.topological.entity_node_map(mesh.topology, integral_type, None, None) - cmap_ = Vcol.topological.entity_node_map(mesh.topology, integral_type, None, None) + rmap_ = get_assembly_entity_node_map(Vrow, mesh, integral_type, None, None) + cmap_ = get_assembly_entity_node_map(Vcol, mesh, integral_type, None, None) maps_and_regions[(i, j)][(rmap_, cmap_)].add(region) return {block_indices: [map_pair + (tuple(region_set), ) for map_pair, region_set in map_pair_to_region_set.items()] for block_indices, map_pair_to_region_set in maps_and_regions.items()} @@ -1499,9 +1582,10 @@ def _apply_bc(self, tensor, bc, u=None): index = 0 if V.index is None else V.index space = V if V.parent is None else V.parent if isinstance(bc, DirichletBC): - if not any(space == fs for fs in spaces): - raise TypeError("bc space does not match the test or trial function space") - if spaces[0] != spaces[1]: + if not any(bc.parent_function_space.topological == fs.topological for fs in spaces): + raise MismatchingFunctionSpaceError( + "bc space does not match the test or trial function space") + if spaces[0].topological != spaces[1].topological: # Not on a diagonal block, we cannot set diagonal entries return @@ -1624,7 +1708,7 @@ def _global_kernel_cache_key(form, local_knl, subdomain_id, all_integer_subdomai all_meshes = extract_domains(form) domain_ids = tuple(mesh.ufl_id() for mesh in all_meshes) - if isinstance(form, ufl.Form): + if isinstance(form, (ufl.Form, ufl.Interpolate)): sig = form.signature() elif isinstance(form, slate.TensorBase): sig = form.expression_hash @@ -1768,9 +1852,14 @@ def _get_dim(self, finat_element): else: return (1,) + def _get_map(self, V): + """Return the appropriate PyOP2 map for a given function space.""" + return get_assembly_entity_node_map(V, self._mesh, self._integral_type, + self._subdomain_id, self._all_integer_subdomain_ids) + def _make_dat_global_kernel_arg(self, V, index=None): finat_element = create_element(V.ufl_element()) - map_arg = V.topological.entity_node_map(self._mesh.topology, self._integral_type, self._subdomain_id, self._all_integer_subdomain_ids)._global_kernel_arg + map_arg = self._get_map(V)._global_kernel_arg if isinstance(finat_element, finat.EnrichedElement) and finat_element.is_mixed: assert index is None subargs = tuple(self._make_dat_global_kernel_arg(Vsub, index=index) @@ -1788,7 +1877,7 @@ def _make_mat_global_kernel_arg(self, Vrow, Vcol): shape = len(relem.elements), len(celem.elements) return op2.MixedMatKernelArg(subargs, shape) else: - rmap_arg, cmap_arg = (V.topological.entity_node_map(self._mesh.topology, self._integral_type, self._subdomain_id, self._all_integer_subdomain_ids)._global_kernel_arg for V in [Vrow, Vcol]) + rmap_arg, cmap_arg = (self._get_map(V)._global_kernel_arg for V in [Vrow, Vcol]) # PyOP2 matrix objects have scalar dims so we flatten them here rdim = numpy.prod(self._get_dim(relem), dtype=int) cdim = numpy.prod(self._get_dim(celem), dtype=int) @@ -1811,6 +1900,18 @@ def _get_map_id(finat_element): return entity_dofs_key(finat_element.entity_dofs()), real_tensorproduct, eperm_key +def _runtime_tabulation_coordinates(arg: kernel_args.TabulationKernelArg, + mesh: MeshGeometry) -> "firedrake.Function": + """Return the reference coordinates that a kernel tabulates its runtime point at.""" + # FIXME: name-matching is a stopgap; get this from a coefficient map handed + # down by compile_expression_dual_evaluation, or a Cofunction/Coargument target. + if arg.loopy_arg.name != RUNTIME_POINT_VARIABLE: + raise ValueError(f"Expecting the runtime tabulation argument {RUNTIME_POINT_VARIABLE}: got {arg.loopy_arg.name}") + if not isinstance(mesh.topology, VertexOnlyMeshTopology): + raise ValueError(f"Runtime tabulation is only supported on a VertexOnlyMesh: got {type(mesh.topology).__name__}") + return mesh.reference_coordinates + + @functools.singledispatch def _as_global_kernel_arg(tsfc_arg, self): raise NotImplementedError @@ -1890,6 +1991,12 @@ def _as_global_kernel_arg_constant(_, self): return op2.GlobalKernelArg((value_size,)) +@_as_global_kernel_arg.register(kernel_args.TabulationKernelArg) +def _as_global_kernel_arg_tabulation(arg, self): + reference_coordinates = _runtime_tabulation_coordinates(arg, self._mesh) + return self._make_dat_global_kernel_arg(reference_coordinates.function_space()) + + @_as_global_kernel_arg.register(kernel_args.ExteriorFacetKernelArg) def _as_global_kernel_arg_exterior_facet(_, self): mesh = next(self._active_exterior_facets) @@ -2048,18 +2155,14 @@ def get_indicess(self): def _filter_bcs(self, row, col): assert len(self._form.arguments()) == 2 and not self._diagonal - if len(self.test_function_space) > 1: - bcrow = tuple(bc for bc in self._bcs - if bc.function_space_index() == row) - else: - bcrow = self._bcs - - if len(self.trial_function_space) > 1: - bccol = tuple(bc for bc in self._bcs - if bc.function_space_index() == col - and isinstance(bc, DirichletBC)) - else: - bccol = tuple(bc for bc in self._bcs if isinstance(bc, DirichletBC)) + bcrow = tuple(bc for bc in self._bcs + if bc.parent_function_space.topological == self.test_function_space.topological + and (len(self.test_function_space) == 1 or bc.function_space_index() == row)) + + bccol = tuple(bc for bc in self._bcs + if isinstance(bc, DirichletBC) + and bc.parent_function_space.topological == self.trial_function_space.topological + and (len(self.trial_function_space) == 1 or bc.function_space_index() == col)) return bcrow, bccol def needs_unrolling(self): @@ -2158,7 +2261,8 @@ def _iterset(self): def _get_map(self, V): """Return the appropriate PyOP2 map for a given function space.""" assert isinstance(V, (WithGeometry, FiredrakeDualSpace, FunctionSpace)) - return V.topological.entity_node_map(self._mesh.topology, self._integral_type, self._subdomain_id, self._all_integer_subdomain_ids) + return get_assembly_entity_node_map(V, self._mesh, self._integral_type, + self._subdomain_id, self._all_integer_subdomain_ids) def _as_parloop_arg(self, tsfc_arg): """Return a :class:`op2.ParloopArg` corresponding to the provided @@ -2238,6 +2342,13 @@ def _as_parloop_arg_constant(arg, self): return op2.GlobalParloopArg(const.dat) +@_as_parloop_arg.register(kernel_args.TabulationKernelArg) +def _as_parloop_arg_tabulation(arg, self): + reference_coordinates = _runtime_tabulation_coordinates(arg, self._mesh) + map_ = self._get_map(reference_coordinates.function_space()) + return op2.DatParloopArg(reference_coordinates.dat, map_) + + @_as_parloop_arg.register(kernel_args.ExteriorFacetKernelArg) def _as_parloop_arg_exterior_facet(_, self): mesh = next(self._active_exterior_facets) diff --git a/firedrake/exceptions.py b/firedrake/exceptions.py index e7c72cab43..c2cfe9a26d 100644 --- a/firedrake/exceptions.py +++ b/firedrake/exceptions.py @@ -43,6 +43,12 @@ def __str__(self): ) +class MismatchingFunctionSpaceError(FiredrakeException): + """Raised when a function space does not match the one expected, such as a + boundary condition applied to a form whose arguments live elsewhere. + """ + + class NonUniqueMeshSequenceError(FiredrakeException): """Raised when calling `.unique()` on a MeshSequence which contains non-unique meshes. diff --git a/firedrake/formmanipulation.py b/firedrake/formmanipulation.py index 28c4c72867..aa8656760b 100644 --- a/firedrake/formmanipulation.py +++ b/firedrake/formmanipulation.py @@ -1,9 +1,10 @@ import numpy import collections +from collections.abc import Iterable from ufl import as_tensor, as_vector, split -from ufl.classes import Form, Zero, FixedIndex, ListTensor, ZeroBaseForm +from ufl.classes import Expr, Form, Interpolate, Zero, FixedIndex, ListTensor, ZeroBaseForm from ufl.algorithms.map_integrands import map_integrand_dags from ufl.algorithms import expand_derivatives from ufl.corealg.map_dag import MultiFunction, map_expr_dags @@ -13,6 +14,7 @@ from firedrake.petsc import PETSc from firedrake.functionspace import MixedFunctionSpace +from firedrake.functionspaceimpl import WithGeometry from firedrake.cofunction import Cofunction from firedrake.ufl_expr import Coargument @@ -30,6 +32,12 @@ class ExtractSubBlock(MultiFunction): """Extract a sub-block from a form.""" + def __init__(self): + super().__init__() + self._arg_cache = {} + self.blocks = {} + self._splitting_interpolate = False + class IndexInliner(MultiFunction): """Inline fixed index of list tensors""" expr = MultiFunction.reuse_if_untouched @@ -82,6 +90,9 @@ def split(self, form, argument_indices): args = form.arguments() self._arg_cache = {} self.blocks = dict(enumerate(map(as_tuple, argument_indices))) + # An outermost Interpolate splits into a smaller Interpolate, while one + # inside an integrand must keep the value shape its neighbours expect. + self._splitting_interpolate = isinstance(form, Interpolate) if len(args) == 0: # Functional can't be split return form @@ -229,9 +240,38 @@ def matrix(self, o): def zero_base_form(self, o): return ZeroBaseForm(tuple(map(self, o.arguments()))) + def _zero_interpolate(self, o: Interpolate) -> ZeroBaseForm | Zero: + """Result of an Interpolate whose operand or target block is Zero.""" + if self._splitting_interpolate: + return self(ZeroBaseForm(o.arguments())) + return Zero(o.ufl_shape) + + @staticmethod + def _select_components(V: WithGeometry, indices: tuple, operand: Expr) -> list: + """Flatten the sub-blocks of ``operand`` whose subspace is in ``indices``.""" + components = [] + cur = 0 + for i, Vi in enumerate(V): + if i in indices: + components.extend(operand[k] for k in range(cur, cur + Vi.value_size)) + cur += Vi.value_size + return components + + @staticmethod + def _embed_components(V: WithGeometry, indices: tuple, values: Iterable) -> list: + """Embed ``values`` into V's full shape, zero outside ``indices``.""" + values = iter(values) + components = [] + for i, Vi in enumerate(V): + if i in indices: + components.extend(next(values) for _ in range(Vi.value_size)) + else: + components.extend(Zero() for _ in range(Vi.value_size)) + return components + def interpolate(self, o, operand): if isinstance(operand, Zero): - return self(ZeroBaseForm(o.arguments())) + return self._zero_interpolate(o) dual_arg, _ = o.argument_slots() if len(dual_arg.arguments()) == 1 or len(dual_arg.arguments()[-1].function_space()) == 1: @@ -249,18 +289,21 @@ def interpolate(self, o, operand): W = sub_dual_arg.function_space() # Unflatten the expression into the target shape - cur = 0 - components = [] - for i, Vi in enumerate(V): - if i in indices: - components.extend(operand[i] for i in range(cur, cur+Vi.value_size)) - cur += Vi.value_size - + components = self._select_components(V, indices, operand) operand = as_tensor(numpy.reshape(components, W.value_shape)) if isinstance(operand, Zero): - return self(ZeroBaseForm(o.arguments())) - - return o._ufl_expr_reconstruct_(operand, sub_dual_arg) + return self._zero_interpolate(o) + + interpolation = o._ufl_expr_reconstruct_(operand, sub_dual_arg) + if self._splitting_interpolate: + return interpolation + + # Inside an integrand the block is one part of a wider expression, so + # pad it back out to V's value shape with zeros. + interpolation_components = ((interpolation[j] for j in numpy.ndindex(interpolation.ufl_shape)) + if interpolation.ufl_shape else iter((interpolation,))) + components = self._embed_components(V, indices, interpolation_components) + return as_tensor(numpy.reshape(components, V.value_shape)) SplitForm = collections.namedtuple("SplitForm", ["indices", "form"]) diff --git a/firedrake/interpolation.py b/firedrake/interpolation.py index ee8a2a2928..ccc86e1919 100644 --- a/firedrake/interpolation.py +++ b/firedrake/interpolation.py @@ -1,35 +1,25 @@ import numpy -import os -import tempfile import abc from functools import cached_property, partial -from typing import Hashable, Literal, Callable, Iterable +from typing import Literal, Callable, Iterable from dataclasses import asdict, dataclass from numbers import Number from ufl.algorithms import extract_arguments, replace from ufl.domain import extract_unique_domain from ufl.classes import Expr -from ufl.duals import is_dual -from ufl.constantvalue import zero, as_ufl +from ufl.constantvalue import as_ufl from ufl.form import ZeroBaseForm, BaseForm from ufl.core.interpolate import Interpolate as UFLInterpolate from pyop2 import op2 -from pyop2.caching import memory_and_disk_cache - from finat.ufl import TensorElement, VectorElement, MixedElement, FiniteElementBase -from finat.element_factory import create_element - -from tsfc.driver import compile_expression_dual_evaluation -from tsfc.ufl_utils import extract_firedrake_constants, hash_expr -from firedrake.utils import IntType, ScalarType, known_pyop2_safe, tuplify -from firedrake.pointeval_utils import runtime_quadrature_element -from firedrake.tsfc_interface import extract_numbered_coefficients, _cachedir -from firedrake.ufl_expr import Argument, Coargument, TrialFunction, TestFunction, action -from firedrake.mesh import MissingPointsBehaviour, VertexOnlyMeshTopology, MeshGeometry, MeshTopology, VertexOnlyMesh +from firedrake.utils import IntType +from firedrake.ufl_expr import Argument, Coargument, TrialFunction, TestFunction, action, extract_domains +from firedrake.mesh import (MissingPointsBehaviour, VertexOnlyMeshTopology, MeshGeometry, + MeshTopology, VertexOnlyMesh) from firedrake.petsc import PETSc from firedrake.halo import _get_mtype from firedrake.functionspaceimpl import WithGeometry @@ -151,6 +141,11 @@ def options(self) -> InterpolateOptions: """ return self._options + def subdomain_data(self): + """Return cell-iteration subdomain data for the target mesh.""" + domain = self.target_space.mesh().unique() + return {domain: {"cell": [self.options.subset]}} + @cached_property def _interpolator(self): """Access the numerical interpolator. @@ -163,8 +158,6 @@ def _interpolator(self): """ arguments = self.arguments() has_mixed_arguments = any(len(arg.function_space()) > 1 for arg in arguments) - if len(arguments) == 2 and has_mixed_arguments: - return MixedInterpolator(self) operand, = self.ufl_operands target_mesh = self.target_space.mesh() @@ -707,132 +700,125 @@ def __init__(self, expr, source_mesh, target_mesh): # Matrix-free assembly of 0-form or 1-form requires INC access if self.access and self.access != op2.INC: raise ValueError("Matfree adjoint interpolation requires INC access") - self.access = op2.INC - elif self.access is None: - # Default access for forward 1-form or 2-form (forward and adjoint) - self.access = op2.WRITE - def _get_tensor(self, mat_type: Literal["aij", "baij"]) -> op2.Mat | Function | Cofunction: - """Return a suitable tensor to interpolate into. + @property + def _needs_adjoint_weighting(self): + return (isinstance(self.dual_arg, Cofunction) + and any(not V.finat_element.is_dg() for V in self.target_space)) - Parameters - ---------- - mat_type - The PETSc matrix type to use when assembling a rank 2 interpolation. - Only ``"aij"`` and ``"baij"`` are currently allowed. + @cached_property + def _weighted_dual_arg(self): + return Function(self.dual_arg.function_space()) - Returns - ------- - op2.Mat | Function | Cofunction - The tensor to interpolate into. - """ - if self.rank == 0: - R = FunctionSpace(self.target_mesh.unique(), "Real", 0) - f = Function(R, dtype=ScalarType) - elif self.rank == 1: - f = Function(self.ufl_interpolate.function_space()) - if self.access in {op2.MIN, op2.MAX}: - finfo = numpy.finfo(f.dat.dtype) - if self.access == op2.MIN: - val = Constant(finfo.max) - else: - val = Constant(finfo.min) - f.assign(val) - elif self.rank == 2: - sparsity = self._get_monolithic_sparsity(mat_type) - f = op2.Mat(sparsity) + @cached_property + def _adjoint_weight(self): + W = self.dual_arg.function_space() + weight = W.make_dat() + if len(W) > 1: + spaces_and_weights = zip(W, weight) else: - raise ValueError(f"Cannot interpolate an expression with {self.rank} arguments") - return f - - def _get_monolithic_sparsity(self, mat_type: Literal["aij", "baij"]) -> op2.Sparsity: - """Returns op2.Sparsity for the interpolation matrix. Only mat_type 'aij' and 'baij' - are currently supported. + spaces_and_weights = ((W, weight),) + + target_mesh = self.target_mesh.unique() + iterset = target_mesh.cell_set if self.subset is None else self.subset + for i, (V, component_weight) in enumerate(spaces_and_weights): + node_map = get_assembly_entity_node_map(V, target_mesh) + size = V.finat_element.space_dimension() * V.block_size + kernel_code = f""" + void multiplicity_{i}(PetscScalar *restrict w) {{ + for (PetscInt i=0; i<{size}; i++) w[i] += 1; + }}""" + kernel = op2.Kernel(kernel_code, f"multiplicity_{i}") + op2.par_loop(kernel, iterset, component_weight(op2.INC, node_map)) + with weight.vec as weight_vec: + weight_vec.reciprocal() + return weight - Parameters - ---------- - mat_type - The PETSc matrix type to use when assembling a rank 2 interpolation. - Only ``"aij"`` and ``"baij"`` are currently allowed. + @cached_property + def _interpolate_to_assemble(self): + """The interpolation that is handed to the assembler. - Returns - ------- - op2.Sparsity - The sparsity pattern for the interpolation matrix. + This is the user's `Interpolate` carrying the assembly options, and the + weighted copy of the dual argument when the adjoint needs one. """ - Vrow = self.interpolate_args[0].function_space() - Vcol = self.interpolate_args[1].function_space() - if len(Vrow) > 1 or len(Vcol) > 1: - raise NotImplementedError("Interpolation matrix with MixedFunctionSpace requires MixedInterpolator") - Vrow_map = get_interp_node_map(self.source_mesh.unique(), self.target_mesh.unique(), Vrow) - Vcol_map = get_interp_node_map(self.source_mesh.unique(), self.target_mesh.unique(), Vcol) - sparsity = op2.Sparsity((Vrow.dof_dset, Vcol.dof_dset), - [(Vrow_map, Vcol_map, None)], # non-mixed - name=f"{Vrow.name}_{Vcol.name}_sparsity", - nest=False, - block_sparse=(mat_type == "baij")) - return sparsity - - def _get_callable(self, tensor=None, bcs=None, mat_type=None, sub_mat_type=None): - mat_type = mat_type or "aij" - if (isinstance(tensor, Cofunction) and isinstance(self.dual_arg, Cofunction)) and set(tensor.dat).intersection(set(self.dual_arg.dat)): - # adjoint one-form case: we need an empty tensor, so if it shares dats with - # the dual_arg we cannot use it directly, so we store it - f = self._get_tensor(mat_type) - copyout = (partial(f.dat.copy, tensor.dat),) - else: - f = tensor or self._get_tensor(mat_type) - copyout = () + options = asdict(self.ufl_interpolate.options) + access = self.access + if not isinstance(self.dual_arg, Coargument): + access = op2.INC + elif access is None: + # Default access for forward 1-form or 2-form (forward and adjoint) + access = op2.WRITE + options.update(subset=self.subset, access=access) + dual_arg = self._weighted_dual_arg if self._needs_adjoint_weighting else self.dual_arg + return self.ufl_interpolate._ufl_expr_reconstruct_(self.operand, v=dual_arg, **options) - op2_tensor = f if isinstance(f, op2.Mat) else f.dat - loops = [] - if self.access is op2.INC: - loops.append(op2_tensor.zero) + def _update_weighted_dual_arg(self): + self.dual_arg.dat.copy(self._weighted_dual_arg.dat) + with self._adjoint_weight.vec_ro as weight, self._weighted_dual_arg.dat.vec as dual: + dual.pointwiseMult(dual, weight) - # Arguments in the operand are allowed to be from a MixedFunctionSpace - # We need to split the target space V and generate separate kernels - if self.rank == 2: - expressions = {(0,): self.ufl_interpolate} - elif isinstance(self.dual_arg, Coargument): - # Split in the coargument - expressions = dict(split_form(self.ufl_interpolate)) - else: - assert isinstance(self.dual_arg, Cofunction) - # Split in the cofunction: split_form can only split in the coargument - # Replace the cofunction with a coargument to construct the Jacobian - interp = self.ufl_interpolate._ufl_expr_reconstruct_(self.operand, self.target_space) - # Split the Jacobian into blocks - interp_split = dict(split_form(interp)) - # Split the cofunction - dual_split = dict(split_form(self.dual_arg)) - # Combine the splits by taking their action - expressions = {i: action(interp_split[i], dual_split[i[-1:]]) for i in interp_split} - - # Interpolate each sub expression into each function space - for indices, sub_expr in expressions.items(): - sub_op2_tensor = op2_tensor[indices[0]] if self.rank == 1 else op2_tensor - loops.extend(_build_interpolation_callables(sub_expr, sub_op2_tensor, self.access, self.subset, bcs)) - - if bcs and self.rank == 1: - loops.extend(partial(bc.apply, f) for bc in bcs) - - loops.extend(copyout) - - def callable() -> Function | Cofunction | PETSc.Mat | Number: - for l in loops: - l() - if self.rank == 0: - return f.dat.data.item() - elif self.rank == 2: - return f.handle # In this case f is an op2.Mat - else: - return f + def _get_callable(self, tensor=None, bcs=None, mat_type=None, sub_mat_type=None): + from firedrake.assemble import get_form_assembler, ParloopFormAssembler + + output = None + preserve_input = False + if isinstance(tensor, Function | Cofunction): + inputs = set() + for coefficient in self._interpolate_to_assemble.coefficients(): + inputs.update(coefficient.dat) + for mesh in extract_domains(self._interpolate_to_assemble): + inputs.update(mesh.coordinates.dat) + if isinstance(self.dual_arg, Cofunction): + inputs.update(self.dual_arg.dat) + if set(tensor.dat) & inputs: + output = tensor + preserve_input = self.access is not None and self.access is not op2.WRITE + + access = self._interpolate_to_assemble.options.access + needs_zeroing = (self.rank == 2 or access is op2.INC) and not preserve_input + assembler = get_form_assembler(self._interpolate_to_assemble, bcs=bcs, + mat_type=mat_type, sub_mat_type=sub_mat_type, + needs_zeroing=needs_zeroing, access=access) + assemble_kwargs = {} + # DirichletBC needs to know now whether it can interpolate its value, + # so it can project instead when it can't. + if isinstance(assembler, ParloopFormAssembler): + assembler.compile() + needs_zeroing |= access is op2.WRITE and not assembler.local_kernels + assemble_kwargs["needs_zeroing"] = needs_zeroing + + copy_input = None + copy_output = None + if output is not None: + tensor = assembler.allocate() + if preserve_input: + copy_input = partial(output.dat.copy, tensor.dat) + copy_output = partial(tensor.dat.copy, output.dat) + elif tensor is None and self.access in {op2.MIN, op2.MAX}: + tensor = assembler.allocate() + finfo = numpy.finfo(tensor.dat.dtype) + value = finfo.max if self.access == op2.MIN else finfo.min + tensor.assign(Constant(value)) + + assembler_tensor = None if self.rank == 2 else tensor + + def callable(): + if self._needs_adjoint_weighting: + self._update_weighted_dual_arg() + if copy_input is not None: + copy_input() + result = assembler.assemble(tensor=assembler_tensor, **assemble_kwargs) + if copy_output is not None: + copy_output() + if isinstance(result, MatrixBase): + return result.petscmat + return output if copy_output is not None else result return callable @property def _allowed_mat_types(self): - return {"aij", "baij", "matfree", None} + return {"aij", "baij", "nest", "matfree", None} class VomOntoVomInterpolator(SameMeshInterpolator): @@ -859,7 +845,11 @@ def _get_callable(self, tensor=None, bcs=None, mat_type=None, sub_mat_type=None) mat_type = mat_type or "matfree" if self.rank == 1: - f = tensor or self._get_tensor(mat_type) + f = tensor or Function(self.ufl_interpolate.function_space()) + if tensor is None and self.access in {op2.MIN, op2.MAX}: + finfo = numpy.finfo(f.dat.dtype) + value = finfo.max if self.access == op2.MIN else finfo.min + f.assign(Constant(value)) self.mat = self._build_python_mat(_get_mtype(f.dat)[0]) if self.ufl_interpolate.is_adjoint: assert isinstance(self.dual_arg, Cofunction) @@ -978,261 +968,36 @@ def _allowed_mat_types(self): return {"aij", "baij", "matfree", None} -@known_pyop2_safe -def _build_interpolation_callables( - expr: Interpolate | ZeroBaseForm, - tensor: op2.Dat | op2.Mat | op2.Global, - access: Literal[op2.WRITE, op2.MIN, op2.MAX, op2.INC], - subset: op2.Subset | None = None, - bcs: Iterable[DirichletBC] | None = None -) -> tuple[Callable, ...]: - """Return a tuple of callables which calculate the interpolation. +def get_assembly_entity_node_map(fs: WithGeometry, target_mesh: MeshGeometry, + integral_type: str = "cell", + subdomain_id: str | int = "everywhere", + all_integer_subdomain_ids: dict | None = None) -> op2.Map | None: + """Return the map between entities of the target mesh and nodes of the function space. - Parameters - ---------- - expr : ufl.Interpolate | ufl.ZeroBaseForm - The symbolic interpolation expression, or a ZeroBaseForm. ZeroBaseForms - are simplified here to avoid code generation when access is WRITE or INC. - tensor : op2.Dat | op2.Mat | op2.Global - Object to hold the result of the interpolation. - access : Literal[op2.WRITE, op2.MIN, op2.MAX, op2.INC] - op2 access descriptor - subset : op2.Subset | None - An optional subset to apply the interpolation over, by default None. - bcs : Iterable[DirichletBC] | None - An optional list of boundary conditions to zero-out in the - output function space. Interpolator rows or columns which are - associated with boundary condition nodes are zeroed out when this is - specified. By default None, by default None. - - Returns - ------- - tuple[Callable, ...] - Tuple of callables which perform the interpolation. + If the function space is not defined on the target mesh then its node map is + composed with a map between target and source cells. """ - if isinstance(expr, ZeroBaseForm): - # Zero simplification, avoid code-generation - if access is op2.INC: - return () - elif access is op2.WRITE: - return (partial(tensor.zero, subset=subset),) - # Unclear how to avoid codegen for MIN and MAX - # Reconstruct the expression as an Interpolate - V = expr.arguments()[-1].function_space().dual() - expr = interpolate(zero(V.value_shape), V) - - if not isinstance(expr, Interpolate): - raise ValueError("Expecting to interpolate a symbolic Interpolate expression.") - - dual_arg, operand = expr.argument_slots() - assert isinstance(dual_arg, Cofunction | Coargument) - V = dual_arg.function_space().dual() - - if access is op2.READ: - raise ValueError("Can't have READ access for output function") - - # NOTE: The par_loop is always over the target mesh cells. - target_mesh = V.mesh() - source_mesh = extract_unique_domain(operand) or target_mesh - target_element = V.ufl_element() - if isinstance(target_mesh.topology, VertexOnlyMeshTopology): - # For interpolation onto a VOM, we use a FInAT QuadratureElement as the - # target element with runtime point set expressions as their - # quadrature rule point set. - rt_var_name = "rt_X" - target_element = runtime_quadrature_element(source_mesh, target_element, - rt_var_name=rt_var_name) - - cell_set = target_mesh.cell_set - if subset is not None: - assert subset.superset == cell_set - cell_set = subset - - parameters = {} - parameters['scalar_type'] = ScalarType - - copyin = () - copyout = () - - # For the matfree adjoint 1-form and the 0-form, the cellwise kernel will add multiple - # contributions from the facet DOFs of the dual argument. - # The incoming Cofunction needs to be weighted by the reciprocal of the DOF multiplicity. - if isinstance(dual_arg, Cofunction) and not create_element(target_element).is_dg(): - # Create a buffer for the weighted Cofunction - W = dual_arg.function_space() - v = Function(W) - expr = expr._ufl_expr_reconstruct_(operand, v=v) - copyin += (partial(dual_arg.dat.copy, v.dat),) - - # Compute the reciprocal of the DOF multiplicity - wdat = W.make_dat() - m_ = get_interp_node_map(source_mesh, target_mesh, W) - wsize = W.finat_element.space_dimension() * W.block_size - kernel_code = f""" - void multiplicity(PetscScalar *restrict w) {{ - for (PetscInt i=0; i<{wsize}; i++) w[i] += 1; - }}""" - kernel = op2.Kernel(kernel_code, "multiplicity") - op2.par_loop(kernel, cell_set, wdat(op2.INC, m_)) - with wdat.vec as w: - w.reciprocal() - - # Create a callable to apply the weight - with wdat.vec_ro as w, v.dat.vec as y: - copyin += (partial(y.pointwiseMult, y, w),) - - kernel = compile_expression(cell_set.comm, expr, target_element, - domain=source_mesh, parameters=parameters) - ast = kernel.ast - oriented = kernel.oriented - needs_cell_sizes = kernel.needs_cell_sizes - coefficient_numbers = kernel.coefficient_numbers - needs_external_coords = kernel.needs_external_coords - name = kernel.name - kernel = op2.Kernel(ast, name, requires_zeroed_output_arguments=(access is not op2.INC), - flop_count=kernel.flop_count, events=(kernel.event,)) - - parloop_args = [kernel, cell_set] - - coefficients = extract_numbered_coefficients(expr, coefficient_numbers) - if needs_external_coords: - coefficients = [source_mesh.coordinates] + coefficients - - if any(c.dat == tensor for c in coefficients): - output = tensor - tensor = op2.Dat(tensor.dataset) - if access is not op2.WRITE: - copyin += (partial(output.copy, tensor), ) - copyout += (partial(tensor.copy, output), ) - - arguments = expr.arguments() - if isinstance(tensor, op2.Global): - parloop_args.append(tensor(access)) - elif isinstance(tensor, op2.Dat): - V_dest = arguments[-1].function_space() - m_ = get_interp_node_map(source_mesh, target_mesh, V_dest) - parloop_args.append(tensor(access, m_)) - else: - assert access == op2.WRITE # Other access descriptors not done for Matrices. - Vrow = arguments[0].function_space() - Vcol = arguments[1].function_space() - assert tensor.handle.getSize() == (Vrow.dim(), Vcol.dim()) - rows_map = get_interp_node_map(source_mesh, target_mesh, Vrow) - columns_map = get_interp_node_map(source_mesh, target_mesh, Vcol) - lgmaps = None - if bcs: - if is_dual(Vrow): - Vrow = Vrow.dual() - if is_dual(Vcol): - Vcol = Vcol.dual() - bc_rows = [bc for bc in bcs if bc.function_space() == Vrow] - bc_cols = [bc for bc in bcs if bc.function_space() == Vcol] - lgmaps = [(Vrow.local_to_global_map(bc_rows), Vcol.local_to_global_map(bc_cols))] - parloop_args.append(tensor(access, (rows_map, columns_map), lgmaps=lgmaps)) - - if oriented: - co = source_mesh.cell_orientations() - parloop_args.append(co.dat(op2.READ, co.cell_node_map())) - - if needs_cell_sizes: - cs = source_mesh.cell_sizes - parloop_args.append(cs.dat(op2.READ, cs.cell_node_map())) - - for coefficient in coefficients: - m_ = get_interp_node_map(source_mesh, target_mesh, coefficient.function_space()) - parloop_args.append(coefficient.dat(op2.READ, m_)) - - for const in extract_firedrake_constants(expr): - parloop_args.append(const.dat(op2.READ)) - - # Finally, add the target mesh reference coordinates if they appear in the kernel if isinstance(target_mesh.topology, VertexOnlyMeshTopology): - if target_mesh is not source_mesh: - # NOTE: TSFC will sometimes drop run-time arguments in generated - # kernels if they are deemed not-necessary. - # FIXME: Checking for argument name in the inner kernel to decide - # whether to add an extra coefficient is a stopgap until - # compile_expression_dual_evaluation - # (a) outputs a coefficient map to indicate argument ordering in - # parloops as `compile_form` does and - # (b) allows the dual evaluation related coefficients to be supplied to - # them rather than having to be added post-hoc (likely by - # replacing `to_element` with a CoFunction/CoArgument as the - # target `dual` which would contain `dual` related - # coefficient(s)) - if any(arg.name == rt_var_name for arg in kernel.code[name].args): - # Add the coordinates of the target mesh quadrature points in the - # source mesh's reference cell as an extra argument for the inner - # loop. (With a vertex only mesh this is a single point for each - # vertex cell.) - target_ref_coords = target_mesh.reference_coordinates - m_ = target_ref_coords.cell_node_map() - parloop_args.append(target_ref_coords.dat(op2.READ, m_)) - - parloop = op2.ParLoop(*parloop_args) - if isinstance(tensor, op2.Mat): - return parloop, tensor.assemble - else: - return copyin + (parloop, ) + copyout - - -def get_interp_node_map(source_mesh: MeshGeometry, target_mesh: MeshGeometry, fs: WithGeometry) -> op2.Map | None: - """Return the map between cells of the target mesh and nodes of the function space. - - If the function space is defined on the source mesh then the node map is composed - with a map between target and source cells. - """ - if isinstance(target_mesh.topology, VertexOnlyMeshTopology): - coeff_mesh = fs.mesh() + source_mesh = fs.mesh() m_ = fs.cell_node_map() - if coeff_mesh is target_mesh or not coeff_mesh: - # NOTE: coeff_mesh is None is allowed e.g. when interpolating from - # a Real space - pass - elif coeff_mesh is source_mesh: - if m_: - # Since the par_loop is over the target mesh cells we need to - # compose a map that takes us from target mesh cells to the - # function space nodes on the source mesh. - if source_mesh.extruded: - # ExtrudedSet cannot be a map target so we need to build - # this ourselves - m_ = vom_cell_parent_node_map_extruded(target_mesh, m_) - else: - m_ = compose_map_and_cache(target_mesh.cell_parent_cell_map, m_) + # NOTE: source_mesh and m_ are None when interpolating from a Real + # space, in the trans-mesh case too. + if source_mesh is not target_mesh and source_mesh and m_: + # Since the par_loop is over the target mesh cells we need to + # compose a map that takes us from target mesh cells to the + # function space nodes on the source mesh. + if source_mesh.extruded: + # ExtrudedSet cannot be a map target so we need to build + # this ourselves + m_ = vom_cell_parent_node_map_extruded(target_mesh, m_) else: - # m_ is allowed to be None when interpolating from a Real space, - # even in the trans-mesh case. - pass - else: - raise ValueError("Have coefficient with unexpected mesh") + m_ = compose_map_and_cache(target_mesh.cell_parent_cell_map, m_) else: - m_ = fs.entity_node_map(target_mesh.topology, "cell", "everywhere", None) + m_ = fs.entity_node_map(target_mesh.topology, integral_type, subdomain_id, + all_integer_subdomain_ids) return m_ -try: - _expr_cachedir = os.environ["FIREDRAKE_TSFC_KERNEL_CACHE_DIR"] -except KeyError: - _expr_cachedir = os.path.join(tempfile.gettempdir(), - f"firedrake-tsfc-expression-kernel-cache-uid{os.getuid()}") - - -def _compile_expression_key(comm, expr, ufl_element, domain, parameters) -> tuple[Hashable, ...]: - """Generate a cache key suitable for :func:`tsfc.compile_expression_dual_evaluation`.""" - dual_arg, operand = expr.argument_slots() - return (hash_expr(operand), type(dual_arg), hash(ufl_element), tuplify(parameters)) - - -@memory_and_disk_cache( - hashkey=_compile_expression_key, - cachedir=_cachedir -) -@PETSc.Log.EventDecorator() -def compile_expression(comm, *args, **kwargs): - return compile_expression_dual_evaluation(*args, **kwargs) - - def compose_map_and_cache(map1: op2.Map, map2: op2.Map | None) -> op2.ComposedMap | None: """ Retrieve a :class:`pyop2.ComposedMap` map from the cache of map1 @@ -1658,8 +1423,8 @@ def _get_sub_interpolators( # Get sub-interpolators and sub-bcs for each block Isub: dict[tuple[int] | tuple[int, int], tuple[Interpolator, list[DirichletBC]]] = {} - for indices, form in split_form(self.ufl_interpolate): - if isinstance(form, ZeroBaseForm): + for indices, block in split_form(self.ufl_interpolate): + if isinstance(block, ZeroBaseForm): # Ensure block sparsity continue sub_bcs = [] @@ -1668,8 +1433,8 @@ def _get_sub_interpolators( sub_bcs.extend(bc for bc in bcs if space_equals(bc.function_space(), subspace)) if needs_action: # Take the action of each sub-cofunction against each block - form = action(form, dual_split[indices[-1:]]) - Isub[indices] = (get_interpolator(form), sub_bcs) + block = action(block, dual_split[indices[-1:]]) + Isub[indices] = (get_interpolator(block), sub_bcs) return Isub diff --git a/firedrake/mg/interface.py b/firedrake/mg/interface.py index 316ed270c8..47674ca6e0 100644 --- a/firedrake/mg/interface.py +++ b/firedrake/mg/interface.py @@ -296,14 +296,6 @@ def inject(fine, coarse): return coarse -def _bc_matches_space(bc, V): - """Return whether a boundary condition is defined on (a subspace of) V.""" - fs = bc.function_space() - while fs.component is not None and fs.parent is not None: - fs = fs.parent - return fs == V - - @PETSc.Log.EventDecorator() def assemble_prolongation_aij(Vc, Vf, bcs=None): """Assemble the explicit AIJ matrix prolonging Vc to Vf. @@ -354,8 +346,8 @@ def assemble_prolongation_aij(Vc, Vf, bcs=None): lgmaps = None if bcs: - row_bcs = [bc for bc in bcs if _bc_matches_space(bc, Vrow)] - col_bcs = [bc for bc in bcs if _bc_matches_space(bc, Vcol)] + row_bcs = [bc for bc in bcs if bc.parent_function_space.topological == Vrow.topological] + col_bcs = [bc for bc in bcs if bc.parent_function_space.topological == Vcol.topological] if row_bcs or col_bcs: lgmaps = [(Vrow.local_to_global_map(row_bcs), Vcol.local_to_global_map(col_bcs))] @@ -383,7 +375,8 @@ def assemble_prolongation_aij(Vc, Vf, bcs=None): if needs_quadrature: interp = interpolate(ufl_expr.TrialFunction(Vf), Vtarget) - Q = assemble(interp, bcs=bcs, mat_type="aij").petscmat + target_bcs = [bc for bc in bcs or () if bc.parent_function_space.topological == Vtarget.topological] + Q = assemble(interp, bcs=target_bcs, mat_type="aij").petscmat result = Q.matMult(result) return AssembledMatrix(arguments, result, bcs=bcs) diff --git a/firedrake/mg/kernels.py b/firedrake/mg/kernels.py index a897c5a621..e1f415965e 100644 --- a/firedrake/mg/kernels.py +++ b/firedrake/mg/kernels.py @@ -2,7 +2,6 @@ import string from collections import defaultdict from pyop2 import op2 -from pyop2.utils import as_tuple from firedrake.utils import IntType, as_cstr, complex_mode, ScalarType from firedrake.functionspacedata import entity_dofs_key from firedrake.functionspaceimpl import FiredrakeDualSpace @@ -20,6 +19,7 @@ import ufl import tsfc +from tsfc import kernel_args import tsfc.kernel_interface.firedrake_loopy as firedrake_interface @@ -34,7 +34,7 @@ from finat.quadrature import make_quadrature from firedrake.pointquery_utils import dX_norm_square, X_isub_dX, init_X, inside_check, is_affine, celldist_l1_c_expr from firedrake.pointquery_utils import to_reference_coords_newton_step as to_reference_coords_newton_step_body -from firedrake.pointeval_utils import runtime_quadrature_element +from tsfc.ufl_utils import runtime_quadrature_element def to_reference_coordinates(ufl_coordinate_element, parameters=None): @@ -131,22 +131,38 @@ def dual_evaluation_kernel(operand, dual_arg, parameters=None, return kernel -def _make_kernel_args(kernel, element, *args): +def _make_kernel_args(kernel, output, coefficient, target_coordinates, + coordinates=None, cell_orientations=None, cell_sizes=None): """Returns a string of argument names to call the kernel. Discards coordinate arguments if they do not appear in the kernel.""" - # NOTE: TSFC will sometimes drop run-time arguments in generated - # kernels if they are deemed not-necessary. - # For further information, see the same note in interpolation.py. - mask = [True] * len(args) - # Drop source mesh quantities if they do not appear in the kernel. - mask[1] = kernel.oriented - mask[2] = kernel.needs_cell_sizes - mask[3] = kernel.needs_external_coords - # Drop the target coordinates if the element is constant. - is_constant = sum(as_tuple(element.degree)) == 0 and not element.complex.is_macrocell() - mask[-1] = not is_constant - kernel_args = ", ".join(arg for arg, include in zip(args, mask) if include) - return kernel_args + # TSFC may omit runtime arguments that it determines are unnecessary from + # generated kernels. Iterate over the generated kernel's actual arguments + # so that the call includes only values that the kernel expects. + coefficients = (coefficient,) if isinstance(coefficient, str) else tuple(coefficient) + coefficient_index = 0 + args = [] + for arg in kernel.arguments: + if isinstance(arg, kernel_args.OutputKernelArg): + value = output + elif isinstance(arg, kernel_args.CoordinatesKernelArg): + value = coordinates + elif isinstance(arg, kernel_args.CellOrientationsKernelArg): + value = cell_orientations + elif isinstance(arg, kernel_args.CellSizesKernelArg): + value = cell_sizes + elif isinstance(arg, kernel_args.CoefficientKernelArg): + value = coefficients[coefficient_index] + coefficient_index += 1 + elif isinstance(arg, kernel_args.TabulationKernelArg): + value = target_coordinates + else: + raise ValueError(f"Unsupported dual-evaluation kernel argument {type(arg).__name__}") + if value is None: + raise ValueError(f"Missing value for dual-evaluation kernel argument {type(arg).__name__}") + args.append(value) + if coefficient_index != len(coefficients): + raise ValueError("Too many coefficient arguments for dual-evaluation kernel") + return ", ".join(args) def _make_element_key(element): @@ -240,7 +256,8 @@ def prolong_kernel(expression, Vf): "evaluate": evaluate_code, "cell_orient": ", const PetscScalar *co" if kernel.oriented else "", "cell_sizes": ", const PetscScalar *cs" if kernel.needs_cell_sizes else "", - "kernel_args": _make_kernel_args(kernel, element, "R", "co+cell", f"cs+cell*{num_verts}", "Xci", "fi", "Xref"), + "kernel_args": _make_kernel_args(kernel, "R", "fi", "Xref", coordinates="Xci", + cell_orientations="co+cell", cell_sizes=f"cs+cell*{num_verts}"), "ncandidate": ncandidate, "Rdim": Vf.block_size, "inside_cell": inside_check(element.cell, eps=1e-8, X="Xref"), @@ -367,7 +384,8 @@ def prolong_matrix_kernel(Vc, Vf): "evaluate": evaluate_code, "cell_orient": ", const PetscScalar *co" if kernel.oriented else "", "cell_sizes": ", const PetscScalar *cs" if kernel.needs_cell_sizes else "", - "kernel_args": _make_kernel_args(kernel, element, "B", "co+cell", f"cs+cell*{num_verts}", "Xci", "Xref"), + "kernel_args": _make_kernel_args(kernel, "B", (), "Xref", coordinates="Xci", + cell_orientations="co+cell", cell_sizes=f"cs+cell*{num_verts}"), "ncandidate": ncandidate, "row_dim": row_dim, "source_cell_inc": source_cell_inc, @@ -462,7 +480,8 @@ def restrict_kernel(Vf, Vc): "evaluate": evaluate_code, "cell_orient": ", const PetscScalar *co" if kernel.oriented else "", "cell_sizes": ", const PetscScalar *cs" if kernel.needs_cell_sizes else "", - "kernel_args": _make_kernel_args(kernel, element, "Ri", "co+cell", f"cs+cell*{num_verts}", "Xc", "b", "Xref"), + "kernel_args": _make_kernel_args(kernel, "Ri", "b", "Xref", coordinates="Xc", + cell_orientations="co+cell", cell_sizes=f"cs+cell*{num_verts}"), "ncandidate": ncandidate, "inside_cell": inside_check(element.cell, eps=1e-8, X="Xref"), "celldist_l1_c_expr": celldist_l1_c_expr(element.cell, X="Xref"), diff --git a/firedrake/pointeval_utils.py b/firedrake/pointeval_utils.py index da8bac0d03..42ba96f618 100644 --- a/firedrake/pointeval_utils.py +++ b/firedrake/pointeval_utils.py @@ -1,10 +1,7 @@ import loopy as lp from firedrake.utils import IntType, as_cstr -from finat.element_factory import as_fiat_cell -from finat.point_set import UnknownPointSet -from finat.quadrature import QuadratureRule -from finat.ufl import MixedElement, FiniteElement, TensorElement +from finat.ufl import MixedElement from ufl.corealg.map_dag import map_expr_dags from ufl.algorithms import extract_arguments, extract_coefficients @@ -17,44 +14,12 @@ import tsfc.kernel_interface.firedrake_loopy as firedrake_interface from tsfc.loopy import generate as generate_loopy from tsfc.parameters import default_parameters +from tsfc.ufl_utils import runtime_quadrature_element # noqa: F401 from firedrake import utils from firedrake.petsc import PETSc -def runtime_quadrature_element(domain, ufl_element, rt_var_name="rt_X"): - """Construct a Quadrature FiniteElement for interpolation onto a - VertexOnlyMesh. The quadrature point is an UnknownPointSet of shape - (1, tdim) where tdim is the topological dimension of domain.ufl_cell(). The - weight is [1.0], since the single local dof in the VertexOnlyMesh function - space corresponds to a point evaluation at the vertex. - - Parameters - ---------- - domain : ufl.AbstractDomain - The source domain. - ufl_element : finat.ufl.finiteelement.FiniteElement - The UFL element of the target FunctionSpace. - rt_var_name : str - String beginning with ``'rt_'`` which is used as the name of the - gem.Variable used to represent the UnknownPointSet. The ``'rt_'`` prefix - forces TSFC to do runtime tabulation. - """ - assert rt_var_name.startswith("rt_") - - cell = domain.ufl_cell() - point_expr = gem.Variable(rt_var_name, (1, cell.topological_dimension)) - point_set = UnknownPointSet(point_expr) - rule = QuadratureRule(point_set, weights=[1.0], ref_el=as_fiat_cell(cell)) - - shape = ufl_element.pullback.physical_value_shape(ufl_element, domain) - rt_element = FiniteElement("Quadrature", cell=cell, degree=0, quad_scheme=rule) - if shape: - symmetry = None if len(shape) < 2 else ufl_element.symmetry() - rt_element = TensorElement(rt_element, shape=shape, symmetry=symmetry) - return rt_element - - @PETSc.Log.EventDecorator() def compile_element(expression, coordinates, parameters=None): """Generates C code for point evaluations. diff --git a/firedrake/supermeshing.py b/firedrake/supermeshing.py index ff3ff04db8..a5cc7088cb 100644 --- a/firedrake/supermeshing.py +++ b/firedrake/supermeshing.py @@ -219,8 +219,6 @@ def likely(cell_A): kernel_A = dual_evaluation_kernel(ufl.Coefficient(V_A), ufl.TestFunction(V_S_A.dual()), name="evaluate_kernel_A") kernel_B = dual_evaluation_kernel(ufl.Coefficient(V_B), ufl.TestFunction(V_S_B.dual()), name="evaluate_kernel_B") kernel_S = dual_evaluation_kernel(ufl.Coefficient(V_S), ufl.TestFunction(V_S.dual()), name="evaluate_kernel_S") - dummy_args = ["dummy_place_holder"] * 3 - M_SS = assemble(inner(TrialFunction(V_S_A), TestFunction(V_S_B)) * dx) M_SS = M_SS.petscmat[:, :] node_locations_A = utils.physical_node_locations(V_S_A).dat.data_ro_with_halos @@ -449,9 +447,9 @@ def likely(cell_A): "evaluate_S": generate_code_v2(kernel_S.ast).device_code(), "evaluate_A": generate_code_v2(kernel_A.ast).device_code(), "evaluate_B": generate_code_v2(kernel_B.ast).device_code(), - "kernel_args_S": _make_kernel_args(kernel_S, V_S.finat_element, "physical_node_location", *dummy_args, "simplex_S", "reference_node_location"), - "kernel_args_A": _make_kernel_args(kernel_A, V_A.finat_element, "&R_AS[i][j]", *dummy_args, "coeffs_A", "reference_nodes_A[j]"), - "kernel_args_B": _make_kernel_args(kernel_B, V_B.finat_element, "&R_BS[i][j]", *dummy_args, "coeffs_B", "reference_nodes_B[j]"), + "kernel_args_S": _make_kernel_args(kernel_S, "physical_node_location", "simplex_S", "reference_node_location"), + "kernel_args_A": _make_kernel_args(kernel_A, "&R_AS[i][j]", "coeffs_A", "reference_nodes_A[j]"), + "kernel_args_B": _make_kernel_args(kernel_B, "&R_BS[i][j]", "coeffs_B", "reference_nodes_B[j]"), "to_reference": str(to_reference_kernel), "num_nodes_A": num_nodes_A, "num_nodes_B": num_nodes_B, diff --git a/firedrake/tsfc_interface.py b/firedrake/tsfc_interface.py index cde7f678a1..cd6e9eaa78 100644 --- a/firedrake/tsfc_interface.py +++ b/firedrake/tsfc_interface.py @@ -84,6 +84,7 @@ def __init__( coefficient_numbers, constant_numbers, dont_split_numbers, + access=op2.INC, diagonal=False ): """A wrapper object for one or more TSFC kernels compiled from a given :class:`~ufl.classes.Form`. @@ -131,6 +132,7 @@ def __init__( events = (kernel.event,) pyop2_kernel = as_pyop2_local_kernel(kernel.ast, kernel.name, len(kernel.arguments), + access=access, flop_count=kernel.flop_count, events=events) kernels.append(KernelInfo(kernel=pyop2_kernel, @@ -150,7 +152,7 @@ def __init__( SplitKernel = collections.namedtuple("SplitKernel", ["indices", "kinfo"]) -def _compile_form_hashkey(form, name, parameters=None, split=True, dont_split=(), diagonal=False): +def _compile_form_hashkey(form, name, parameters=None, split=True, dont_split=(), diagonal=False, access=op2.INC): return ( form.signature(), name, @@ -158,6 +160,7 @@ def _compile_form_hashkey(form, name, parameters=None, split=True, dont_split=() split, _make_dont_split_numbers(dont_split, form), diagonal, + access, ) @@ -168,7 +171,7 @@ def _compile_form_hashkey(form, name, parameters=None, split=True, dont_split=() cachedir=_cachedir ) @PETSc.Log.EventDecorator() -def compile_form(form, name, parameters=None, split=True, dont_split=(), diagonal=False): +def compile_form(form, name, parameters=None, split=True, dont_split=(), diagonal=False, access=op2.INC): """Compile a form using TSFC. Parameters @@ -188,6 +191,8 @@ def compile_form(form, name, parameters=None, split=True, dont_split=(), diagona Coefficients that are not to be split into components by form compiler. diagonal : bool If assembling a matrix is it diagonal? + access : pyop2.types.access.Access + Access mode for the output tensor. Returns ------- @@ -207,7 +212,7 @@ def compile_form(form, name, parameters=None, split=True, dont_split=(), diagona """ # Check that we get a Form - if not isinstance(form, Form): + if not isinstance(form, (Form, ufl.Interpolate)): raise RuntimeError("Unable to convert object to a UFL form: %s" % repr(form)) if parameters is None: @@ -256,6 +261,7 @@ def compile_form(form, name, parameters=None, split=True, dont_split=(), diagona coefficient_numbers, constant_numbers, dont_split_numbers, + access, diagonal, ) for kinfo in tsfc_kernel.kernels: @@ -269,7 +275,8 @@ def _real_mangle(form): """If the form contains arguments in the Real function space, replace these with literal 1 before passing to tsfc.""" a = form.arguments() - reals = [x.ufl_element().family() == "Real" for x in a] + # Real Coarguments are not integrated against, they are only contracted in Interpolate. + reals = [x.ufl_element().family() == "Real" and not isinstance(x, ufl.Coargument) for x in a] if not any(reals): return form replacements = {} diff --git a/firedrake/ufl_expr.py b/firedrake/ufl_expr.py index f71d111981..7503b820b6 100644 --- a/firedrake/ufl_expr.py +++ b/firedrake/ufl_expr.py @@ -391,6 +391,8 @@ def extract_domains(f): return list(set(mesh._meshes)) else: return [mesh] + elif isinstance(f, ufl.core.base_form_operator.BaseFormOperator): + return f.ufl_domains() elif isinstance(f, (ufl.form.FormSum, ufl.Action)): # ufl.domain.extract_domains does not work. if f._domains is None: diff --git a/tests/firedrake/multigrid/test_embedded_transfer.py b/tests/firedrake/multigrid/test_embedded_transfer.py index ac6e9abeb0..6dd10d2192 100644 --- a/tests/firedrake/multigrid/test_embedded_transfer.py +++ b/tests/firedrake/multigrid/test_embedded_transfer.py @@ -98,6 +98,7 @@ def check_transfer(op, V): assert errornorm(expr(Vc), uc) < 1E-13 +@pytest.mark.parallel([1, 3]) @pytest.mark.parametrize("op", ["prolong", "restrict", "inject"]) def test_transfer(op, V): check_transfer(op, V) diff --git a/tests/firedrake/regression/test_adjoint_operators.py b/tests/firedrake/regression/test_adjoint_operators.py index 57faf80477..76331e9893 100644 --- a/tests/firedrake/regression/test_adjoint_operators.py +++ b/tests/firedrake/regression/test_adjoint_operators.py @@ -83,6 +83,20 @@ def test_interpolate_with_arguments(rg): assert taylor_test(rf, f, h) > 1.9 +@pytest.mark.skipcomplex +def test_interpolate_in_form(rg): + mesh = UnitSquareMesh(3, 3) + V = FunctionSpace(mesh, "CG", 1) + W = FunctionSpace(mesh, "DG", 0) + x, y = SpatialCoordinate(mesh) + f = Function(V).interpolate(x + 2 * y) + + J = assemble(interpolate(f, W) ** 2 * dx) + rf = ReducedFunctional(J, Control(f)) + + assert taylor_test(rf, f, rg.uniform(V)) > 1.9 + + @pytest.mark.skipcomplex # Taping for complex-valued 0-forms not yet done def test_interpolate_scalar_valued(rg): mesh = IntervalMesh(10, 0, 1) diff --git a/tests/firedrake/regression/test_bcs.py b/tests/firedrake/regression/test_bcs.py index 9e43ba805b..14462504a9 100644 --- a/tests/firedrake/regression/test_bcs.py +++ b/tests/firedrake/regression/test_bcs.py @@ -56,7 +56,7 @@ def test_assemble_bcs_wrong_fs(V, measure): u, v = TrialFunction(V), TestFunction(V) W = FunctionSpace(V.mesh(), "CG", 2) - with pytest.raises(RuntimeError): + with pytest.raises(MismatchingFunctionSpaceError): assemble(inner(u, v)*measure, bcs=[DirichletBC(W, 32, 1)]) @@ -65,7 +65,7 @@ def test_assemble_bcs_wrong_fs_interior(V): u, v = TrialFunction(V), TestFunction(V) W = FunctionSpace(V.mesh(), "CG", 2) n = FacetNormal(V.mesh()) - with pytest.raises(RuntimeError): + with pytest.raises(MismatchingFunctionSpaceError): assemble(inner(jump(u, n), jump(v, n))*dS, bcs=[DirichletBC(W, 32, 1)]) diff --git a/tests/firedrake/regression/test_interp_dual.py b/tests/firedrake/regression/test_interp_dual.py index b0ce971a95..bbf36be7be 100644 --- a/tests/firedrake/regression/test_interp_dual.py +++ b/tests/firedrake/regression/test_interp_dual.py @@ -64,7 +64,36 @@ def test_assemble_interp_operator(V2, f1): assert np.allclose(a.dat.data, b.dat.data) -def test_assemble_interp_matrix(V1, V2, f1): +@pytest.fixture(params=("scalar", "vector", "mixed")) +def interpolation_matrix_case(request, mesh, V1, V2, f1): + if request.param == "scalar": + return V1, V2, f1, None + + elif request.param == "vector": + x, y = SpatialCoordinate(mesh) + expression = as_vector((x + 2*y, 2*x - y)) + source = VectorFunctionSpace(mesh, "CG", 1, dim=2) + target = VectorFunctionSpace(mesh, "DG", 1, dim=2) + return source, target, Function(source).interpolate(expression), None + + elif request.param == "mixed": + x, y = SpatialCoordinate(mesh) + expression = as_vector((x + 2*y, 2*x - y)) + X = VectorFunctionSpace(mesh, "CG", 2) + Y = VectorFunctionSpace(mesh, "DG", 1) + V = FunctionSpace(mesh, "CG", 1) + source = X * V + target = Y * V + function = Function(source) + function.sub(0).interpolate(expression) + function.sub(1).interpolate(x - y) + return source, target, function, "nest" + + +def test_form_interp_composition(interpolation_matrix_case): + from firedrake.assemble import ExplicitMatrixAssembler, get_assembler + + V1, V2, f1, mat_type = interpolation_matrix_case # -- I(v1, V2) -- # v1 = TrialFunction(V1) Iv1 = interpolate(v1, V2) @@ -75,7 +104,7 @@ def test_assemble_interp_matrix(V1, V2, f1): assert b.function_space() == V2 # Get the interpolation matrix - a = assemble(Iv1) + a = assemble(Iv1, mat_type=mat_type) assert a.arguments()[0].function_space() == V2.dual() assert a.arguments()[1].function_space() == V1 assert a.petscmat.getSize() == (V2.dim(), V1.dim()) @@ -84,7 +113,20 @@ def test_assemble_interp_matrix(V1, V2, f1): # and b the interpolation of f1 into V2. res = assemble(action(a, f1)) assert res.function_space() == V2 - assert np.allclose(res.dat.data, b.dat.data) + for result, expected in zip(res.subfunctions, b.subfunctions): + assert np.allclose(result.dat.data, expected.dat.data) + + v = TestFunction(V2) + form = inner(interpolate(TrialFunction(V1), V2), v) * dx + assembler = get_assembler(form, mat_type=mat_type) + assert isinstance(assembler, ExplicitMatrixAssembler) + operator = assembler.assemble() + actual = assemble(action(operator, f1)) + interpolated = assemble(interpolate(f1, V2)) + expected = assemble(inner(interpolated, v) * dx) + + for result, expected in zip(actual.subfunctions, expected.subfunctions): + assert np.allclose(result.dat.data, expected.dat.data) def test_assemble_interp_tlm(V1, V2, f1): @@ -131,6 +173,25 @@ def test_assemble_interp_adjoint_model(V1, V2): assert np.allclose(res.dat.data, Ivfstar.dat.data) +def test_adjoint_interp_direct_sum(): + # A restricted element is a direct sum, so it tabulates into a Concatenate + # that the dual argument has to contract one summand at a time. + mesh = UnitSquareMesh(2, 2, quadrilateral=True) + element = FiniteElement("Lagrange", mesh.ufl_cell(), 3) + Vf = FunctionSpace(mesh, RestrictedElement(element, "facet")) + Vc = FunctionSpace(mesh, RestrictedElement(element.reconstruct(degree=1), "facet")) + + uc = Function(Vc) + uc.dat.data_wo[...] = np.arange(1, 1 + uc.dat.data_ro.size) + rf = Cofunction(Vf.dual()) + rf.dat.data_wo[...] = np.arange(1, 1 + rf.dat.data_ro.size) + + # Restriction is the adjoint of prolongation. + uf = assemble(interpolate(uc, Vf)) + rc = assemble(interpolate(TestFunction(Vc), rf)) + assert np.isclose(assemble(action(rf, uf)), assemble(action(rc, uc))) + + def test_assemble_interp_adjoint_complex(mesh, V1, V2, f1): if complex_mode: f1 = Constant(3 - 5.j) * f1 @@ -395,3 +456,201 @@ def test_assemble_action_adjoint(V1, V2): assert isinstance(res4, Cofunction) assert res4.function_space() == V1.dual() assert np.allclose(res.dat.data, res4.dat.data) + + +@pytest.mark.parallel(2) +def test_form_interp_reuse(): + from firedrake.assemble import OneFormAssembler, get_assembler + + mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(mesh, "CG", 2) + W = FunctionSpace(mesh, "DG", 1) + x, y = SpatialCoordinate(mesh) + u = Function(V) + v = TestFunction(W) + interpolation = interpolate(u, W) + form = inner(interpolation, v) * dx + assembler = get_assembler(form) + + assert isinstance(assembler, OneFormAssembler) + for scale in (1, 3): + u.interpolate(scale * (x + y)) + actual = assembler.assemble() + expected = assemble(inner(assemble(interpolation), v) * dx) + assert np.allclose(actual.dat.data, expected.dat.data) + + +def test_form_interp_mapped_derivative(): + mesh = UnitSquareMesh(2, 2) + V = VectorFunctionSpace(mesh, "CG", 2) + W = FunctionSpace(mesh, "RT", 1) + x, y = SpatialCoordinate(mesh) + u = Function(V).interpolate(as_vector((x**2 + y, x - y**2))) + interpolation = interpolate(u, W) + + actual = assemble(inner(grad(interpolation), grad(interpolation)) * dx) + interpolated = assemble(interpolation) + expected = assemble(inner(grad(interpolated), grad(interpolated)) * dx) + assert np.isclose(actual, expected) + + +def test_form_interp_bilinear(): + mesh = UnitIntervalMesh(3) + V = FunctionSpace(mesh, "CG", 1) + W = FunctionSpace(mesh, "DG", 0) + u = TrialFunction(V) + v = TestFunction(W) + operator = assemble(inner(interpolate(u, W), v) * dx) + + x, = SpatialCoordinate(mesh) + f = Function(V).interpolate(x + 1) + actual = assemble(action(operator, f)) + expected = assemble(inner(assemble(interpolate(f, W)), v) * dx) + assert np.allclose(actual.dat.data, expected.dat.data) + + +@pytest.mark.parametrize("family", ("RTCE", "RTCF", "Q")) +def test_form_interp_direct_sum(family): + # A direct sum blocks its tabulation and its dual basis along the same + # summands, and each summand dual evaluates on points of its own. An + # interpolation into a nodal space is the identity on that space, so it + # catches a block that contracts against the wrong points. + mesh = UnitSquareMesh(2, 2, quadrilateral=True) + V = FunctionSpace(mesh, family, 2) + u = TrialFunction(V) + v = TestFunction(V) + expected = assemble(inner(u, v) * dx) + actual = assemble(inner(interpolate(u, V), v) * dx) + assert np.allclose(actual.M.values, expected.M.values) + + +def test_form_interp_interior_facet(): + mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(mesh, "CG", 2) + W = FunctionSpace(mesh, "DG", 1) + x, y = SpatialCoordinate(mesh) + u = Function(V).interpolate(x**2 + y) + v = TestFunction(W) + interpolation = interpolate(u, W) + + actual = assemble(jump(interpolation) * jump(v) * dS) + interpolated = assemble(interpolation) + expected = assemble(jump(interpolated) * jump(v) * dS) + assert np.allclose(actual.dat.data, expected.dat.data) + + +def test_form_interp_partial_fusion(): + from firedrake.assemble import BaseFormAssembler, get_assembler + + source_mesh = UnitSquareMesh(1, 1) + target_mesh = UnitSquareMesh(1, 1) + V = FunctionSpace(source_mesh, "CG", 1) + W = FunctionSpace(target_mesh, "CG", 1) + Q = FunctionSpace(target_mesh, "DG", 0) + v = TestFunction(Q) + + xs, ys = SpatialCoordinate(source_mesh) + u = Function(V).interpolate(xs + ys) + xt, yt = SpatialCoordinate(target_mesh) + w = Function(W).interpolate(xt * yt) + cross_mesh = interpolate(u, W) + same_mesh = interpolate(w, Q) + form = (inner(same_mesh, v) + inner(cross_mesh, v)) * dx(domain=target_mesh) + + assert isinstance(get_assembler(form), BaseFormAssembler) + # Only the cross-mesh interpolation is assembled on its own. + assert BaseFormAssembler.base_form_operands(form) == [cross_mesh] + + actual = assemble(form) + expected = assemble((inner(assemble(same_mesh), v) + + inner(assemble(cross_mesh), v)) * dx(domain=target_mesh)) + assert np.allclose(actual.dat.data, expected.dat.data) + + +def test_nested_interp_shares_kernel(): + from firedrake.assemble import BaseFormAssembler + + mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(mesh, "CG", 1) + W = FunctionSpace(mesh, "DG", 1) + Q = FunctionSpace(mesh, "CG", 2) + x, y = SpatialCoordinate(mesh) + f = Function(Q).interpolate(x*x + y) + v = TestFunction(W) + + interpolation = interpolate(f, V) + nested = interpolate(interpolation, W) + assert interpolation not in BaseFormAssembler.base_form_operands(nested) + + expected = assemble(interpolate(assemble(interpolation), W)) + assert np.allclose(assemble(nested).dat.data, expected.dat.data) + + actual = assemble(inner(nested, v) * dx) + assert np.allclose(actual.dat.data, assemble(inner(expected, v) * dx).dat.data) + + +def test_nested_cross_mesh_interp_assembles_operand(): + from firedrake.assemble import BaseFormAssembler + + source_mesh = UnitSquareMesh(3, 3) + target_mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(source_mesh, "CG", 2) + W = FunctionSpace(source_mesh, "CG", 1) + Q = FunctionSpace(target_mesh, "CG", 1) + u = Function(V).interpolate(SpatialCoordinate(source_mesh)[0]) + + # The operand is interpolated on the source mesh, so it cannot share the + # kernel of an interpolation that targets the other mesh. + interpolation = interpolate(u, W) + nested = interpolate(interpolation, Q) + assert interpolation in BaseFormAssembler.base_form_operands(nested) + + expected = assemble(interpolate(assemble(interpolation), Q)) + assert np.allclose(assemble(nested).dat.data, expected.dat.data) + + +def test_nested_submesh_interp_shares_kernel(): + from firedrake.assemble import BaseFormAssembler + + mesh = RectangleMesh(3, 1, 3., 1., quadrilateral=True) + x, y = SpatialCoordinate(mesh) + DG0 = FunctionSpace(mesh, "DG", 0) + left = Function(DG0).interpolate(conditional(x < 2., 1, 0)) + right = Function(DG0).interpolate(conditional(x > 1., 1, 0)) + mesh = RelabeledMesh(mesh, [left, right], [111, 222]) + left_mesh = Submesh(mesh, mesh.topological_dimension, 111) + right_mesh = Submesh(mesh, mesh.topological_dimension, 222) + V = FunctionSpace(left_mesh, "CG", 1) + W = FunctionSpace(right_mesh, "CG", 1) + x, y = SpatialCoordinate(left_mesh) + + interpolation = interpolate(x + y, V) + nested = interpolate(interpolation, W, allow_missing_dofs=True) + assert interpolation not in BaseFormAssembler.base_form_operands(nested) + + expected = assemble(interpolate(assemble(interpolation), W, allow_missing_dofs=True)) + assert np.allclose(assemble(nested).dat.data, expected.dat.data) + + +@pytest.mark.parametrize("target", ["same mesh", "submesh", "point cloud"]) +def test_tsfc_interp_accepts_related_domains(target): + """TSFC lowers an interpolation on the cells of its source mesh.""" + from tsfc import compile_form + + mesh = RectangleMesh(3, 1, 3., 1., quadrilateral=True) + x, y = SpatialCoordinate(mesh) + DG0 = FunctionSpace(mesh, "DG", 0) + left = Function(DG0).interpolate(conditional(x < 2., 1, 0)) + mesh = RelabeledMesh(mesh, [left], [111]) + + x, y = SpatialCoordinate(mesh) + f = Function(FunctionSpace(mesh, "CG", 2)).interpolate(x + y) + if target == "same mesh": + V = FunctionSpace(mesh, "CG", 1) + elif target == "submesh": + V = FunctionSpace(Submesh(mesh, mesh.topological_dimension, 111), "CG", 1) + else: + V = FunctionSpace(VertexOnlyMesh(mesh, [[0.5, 0.5], [2.5, 0.5]]), "DG", 0) + + kernel, = compile_form(interpolate(f, V, allow_missing_dofs=True), prefix="interp") + assert kernel.name == "interp_cell_integral" diff --git a/tests/firedrake/regression/test_interpolate.py b/tests/firedrake/regression/test_interpolate.py index 4539c20504..50c5cdf4ae 100644 --- a/tests/firedrake/regression/test_interpolate.py +++ b/tests/firedrake/regression/test_interpolate.py @@ -19,6 +19,18 @@ def test_constant(): assert np.allclose(1.0, f.dat.data) +def test_zero_expression(): + mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(mesh, "CG", 1) + x, y = SpatialCoordinate(mesh) + c = Constant(2.0) + + f = Function(V).assign(1) + f.interpolate((c * y).dx(0)) + + assert np.allclose(f.dat.data_ro, 0.0) + + def test_function(): m = UnitTriangleMesh() x = SpatialCoordinate(m) @@ -34,6 +46,21 @@ def test_function(): assert np.allclose(g.dat.data, h.dat.data) +@pytest.mark.parallel([1, 2]) +@pytest.mark.parametrize( + ("access", "initial", "expected"), + [(op2.INC, 2.0, 4.0), (op2.MIN, 3.0, 3.0), (op2.MAX, -3.0, -3.0)], + ids=["inc", "min", "max"], +) +def test_in_place_interpolation_preserves_reduction(access, initial, expected): + mesh = UnitSquareMesh(1, 1) + V = FunctionSpace(mesh, "DG", 0) + f = Function(V).assign(initial) + + assert f.interpolate(f, access=access) is f + assert np.allclose(f.dat.data_ro, expected) + + def test_mixed_expression(): m = UnitTriangleMesh() x = SpatialCoordinate(m) @@ -606,6 +633,57 @@ def test_mixed_matrix(mode, mat_type): assert np.allclose(x.dat.data, y.dat.data) +@pytest.mark.parallel([1, 2]) +@pytest.mark.parametrize("mode", ["forward", "adjoint"]) +def test_mixed_matrix_q_rtce(mode): + mesh = UnitSquareMesh(1, 1, quadrilateral=True) + source = FunctionSpace(mesh, "Q", 1) * FunctionSpace(mesh, "RTCE", 1) + target = FunctionSpace(mesh, "Q", 2) * FunctionSpace(mesh, "RTCE", 2) + + if mode == "forward": + I = Interpolate(TrialFunction(source), TestFunction(target.dual())) + a = assemble(I) + u = Function(source) + u.subfunctions[0].assign(1) + u.subfunctions[1].assign(2) + result_matfree = assemble(Interpolate(u, TestFunction(target.dual()))) + else: + I = Interpolate(TestFunction(source), TrialFunction(target.dual())) + a = assemble(I) + u = Cofunction(target.dual()) + u.subfunctions[0].assign(1) + u.subfunctions[1].assign(2) + result_matfree = assemble(Interpolate(TestFunction(source), u)) + + result_explicit = assemble(action(a, u)) + for x, y in zip(result_explicit.subfunctions, result_matfree.subfunctions): + assert np.allclose(x.dat.data, y.dat.data) + + +def test_mixed_matrix_direct_sum(): + mesh = UnitSquareMesh(3, 3, quadrilateral=True) + V1 = VectorFunctionSpace(mesh, "CG", 2) + V2 = FunctionSpace(mesh, "CG", 1) + # RTCF is a direct sum, so it tabulates into a Concatenate that only its + # own block of the dual argument can split. + V3 = FunctionSpace(mesh, "RTCF", 1) + V4 = FunctionSpace(mesh, "DG", 0) + + Z = V1 * V2 + W = V3 * V4 + + a = assemble(Interpolate(TestFunction(Z), TrialFunction(W.dual()))) + + u = Function(W.dual()) + u.subfunctions[0].assign(1) + u.subfunctions[1].assign(2) + + result_explicit = assemble(action(a, u)) + result_matfree = assemble(Interpolate(TestFunction(Z), u)) + for x, y in zip(result_explicit.subfunctions, result_matfree.subfunctions): + assert np.allclose(x.dat.data, y.dat.data) + + @pytest.mark.parallel(2) @pytest.mark.parametrize("mode", ["forward", "adjoint"]) @pytest.mark.parametrize("family,degree", [("CG", 1), ("DG", 0)]) @@ -641,6 +719,24 @@ def test_interpolator_reuse(family, degree, mode): assert np.allclose(result.dat.data, expected) +@pytest.mark.parallel([1, 3]) +def test_same_space_interp_bcs(): + mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(mesh, "CG", 1) + rg = RandomGenerator(PCG64(seed=123456789)) + w = rg.uniform(V) + + # Source and target agree, so the interpolation has a diagonal to carry + # the boundary rows, just as a Form on the same spaces does. + I = assemble(interpolate(2 * TrialFunction(V), V), bcs=[DirichletBC(V, 0, 1)]) + result = assemble(action(I, w)) + + expected = Function(V).assign(2 * w) + DirichletBC(V, w, 1).apply(expected) + + assert np.allclose(result.dat.data, expected.dat.data) + + def test_mixed_space_bcs(): mesh = UnitSquareMesh(2, 2) V = FunctionSpace(mesh, "CG", 1) diff --git a/tests/firedrake/regression/test_interpolation_manual.py b/tests/firedrake/regression/test_interpolation_manual.py index cedaee54be..3c32838a8d 100644 --- a/tests/firedrake/regression/test_interpolation_manual.py +++ b/tests/firedrake/regression/test_interpolation_manual.py @@ -260,7 +260,7 @@ def test_mixed_space_interpolation(): for j in range(2): sub_mat = I.petscmat.getNestSubMatrix(i, j) if i != j: - assert not sub_mat + assert sub_mat.norm() == 0.0 continue else: res_block = assemble(interpolate(TrialFunction(U.sub(j)), W.sub(i))) diff --git a/tests/firedrake/regression/test_interpolation_operators.py b/tests/firedrake/regression/test_interpolation_operators.py index 6c97264ded..fb4515b197 100644 --- a/tests/firedrake/regression/test_interpolation_operators.py +++ b/tests/firedrake/regression/test_interpolation_operators.py @@ -1,6 +1,6 @@ from firedrake import * from firedrake.interpolation import ( - MixedInterpolator, SameMeshInterpolator, CrossMeshInterpolator, + SameMeshInterpolator, CrossMeshInterpolator, get_interpolator, VomOntoVomInterpolator, ) from firedrake.matrix import ImplicitMatrix, Matrix @@ -73,9 +73,12 @@ def test_same_mesh_mattype(value_shape, mat_type, mode): res2 = assemble(action(adjoint(forward_I_mat), f)) assert np.allclose(res2.dat.data, exact.dat.data) - with pytest.raises(NotImplementedError): - # MatNest only implemented for interpolation between MixedFunctionSpaces - assemble(interp, mat_type="nest") + # A nest of the one block these unmixed spaces make collapses to that + # block's own type, just as it does for a Form. + nest_mat = assemble(interp, mat_type="nest") + assert nest_mat.petscmat.type == prefix + ("baij" if value_shape == "vector" else "aij") + res = assemble(action(nest_mat, f)) + assert np.allclose(res.dat.data, exact.dat.data) @pytest.mark.parametrize("value_shape", ["scalar", "vector"], ids=lambda v: f"fs_type={v}") @@ -188,7 +191,7 @@ def test_mixed_same_mesh_mattype(value_shape, mat_type, sub_mat_type): expr = as_vector([x**2, x**2, y**2, y**2]) interp = interpolate(TrialFunction(U), W) - assert isinstance(get_interpolator(interp), MixedInterpolator) + assert isinstance(get_interpolator(interp), SameMeshInterpolator) I_mat = assemble(interp, mat_type=mat_type, sub_mat_type=sub_mat_type) assert isinstance(I_mat, ImplicitMatrix if mat_type == "matfree" else Matrix) @@ -198,17 +201,19 @@ def test_mixed_same_mesh_mattype(value_shape, mat_type, sub_mat_type): assert I_mat.petscmat.type == "seqaij" else: assert I_mat.petscmat.type == "nest" + if value_shape == "scalar": + # Always seqaij for scalar + sub_type = "seqaij" + else: + # A blocked space makes matnest default to baij + sub_type = "seq" + (sub_mat_type if sub_mat_type else "baij") for (i, j) in [(0, 0), (0, 1), (1, 0), (1, 1)]: + # Every block is assembled, as it is for a Form. The components + # do not mix, so the off-diagonal ones assemble to zero. sub_mat = I_mat.petscmat.getNestSubMatrix(i, j) + assert sub_mat.type == sub_type if i != j: - assert not sub_mat - continue - if value_shape == "scalar": - # Always seqaij for scalar - assert sub_mat.type == "seqaij" - else: - # matnest sub_mat_type defaults to aij - assert sub_mat.type == "seq" + (sub_mat_type if sub_mat_type else "aij") + assert sub_mat.norm() == 0.0 f = Function(U).interpolate(expr) exact = Function(W).interpolate(expr) @@ -216,5 +221,6 @@ def test_mixed_same_mesh_mattype(value_shape, mat_type, sub_mat_type): for resi, exi in zip(res.subfunctions, exact.subfunctions): assert np.allclose(resi.dat.data, exi.dat.data) - with pytest.raises(NotImplementedError): + with pytest.raises(ValueError, match="BAIJ matrix type makes no sense"): + # A mixed space has no block structure to give BAIJ, as for a Form. assemble(interp, mat_type="baij") diff --git a/tests/firedrake/slate/test_assemble_tensors.py b/tests/firedrake/slate/test_assemble_tensors.py index 11f043375a..a166f57cca 100644 --- a/tests/firedrake/slate/test_assemble_tensors.py +++ b/tests/firedrake/slate/test_assemble_tensors.py @@ -367,3 +367,33 @@ def test_nested_block(mesh, degree): result = assemble(Block(Jp, (0, 0))).petscmat assert np.allclose(result[:, :], expect[:, :]) + + +def test_slate_tensor_interp(): + mesh = UnitSquareMesh(2, 2) + V = FunctionSpace(mesh, "CG", 2) + W = FunctionSpace(mesh, "DG", 1) + x, y = SpatialCoordinate(mesh) + u = Function(V).interpolate(x**2 + y) + w = TrialFunction(W) + v = TestFunction(W) + interpolation = interpolate(u, W) + + mass = Tensor(inner(w, v) * dx) + actual = assemble(mass.inv * Tensor(inner(interpolation, v) * dx)) + + interpolated = assemble(interpolation) + expected = assemble(mass.inv * Tensor(inner(interpolated, v) * dx)) + assert np.allclose(actual.dat.data, expected.dat.data) + + +def test_slate_tensor_nonfusible_interp(): + source_mesh = UnitSquareMesh(1, 1) + target_mesh = UnitSquareMesh(1, 1) + V = FunctionSpace(source_mesh, "CG", 1) + W = FunctionSpace(target_mesh, "CG", 1) + v = TestFunction(W) + form = inner(interpolate(Function(V), W), v) * dx(domain=target_mesh) + + with pytest.raises(NotImplementedError): + assemble(Tensor(form)) diff --git a/tests/firedrake/submesh/test_submesh_interpolate.py b/tests/firedrake/submesh/test_submesh_interpolate.py index c3b084a878..75f101efd6 100644 --- a/tests/firedrake/submesh/test_submesh_interpolate.py +++ b/tests/firedrake/submesh/test_submesh_interpolate.py @@ -57,6 +57,93 @@ def _test_submesh_interpolate_cell_cell(mesh, subdomain_cond, fe_fesub): assert assemble(inner(g - f, g - f) * dx(label_value)).real < 1e-14 +@pytest.mark.parallel([1, 3]) +def test_submesh_form_interp(): + from firedrake.assemble import OneFormAssembler, get_assembler + + mesh = UnitSquareMesh(4, 4) + x, y = SpatialCoordinate(mesh) + submesh = make_submesh(mesh, conditional(x < 0.51, 1, 0), 999) + V = FunctionSpace(mesh, "CG", 2) + W = FunctionSpace(submesh, "CG", 1) + f = Function(V).interpolate(x + 2*y) + xs, ys = SpatialCoordinate(submesh) + expected = Function(W).interpolate(xs + 2*ys) + + actual = assemble(interpolate(f, W)) + assert np.allclose(actual.dat.data_ro_with_halos, expected.dat.data_ro_with_halos) + + operator = assemble(interpolate(TrialFunction(V), W)) + actual = assemble(action(operator, f)) + assert np.allclose(actual.dat.data_ro_with_halos, expected.dat.data_ro_with_halos) + + v = TestFunction(W) + subdx = Measure("dx", submesh, intersect_measures=(Measure("dx", mesh),)) + form = inner(interpolate(f, W), v) * subdx + assembler = get_assembler(form) + assert isinstance(assembler, OneFormAssembler) + actual = assembler.assemble() + expected = assemble(inner(expected, v) * dx(submesh)) + assert np.allclose(actual.dat.data_ro, expected.dat.data_ro) + + +@pytest.mark.parallel([1, 3]) +@pytest.mark.parametrize('covered', [True, False]) +def test_submesh_form_interp_coverage(covered): + # A Form iterates over its own cells, so it only fuses an interpolation + # whose source mesh supplies every target cell. + from firedrake.assemble import BaseFormAssembler + + mesh = RectangleMesh(4, 2, 2., 1., quadrilateral=True) + x, y = SpatialCoordinate(mesh) + DG0 = FunctionSpace(mesh, "DG", 0) + left = Function(DG0).interpolate(conditional(x < 1., 1, 0)) + mesh = RelabeledMesh(mesh, [left], [111]) + subm = Submesh(mesh, mesh.topological_dimension, 111) + # The parent supplies every submesh cell, but the submesh supplies only + # those parent cells that it was cut from. + source, target = (mesh, subm) if covered else (subm, mesh) + Vsource = FunctionSpace(source, "CG", 1) + Vtarget = FunctionSpace(target, "CG", 1) + xs, ys = SpatialCoordinate(source) + f = Function(Vsource).interpolate(xs + 2 * ys) + v = TestFunction(Vtarget) + interp = interpolate(f, Vtarget, allow_missing_dofs=not covered) + measure = Measure("dx", target, intersect_measures=(Measure("dx", source),)) + form = inner(interp, v) * measure + + assert (interp not in BaseFormAssembler.base_form_operands(form)) == covered + actual = assemble(form) + expected = assemble(inner(assemble(interp), v) * dx(target)) + assert np.allclose(actual.dat.data_ro, expected.dat.data_ro) + + +@pytest.mark.parallel([1, 3]) +def test_submesh_form_interp_facet_trace(): + # An interpolation that crosses a codimension does not map cells to cells, + # so the Form that holds it assembles it on its own. + from firedrake.assemble import BaseFormAssembler + + mesh = UnitCubeMesh(4, 4, 4) + x, y, z = SpatialCoordinate(mesh) + trace = Function(FunctionSpace(mesh, "HDiv Trace", 0)) + trace.interpolate(conditional(x > .999, 1., 0.)) + mesh = RelabeledMesh(mesh, [trace], [999]) + subm = Submesh(mesh, mesh.topological_dimension - 1, 999) + V = FunctionSpace(mesh, "CG", 1) + T = FunctionSpace(subm, "CG", 1) + f = Function(V).interpolate(y + 2 * z) + v = TestFunction(T) + interp = interpolate(f, T) + measure = Measure("dx", subm, intersect_measures=(Measure("dx", mesh),)) + form = inner(interp, v) * measure + + assert interp in BaseFormAssembler.base_form_operands(form) + actual = assemble(form) + expected = assemble(inner(assemble(interp), v) * dx(subm)) + assert np.allclose(actual.dat.data_ro, expected.dat.data_ro) + + @pytest.mark.parametrize('nelem', [2, 4, 8, None]) @pytest.mark.parametrize('fe_fesub', [[("DQ", 0), ("DQ", 0)], [("Q", 4), ("Q", 5)]]) diff --git a/tests/tsfc/test_dual_evaluation.py b/tests/tsfc/test_dual_evaluation.py index a2aa366b47..07d2b8e70a 100644 --- a/tests/tsfc/test_dual_evaluation.py +++ b/tests/tsfc/test_dual_evaluation.py @@ -18,6 +18,19 @@ def test_ufl_only_simple(): assert kernel.needs_external_coords is False +def test_ufl_only_nested_interpolate(): + mesh = ufl.Mesh(finat.ufl.VectorElement("P", ufl.triangle, 1)) + V = ufl.FunctionSpace(mesh, finat.ufl.VectorElement("P", ufl.triangle, 2)) + W = ufl.FunctionSpace(mesh, finat.ufl.FiniteElement("RT", ufl.triangle, 1)) + X = ufl.FunctionSpace(mesh, finat.ufl.VectorElement("P", ufl.triangle, 2)) + v = ufl.Coefficient(V) + expression = ufl.Interpolate(ufl.Interpolate(v, W), X) + + kernel = compile_expression_dual_evaluation(expression, X.ufl_element()) + + assert kernel.needs_external_coords is True + + def test_ufl_only_spatialcoordinate(): mesh = ufl.Mesh(finat.ufl.VectorElement("P", ufl.triangle, 1)) V = ufl.FunctionSpace(mesh, finat.ufl.FiniteElement("P", ufl.triangle, 2)) @@ -75,7 +88,7 @@ def dual_argument_kernel(cell, degree, restriction=None): Returns ------- - ExpressionKernel + Kernel The compiled dual evaluation kernel. """ mesh = ufl.Mesh(finat.ufl.VectorElement("Q", cell, 1)) diff --git a/tests/tsfc/test_interpolation_factorisation.py b/tests/tsfc/test_interpolation_factorisation.py index c76753a890..a1809196f6 100644 --- a/tests/tsfc/test_interpolation_factorisation.py +++ b/tests/tsfc/test_interpolation_factorisation.py @@ -1,12 +1,13 @@ from functools import partial import numpy import pytest +import ufl -from ufl import (Mesh, FunctionSpace, Coefficient, +from ufl import (Mesh, MeshSequence, FunctionSpace, interval, quadrilateral, hexahedron) -from finat.ufl import FiniteElement, VectorElement, TensorElement +from finat.ufl import FiniteElement, VectorElement, TensorElement, MixedElement -from tsfc import compile_expression_dual_evaluation +from tsfc import compile_form @pytest.fixture(params=[interval, quadrilateral, hexahedron], @@ -25,38 +26,82 @@ def element(request, mesh): return partial(request.param, family, mesh.ufl_cell()) -def flop_count(mesh, source, target): - Vtarget = FunctionSpace(mesh, target) - Vsource = FunctionSpace(mesh, source) - expr = Coefficient(Vsource) - kernel = compile_expression_dual_evaluation(expr, Vtarget.ufl_element()) +def interpolate_expression(domain, source, target, dual): + Vsource = FunctionSpace(domain, source) + Vtarget = FunctionSpace(domain, target) + if dual: + return ufl.Interpolate( + ufl.Argument(Vsource, 0), ufl.Cofunction(Vtarget.dual()) + ) + return ufl.Interpolate( + ufl.Coefficient(Vsource), ufl.Coargument(Vtarget.dual(), 0) + ) + + +def interpolate_flop_count(domain, source, target, dual): + kernel, = compile_form( + interpolate_expression(domain, source, target, dual), + parameters={"mode": "spectral"}, + ) return kernel.flop_count -def test_sum_factorisation(mesh, element): +@pytest.mark.parametrize("dual", (False, True), ids=("primal", "dual")) +def test_sum_factorisation(mesh, element, dual): # Interpolation between sum factorisable elements should cost # O(p^{d+1}) - degrees = numpy.asarray([2**n - 1 for n in range(2, 9)]) + degrees = numpy.asarray([4, 8, 16]) flops = [] for lo, hi in zip(degrees - 1, degrees): - flops.append(flop_count(mesh, element(int(lo)), element(int(hi)))) + flops.append(interpolate_flop_count( + mesh, element(int(lo)), element(int(hi)), dual + )) flops = numpy.asarray(flops) rates = numpy.diff(numpy.log(flops)) / numpy.diff(numpy.log(degrees)) - assert (rates < (mesh.topological_dimension+1)).all() + assert (rates < mesh.topological_dimension + 1).all() -def test_sum_factorisation_scalar_tensor(mesh, element): +@pytest.mark.parametrize("dual", (False, True), ids=("primal", "dual")) +def test_sum_factorisation_scalar_tensor(mesh, element, dual): # Interpolation into tensor elements should cost value_shape # more than the equivalent scalar element. - degree = 2**7 - 1 + degree = 16 source = element(degree - 1) target = element(degree) - tensor_flops = flop_count(mesh, source, target) + tensor_flops = interpolate_flop_count(mesh, source, target, dual) expect = FunctionSpace(mesh, target).value_size if isinstance(target, FiniteElement): scalar_flops = tensor_flops else: target = target.sub_elements[0] source = source.sub_elements[0] - scalar_flops = flop_count(mesh, source, target) + scalar_flops = interpolate_flop_count(mesh, source, target, dual) assert numpy.allclose(tensor_flops / scalar_flops, expect, rtol=1e-2) + + +def q_rtce_elements(degree): + return (FiniteElement("Q", quadrilateral, degree), + FiniteElement("RTCE", quadrilateral, degree)) + + +@pytest.mark.parametrize("dual", (False, True), ids=("primal", "dual")) +def test_sum_factorisation_mixed_q_rtce(dual): + mesh = Mesh(VectorElement("Q", quadrilateral, 1)) + mixed_mesh = MeshSequence([mesh, mesh]) + degrees = numpy.asarray([4, 8, 16]) + mixed_flops = [] + component_flops = [] + for degree in degrees: + source = q_rtce_elements(int(degree - 1)) + target = q_rtce_elements(int(degree)) + mixed_flops.append(interpolate_flop_count( + mixed_mesh, MixedElement(*source), MixedElement(*target), dual + )) + component_flops.append(sum( + interpolate_flop_count(mesh, source_element, target_element, dual) + for source_element, target_element in zip(source, target, strict=True) + )) + + numpy.testing.assert_equal(mixed_flops, component_flops) + rates = numpy.diff(numpy.log(mixed_flops)) / numpy.diff(numpy.log(degrees)) + assert (rates < quadrilateral.topological_dimension + 1).all() diff --git a/tests/tsfc/test_sum_factorisation.py b/tests/tsfc/test_sum_factorisation.py index 85d9729e81..b66422d4a9 100644 --- a/tests/tsfc/test_sum_factorisation.py +++ b/tests/tsfc/test_sum_factorisation.py @@ -3,7 +3,8 @@ from ufl import (Mesh, FunctionSpace, TestFunction, TrialFunction, TensorProductCell, dx, action, interval, triangle, - quadrilateral, hexahedron, curl, dot, div, grad) + quadrilateral, hexahedron, curl, dot, div, grad, inner, + Interpolate) from finat.ufl import (FiniteElement, VectorElement, EnrichedElement, TensorProductElement, HCurlElement, HDivElement) @@ -64,6 +65,15 @@ def split_vector_laplace(cell, degree): return [dot(u, grad(tau))*dx, dot(grad(sigma), v)*dx, dot(curl(u), curl(v))*dx] +def interpolated_vector_laplace(cell, degree): + m = Mesh(VectorElement('CG', cell, 1)) + Q = FunctionSpace(m, FiniteElement('Q', cell, degree)) + NCE = FunctionSpace(m, FiniteElement('NCE', cell, degree)) + dtest = Interpolate(grad(TestFunction(Q)), NCE) + dtrial = TrialFunction(NCE) + return inner(dtrial, dtest) * dx + + def count_flops(form): kernel, = compile_form(form, parameters=dict(mode='spectral')) flops = kernel.flop_count @@ -190,6 +200,16 @@ def test_vector_laplace_action(cell, order): assert (rates < order).all() +@pytest.mark.parametrize(('cell', 'order'), + [(TensorProductCell(quadrilateral, interval), 7)]) +def test_interpolated_vector_laplace(cell, order): + degrees = numpy.arange(3, 8) + flops = [count_flops(interpolated_vector_laplace(cell, int(degree))) + for degree in degrees] + rates = numpy.diff(numpy.log(flops)) / numpy.diff(numpy.log(degrees)) + assert (rates < order).all() + + @pytest.mark.parametrize(('cell', 'equivalent_cell'), [(quadrilateral, TensorProductCell(interval, interval)), (hexahedron, TensorProductCell(quadrilateral, interval)), diff --git a/tsfc/driver.py b/tsfc/driver.py index b287f87e63..77c23b8103 100644 --- a/tsfc/driver.py +++ b/tsfc/driver.py @@ -1,31 +1,20 @@ import collections import time import sys -from itertools import chain -import numpy -from finat.physically_mapped import NeedsCoordinateMappingElement import ufl from ufl.algorithms import extract_coefficients -from ufl.algorithms.analysis import has_type -from ufl.algorithms.apply_coefficient_split import CoefficientSplitter -from ufl.classes import Form, GeometricQuantity -from ufl.domain import extract_unique_domain, extract_domains - -import gem -import gem.impero_utils as impero_utils -from gem.unconcatenate import unconcatenate +from ufl.algorithms.apply_coefficient_split import build_coefficient_split +from ufl.classes import Form +from ufl.domain import extract_unique_domain, extract_domains, join_domains import finat -from finat.element_factory import as_fiat_cell -from tsfc import fem, ufl_utils +from tsfc import ufl_utils from tsfc.logging import logger -from tsfc.modified_terminals import analyse_modified_terminal from tsfc.parameters import default_parameters, is_complex -from tsfc.ufl_utils import apply_mapping, extract_firedrake_constants, simplify_abs +from tsfc.ufl_utils import extract_firedrake_constants import tsfc.kernel_interface.firedrake_loopy as firedrake_interface_loopy -from tsfc.kernel_interface.common import get_index_ordering, pick_mode from tsfc.exceptions import MismatchingDomainError @@ -79,6 +68,13 @@ def compile_form(form, prefix="form", parameters=None, dont_split_numbers=(), di """ cpu_time = time.time() + if isinstance(form, ufl.Interpolate): + kernel = compile_expression_dual_evaluation( + form, form.ufl_element(), parameters=parameters, + name=f"{prefix}_cell_integral", + ) + return [] if kernel is None else [kernel] + assert isinstance(form, Form) GREEN = "\033[1;37;32m%s\033[0m" @@ -112,6 +108,27 @@ def compile_form(form, prefix="form", parameters=None, dont_split_numbers=(), di return kernels +def make_kernel_builder(integral_data_info: TSFCIntegralDataInfo, + constants: tuple, + parameters: dict, + diagonal: bool = False) -> firedrake_interface_loopy.KernelBuilder: + """Create a kernel builder holding every mesh quantity its integral may read. + + The caller sets the coordinates that it needs. + """ + builder = firedrake_interface_loopy.KernelBuilder( + integral_data_info, parameters["scalar_type"], diagonal=diagonal + ) + domains = tuple(integral_data_info.domain_integral_type_map) + builder.set_cell_orientations(domains) + builder.set_cell_sizes(domains) + builder.set_coefficients() + # TODO: We do not want pass constants to kernels that do not need them + # so we should attach the constants to integral data instead + builder.set_constants(constants) + return builder + + def compile_integral(integral_data, form_data, prefix, parameters, *, diagonal=False): """Compiles a UFL integral into an assembly kernel. @@ -123,7 +140,6 @@ def compile_integral(integral_data, form_data, prefix, parameters, *, diagonal=F :returns: a kernel constructed by the kernel interface """ parameters = preprocess_parameters(parameters) - scalar_type = parameters["scalar_type"] integral_type = integral_data.integral_type arguments = form_data.preprocessed_form.arguments() if integral_type.startswith("interior_facet") and diagonal and any(a.function_space().finat_element.is_dg() for a in arguments): @@ -158,21 +174,10 @@ def compile_integral(integral_data, form_data, prefix, parameters, *, diagonal=F coefficient_split=coefficient_split, coefficient_numbers=coefficient_numbers, ) - - builder = firedrake_interface_loopy.KernelBuilder( - integral_data_info, - scalar_type, - diagonal=diagonal, + builder = make_kernel_builder( + integral_data_info, form_data.constants, parameters, diagonal=diagonal, ) - builder.set_entity_numbers(all_meshes) - builder.set_entity_orientations(all_meshes) - builder.set_coordinates(all_meshes) - builder.set_cell_orientations(all_meshes) - builder.set_cell_sizes(all_meshes) - builder.set_coefficients() - # TODO: We do not want pass constants to kernels that do not need them - # so we should attach the constants to integral data instead - builder.set_constants(form_data.constants) + builder.set_coordinates(tuple(integral_data_info.domain_integral_type_map)) ctx = builder.create_context() for integral in integral_data.integrals: params = parameters.copy() @@ -180,7 +185,7 @@ def compile_integral(integral_data, form_data, prefix, parameters, *, diagonal=F integrand_exprs = builder.compile_integrand(integral.integrand(), params, ctx) integral_exprs = builder.construct_integrals(integrand_exprs, params) builder.stash_integrals(integral_exprs, params, ctx) - return builder.construct_kernel(kernel_name, ctx, parameters["add_petsc_events"]) + return builder.construct_kernel(kernel_name, ctx, log=parameters["add_petsc_events"]) def validate_domains(form): @@ -224,8 +229,7 @@ def preprocess_parameters(parameters): def compile_expression_dual_evaluation(expression, ufl_element, *, - domain=None, interface=None, - parameters=None, name=None): + domain=None, parameters=None, name=None): """Compile a UFL expression to be evaluated against a compile-time known reference element's dual basis. Useful for interpolating UFL expressions into e.g. N1curl spaces. @@ -233,225 +237,44 @@ def compile_expression_dual_evaluation(expression, ufl_element, *, :arg expression: UFL expression :arg ufl_element: The UFL element of the target space. :arg domain: optional UFL domain the expression is defined on (required when expression contains no domain). - :arg interface: backend module for the kernel interface :arg parameters: parameters object - :returns: Loopy-based ExpressionKernel object. + :returns: Loopy-based Kernel object. """ - if parameters is None: - parameters = default_parameters() - else: - _ = default_parameters() - _.update(parameters) - parameters = _ - - # Determine whether in complex mode - complex_mode = is_complex(parameters["scalar_type"]) - - orig_coefficients = extract_coefficients(expression) - if isinstance(expression, ufl.Interpolate): - v, operand = expression.argument_slots() - else: - operand = expression - v = ufl.FunctionSpace(extract_unique_domain(operand), ufl_element) - - # Map into reference space - operand = apply_mapping(operand, ufl_element, domain) - - # Apply UFL preprocessing - operand = ufl_utils.preprocess_expression(operand, complex_mode=complex_mode) - operand = simplify_abs(operand, complex_mode) - - # Reconstructed Interpolate with mapped operand - expression = ufl.Interpolate(operand, v) - - # Initialise kernel builder - if interface is None: - # Delayed import, loopy is a runtime dependency - from tsfc.kernel_interface.firedrake_loopy import ExpressionKernelBuilder as interface - - builder = interface(parameters["scalar_type"]) + parameters = preprocess_parameters(parameters) + if not isinstance(expression, ufl.Interpolate): + V = ufl.FunctionSpace(extract_unique_domain(expression) or domain, ufl_element) + expression = ufl.Interpolate(expression, V) arguments = expression.arguments() - argument_multiindices = {arg.number(): builder.create_element(arg.ufl_element()).get_indices() - for arg in arguments} - assert len(argument_multiindices) == len(arguments) - - # Replace coordinates (if any) unless otherwise specified by kwarg - if domain is None: - domain = extract_unique_domain(expression) - assert domain is not None - builder._domain_integral_type_map = {domain: "cell"} - builder._entity_ids = {domain: (0,)} - - # Collect required coefficients and determine numbering - coefficients = extract_coefficients(expression) - coefficient_numbers = tuple(map(orig_coefficients.index, coefficients)) - builder.set_coefficient_numbers(coefficient_numbers) - # Need this ad-hoc fix for now. - for c in coefficients: - d = extract_unique_domain(c) - builder._domain_integral_type_map[d] = "cell" - - elements = [f.ufl_element() for f in (*coefficients, *arguments)] - - needs_external_coords = False - if has_type(expression, GeometricQuantity) or any(map(fem.needs_coordinate_mapping, elements)): - # Create a fake coordinate coefficient for a domain. - coords_coefficient = ufl.Coefficient(ufl.FunctionSpace(domain, domain.ufl_coordinate_element())) - builder.domain_coordinate[domain] = coords_coefficient - builder.set_cell_orientations((domain, )) - builder.set_cell_sizes((domain, )) - coefficients = [coords_coefficient] + coefficients - needs_external_coords = True - builder.set_coefficients(coefficients) - - constants = extract_firedrake_constants(expression) - builder.set_constants(constants) - - # Split mixed coefficients - coeff_splitter = CoefficientSplitter(builder.coefficient_split) - expression = coeff_splitter(expression) - - # Set up kernel config for translation of UFL expression to gem - kernel_cfg = dict(interface=builder, - ufl_cell=domain.ufl_cell(), - integration_dim=as_fiat_cell(domain.ufl_cell()).get_dimension(), - # FIXME: change if we ever implement - # interpolation on facets. - argument_multiindices=argument_multiindices, - index_cache={}, - scalar_type=parameters["scalar_type"]) - - # Create the finat element for the target space - try: - to_element = builder.create_element(ufl_element) - except KeyError: - # FInAT only elements - raise NotImplementedError(f"Don't know how to create FIAT element for {ufl_element}") - - # Allow interpolation onto QuadratureElements to refer to the quadrature - # rule they represent - if isinstance(to_element, finat.QuadratureElement): - kernel_cfg["quadrature_rule"] = to_element._rule - + domains = expression.ufl_domains() dual_arg, operand = expression.argument_slots() + target_domains = join_domains([dual_arg.ufl_function_space().ufl_domain()]) + if len(target_domains) != 1: + raise NotImplementedError("Interpolation onto multiple distinct meshes is not supported") + target_domain, = target_domains + source_domain = domain or extract_unique_domain(operand) or target_domain + if target_domain.topological_dimension == 0 and source_domain.topological_dimension > 0: + ufl_element = ufl_utils.runtime_quadrature_element(source_domain, ufl_element) + + original_coefficients = extract_coefficients(expression) + coefficients = extract_coefficients(expression) + integral_data_info = TSFCIntegralDataInfo( + domain=source_domain, + integral_type="cell", + subdomain_id=("everywhere",), + domain_number=domains.index(target_domain), + domain_integral_type_map={mesh: "cell" for mesh in domains}, + arguments=arguments, + coefficients=coefficients, + coefficient_split=build_coefficient_split( + c for c in coefficients if type(c.ufl_element()) is finat.ufl.MixedElement + ), + coefficient_numbers=tuple(map(original_coefficients.index, coefficients)), + ) - # Create callable for translation of UFL expression to gem - fn = DualEvaluationCallable(operand, kernel_cfg) - - # Get the gem expression for dual evaluation and corresponding basis - # indices needed for compilation of the expression - if isinstance(to_element, NeedsCoordinateMappingElement): - ctx = fem.PointSetContext(**kernel_cfg) - mt = analyse_modified_terminal(ufl.Coefficient(dual_arg.ufl_function_space().dual())) - coordinate_mapping = fem.CoordinateMapping(mt, ctx) - else: - coordinate_mapping = None - evaluation, point_indices, basis_indices = to_element.dual_evaluation(fn, coordinate_mapping) - quadrature_multiindex = tuple(point_indices) - - # Compute the action against the dual argument - if isinstance(dual_arg, ufl.Cofunction): - gem_dual = builder.coefficient_map[dual_arg] - if complex_mode: - evaluation = gem.MathFunction('conj', evaluation) - # The dual argument contracts over the nodes. Split the dual basis - # along its Concatenate nodes first, as assembly does for coefficient - # evaluation. Each block then sums over its own basis indices, rather - # than over the concatenated index, which nothing can be split along. - var, = gem.optimise.remove_componenttensors([gem_dual[basis_indices]]) - summands = [] - for v, expr in unconcatenate([(var, evaluation)], kernel_cfg["index_cache"]): - quadrature_multiindex += v.index_ordering() - summands.append(gem.IndexSum(gem.Product(expr, v), v.index_ordering())) - evaluation = gem.optimise.make_sum(summands) - basis_indices = () - else: - argument_multiindices[dual_arg.number()] = basis_indices - - argument_multiindices = dict(sorted(argument_multiindices.items())) - - # Build kernel body - return_indices = tuple(chain.from_iterable(argument_multiindices.values())) - return_shape = tuple(i.extent for i in return_indices) - return_var = gem.Variable('A', (numpy.prod(return_shape, dtype=int),)) - return_expr = gem.Indexed(gem.reshape(return_var, return_shape), return_indices) - return_expr, = gem.optimise.remove_componenttensors([return_expr]) - - # Contract over the points with the same GEM optimisations as in assembly. - mode = pick_mode(parameters["mode"]) - reps = mode.Integrals([evaluation], quadrature_multiindex, - tuple(argument_multiindices.values()), parameters) - assignments = list(mode.flatten([(return_expr, reps)], kernel_cfg["index_cache"])) - return_variables, expressions = zip(*assignments) - # Argument factorisation does not cancel every Delta here, so lower them. - finalise_options = dict(mode.finalise_options, replace_delta=True) - expressions = impero_utils.preprocess_gem(expressions, **finalise_options) - index_ordering = get_index_ordering(quadrature_multiindex, return_variables) - impero_c = impero_utils.compile_gem(list(zip(return_variables, expressions)), index_ordering) - index_names = {idx: f"p{i}" for (i, idx) in enumerate(basis_indices)} - # Handle kernel interface requirements - builder.register_requirements(expressions) - builder.set_output(return_var) - # Build kernel tuple - return builder.construct_kernel(impero_c, index_names, needs_external_coords, parameters["add_petsc_events"], name=name) - - -class DualEvaluationCallable(object): - """ - Callable representing a function to dual evaluate. - - When called, this takes in a - :class:`finat.point_set.AbstractPointSet` and returns a GEM - expression for evaluation of the function at those points. - - :param expression: UFL expression for the function to dual evaluate. - :param kernel_cfg: A kernel configuration for creation of a - :class:`GemPointContext` or a :class:`PointSetContext` - - Not intended for use outside of - :func:`compile_expression_dual_evaluation`. - """ - def __init__(self, expression, kernel_cfg): - self.expression = expression - self.kernel_cfg = kernel_cfg - - def __call__(self, ps): - """The function to dual evaluate. - - :param ps: The :class:`finat.point_set.AbstractPointSet` for - evaluating at - :returns: a gem expression representing the evaluation of the - input UFL expression at the given point set ``ps``. - For point set points with some shape ``(*value_shape)`` - (i.e. ``()`` for scalar points ``(x)`` for vector points - ``(x, y)`` for tensor points etc) then the gem expression - has shape ``(*value_shape)`` and free indices corresponding - to the input :class:`finat.point_set.AbstractPointSet`'s - free indices alongside any input UFL expression free - indices. - """ - - if not isinstance(ps, finat.point_set.AbstractPointSet): - raise ValueError("Callable argument not a point set!") - - # Avoid modifying saved kernel config - kernel_cfg = self.kernel_cfg.copy() - - if isinstance(ps, finat.point_set.UnknownPointSet): - # Run time known points - kernel_cfg.update(point_indices=ps.indices, point_expr=ps.expression) - # GemPointContext's aren't allowed to have quadrature rules - kernel_cfg.pop("quadrature_rule", None) - translation_context = fem.GemPointContext(**kernel_cfg) - else: - # Compile time known points - kernel_cfg.update(point_set=ps) - translation_context = fem.PointSetContext(**kernel_cfg) - - gem_expr, = fem.compile_ufl(self.expression, translation_context, point_sum=False) - # In some cases ps.indices may be dropped from expr, but nothing - # new should now appear - argument_multiindices = kernel_cfg["argument_multiindices"].values() - assert set(gem_expr.free_indices) <= set(chain(ps.indices, *argument_multiindices)) - - return gem_expr + builder = make_kernel_builder( + integral_data_info, extract_firedrake_constants(expression), parameters, + ) + ctx = builder.create_context() + reps = builder.compile_interpolate(expression, ufl_element, parameters, ctx) + builder.stash_integrals(reps, parameters, ctx) + return builder.construct_kernel(name, ctx, log=parameters["add_petsc_events"]) diff --git a/tsfc/fem.py b/tsfc/fem.py index 95dc093c5c..c8433ab347 100644 --- a/tsfc/fem.py +++ b/tsfc/fem.py @@ -3,8 +3,10 @@ import collections import itertools +from itertools import chain from functools import cached_property, singledispatch +import finat import gem import numpy import ufl @@ -13,12 +15,14 @@ from FIAT.reference_element import TensorProductCell from finat.physically_mapped import (NeedsCoordinateMappingElement, PhysicalGeometry) +from finat.finiteelementbase import FiniteElementBase from finat.point_set import PointSet, PointSingleton +from finat.point_set import AbstractPointSet, UnknownPointSet from finat.quadrature import make_quadrature from finat.element_factory import as_fiat_cell, create_element from gem.node import traversal from gem.optimise import constant_fold_zero, ffc_rounding -from gem.unconcatenate import unconcatenate +from gem.unconcatenate import split_contraction, unconcatenate from ufl.classes import (Argument, CellCoordinate, CellEdgeVectors, CellFacetJacobian, CellOrientation, CellOrigin, CellVertices, CellVolume, Coefficient, FacetArea, @@ -37,6 +41,7 @@ from tsfc.kernel_interface import ProxyKernelInterface from tsfc.kernel_interface.common import lower_integral_type from tsfc.modified_terminals import (analyse_modified_terminal, + ModifiedTerminal, construct_modified_terminal) from tsfc.parameters import is_complex from tsfc.ufl_utils import (ModifiedTerminalMixin, PickRestriction, @@ -114,6 +119,21 @@ def entity_selector(self, callback, domain, restriction): def index_cache(self): return {} + def dual_evaluation_config(self, domain, restriction): + """Kernel config for dual-evaluating a nested Interpolate at ``domain``. + + :arg domain: the domain the nested Interpolate is evaluated on. + :arg restriction: restriction of the modified terminal wrapping it. + """ + return dict( + interface=CellVolumeKernelInterface(self, domain, restriction), + ufl_cell=domain.ufl_cell(), + integration_dim=as_fiat_cell(domain.ufl_cell()).get_dimension(), + argument_multiindices=self.argument_multiindices, + index_cache=self.index_cache, + scalar_type=self.scalar_type, + ) + @cached_property def translator(self): # NOTE: reference cycle! @@ -171,6 +191,10 @@ def coefficient(self, ufl_coefficient, r): assert r is None return self._wrapee.coefficient(ufl_coefficient, self.restriction) + def coefficient_components(self, ufl_coefficient, r): + assert r is None + return self._wrapee.coefficient_components(ufl_coefficient, self.restriction) + class CoordinateMapping(PhysicalGeometry): """Callback class that provides physical geometry to FInAT elements. @@ -336,6 +360,101 @@ def needs_coordinate_mapping(element): return isinstance(create_element(element), NeedsCoordinateMappingElement) +def dual_evaluate(operand: ufl.core.expr.Expr, dual_arg: ufl.Coargument | ufl.Cofunction, + to_element: FiniteElementBase, + kernel_cfg: dict) -> list[tuple[gem.Node, tuple, tuple]]: + """Translate an interpolation operand and evaluate its target dual basis. + + Parameters + ---------- + operand + Expression to evaluate against the target dual basis. + dual_arg + Dual argument from the interpolation that owns ``operand``. + to_element + Target FInAT element. + kernel_cfg + Configuration for the point-evaluation translation context. + + Returns + ------- + list[tuple] + One triple for each summand of the dual basis: the GEM expression for + the local interpolated values, the quadrature indices to sum it over, + and the basis indices that remain uncontracted. + + """ + if isinstance(to_element, finat.QuadratureElement): + kernel_cfg = dict(kernel_cfg, quadrature_rule=to_element._rule) + + fn = DualEvaluationCallable(operand, kernel_cfg) + + if isinstance(to_element, NeedsCoordinateMappingElement): + ctx = PointSetContext(**kernel_cfg) + coefficient = ufl.Coefficient(dual_arg.ufl_function_space().dual()) + coordinate_mapping = CoordinateMapping(analyse_modified_terminal(coefficient), ctx) + else: + coordinate_mapping = None + if isinstance(dual_arg, ufl.Cofunction): + gem_duals = kernel_cfg["interface"].coefficient_components(dual_arg, None) + else: + gem_duals = () + + if not gem_duals: + evaluation, quadrature_indices, basis_indices = to_element.dual_evaluation(fn, coordinate_mapping) + return [(evaluation, tuple(quadrature_indices), basis_indices)] + + # A mixed dual argument has one component per sub-element. + elements = to_element.elements if len(gem_duals) > 1 else (to_element,) + component_summands = [] + for element, gem_dual in zip(elements, gem_duals, strict=True): + evaluation, quadrature_indices, basis_indices = element.dual_evaluation(fn, coordinate_mapping) + if is_complex(kernel_cfg["scalar_type"]): + evaluation = gem.MathFunction("conj", evaluation) + # The dual argument contracts over the nodes, so the basis indices + # reduce here instead of indexing the return value. A direct sum + # tabulates into a Concatenate that only its own component can split. + dual, = gem.optimise.remove_componenttensors([gem_dual[basis_indices]]) + for var, expr in unconcatenate([(dual, evaluation)], kernel_cfg["index_cache"]): + component_summands.append((tuple(quadrature_indices), var, expr)) + + evaluations = [] + for quadrature_indices, var, expr in component_summands: + product = gem.Product(expr, var) + summed_indices = tuple( + index for index in chain(quadrature_indices, var.index_ordering()) + if index in product.free_indices + ) + evaluations.append((product, summed_indices, ())) + return evaluations + + +class DualEvaluationCallable: + """Translate an expression at points requested by a FInAT dual basis.""" + + def __init__(self, operand: ufl.core.expr.Expr, kernel_cfg: dict) -> None: + self.operand = operand + self.kernel_cfg = kernel_cfg + + def __call__(self, point_set: AbstractPointSet) -> gem.Node: + if not isinstance(point_set, AbstractPointSet): + raise ValueError("Callable argument not a point set!") + + kernel_cfg = self.kernel_cfg.copy() + if isinstance(point_set, UnknownPointSet): + kernel_cfg.update(point_indices=point_set.indices, + point_expr=point_set.expression) + kernel_cfg.pop("quadrature_rule", None) + translation_context = GemPointContext(**kernel_cfg) + else: + kernel_cfg.update(point_set=point_set) + translation_context = PointSetContext(**kernel_cfg) + + gem_expr, = compile_ufl(self.operand, translation_context, point_sum=False) + assert set(gem_expr.free_indices) <= set(chain(point_set.indices, *kernel_cfg["argument_multiindices"])) + return gem_expr + + @serial_cache(hashkey=lambda *args: args) def get_quadrature_rule(fiat_cell, integration_dim, quadrature_degree, scheme): integration_cell = fiat_cell.construct_subcomplex(integration_dim) @@ -748,11 +867,36 @@ def translate_constant_value(terminal, mt, ctx): return ctx.constant(terminal) +@translate.register(ufl.Interpolate) +def translate_interpolate(terminal: ufl.Interpolate, mt: ModifiedTerminal, ctx: ContextBase) -> gem.Node: + dual_arg, operand = terminal.argument_slots() + domain = extract_unique_domain(operand) or dual_arg.ufl_function_space().ufl_domain() + element = ctx.create_element(terminal.ufl_element(), restriction=mt.restriction) + kernel_cfg = ctx.dual_evaluation_config(domain, mt.restriction) + vec = gem.Sum(*( + gem.ComponentTensor(evaluation, basis_indices) + for evaluation, _, basis_indices + in dual_evaluate(operand, dual_arg, element, kernel_cfg) + )) + return evaluate_element_values(terminal, mt, ctx, vec, element, beta=element.get_indices()) + + @translate.register(Coefficient) def translate_coefficient(terminal, mt, ctx): - domain = extract_unique_domain(terminal) vec = ctx.coefficient(terminal, mt.restriction) element = ctx.create_element(terminal.ufl_element(), restriction=mt.restriction) + return evaluate_element_values(terminal, mt, ctx, vec, element) + + +def evaluate_element_values(terminal: ufl.core.expr.Expr, mt: ModifiedTerminal, ctx: ContextBase, + vec: gem.Node, element: FiniteElementBase, + beta: tuple | None = None) -> gem.Node: + """Evaluate the function whose values ``vec`` holds, at the current points. + + ``beta`` are the basis indices of ``vec``. They default to the indices that + the context caches for ``terminal``. + """ + domain = extract_unique_domain(terminal) # Collect FInAT tabulation for all entities per_derivative = collections.defaultdict(list) @@ -780,16 +924,42 @@ def take_singleton(xs): for alpha, tables in per_derivative.items()} # Coefficient evaluation - beta = ctx.index_cache.setdefault(terminal.ufl_element(), element.get_indices()) + if beta is None: + beta = ctx.index_cache.setdefault(terminal.ufl_element(), element.get_indices()) zeta = element.get_value_indices() vec_beta, = gem.optimise.remove_componenttensors([gem.Indexed(vec, beta)]) + # The value indices, the form's own quadrature points and the argument + # indices are bound outside this contraction, so they stay free. + unsummed_indices = set(chain(zeta, ctx.point_indices, *ctx.argument_multiindices)) + unsummed_indices.update(ctx.unsummed_coefficient_indices) + # A dat is a view into a kernel argument, which unconcatenate can slice + # into the blocks of a direct sum. Anything else is computed here. + aggregate = vec_beta.children[0] if isinstance( + vec_beta, (gem.Indexed, gem.FlexiblyIndexed)) else None value_dict = {} for alpha, table in per_derivative.items(): table_qi = gem.Indexed(table, beta + zeta) + if not isinstance(aggregate, gem.Variable): + # An interpolated value is a computed expression rather than an + # indexed dat, so no assignment variable carries beta. The + # contraction over beta splits the Concatenate that a direct sum + # tabulates into. Each block then contracts over the + # interpolation points that its own summand evaluates on. + summands = [] + for expr, indices in split_contraction(gem.Product(vec_beta, table_qi), + beta, ctx.index_cache): + indices = tuple(i for i in dict.fromkeys(chain(indices, expr.free_indices)) + if i in expr.free_indices and i not in unsummed_indices) + summands.append(gem.optimise.contraction(gem.IndexSum(expr, indices))) + value_dict[alpha] = gem.ComponentTensor(gem.optimise.make_sum(summands), zeta) + continue + summands = [] for var, expr in unconcatenate([(vec_beta, table_qi)], ctx.index_cache): - indices = tuple(i for i in var.index_ordering() if i not in ctx.unsummed_coefficient_indices) - value = gem.IndexSum(gem.Product(expr, var), indices) + product = gem.Product(expr, var) + indices = tuple(i for i in dict.fromkeys(chain(var.index_ordering(), beta)) + if i not in unsummed_indices and i in product.free_indices) + value = gem.IndexSum(product, indices) summands.append(gem.optimise.contraction(value)) optimised_value = gem.optimise.make_sum(summands) value_dict[alpha] = gem.ComponentTensor(optimised_value, zeta) @@ -797,7 +967,9 @@ def take_singleton(xs): # Change from FIAT to UFL arrangement result = fiat_to_ufl(value_dict, mt.local_derivatives) assert result.shape == mt.expr.ufl_shape - assert set(result.free_indices) - ctx.unsummed_coefficient_indices <= set(ctx.point_indices) + allowed_indices = set(chain(ctx.point_indices, *ctx.argument_multiindices)) + unexpected_indices = set(result.free_indices) - ctx.unsummed_coefficient_indices - allowed_indices + assert not unexpected_indices, unexpected_indices # Detect Jacobian of affine cells if not result.free_indices and all(numpy.count_nonzero(node.array) <= 2 diff --git a/tsfc/kernel_interface/__init__.py b/tsfc/kernel_interface/__init__.py index 3c20720c33..abd59fd892 100644 --- a/tsfc/kernel_interface/__init__.py +++ b/tsfc/kernel_interface/__init__.py @@ -17,6 +17,10 @@ def coefficient(self, ufl_coefficient, restriction): """A function that maps :class:`ufl.Coefficient`s to GEM expressions.""" + @abstractmethod + def coefficient_components(self, ufl_coefficient, restriction): + """Return GEM expressions for a coefficient's stored components.""" + @abstractmethod def constant(self, const): """Return the GEM expression corresponding to the constant.""" diff --git a/tsfc/kernel_interface/common.py b/tsfc/kernel_interface/common.py index c46ba3e8f4..46e3254564 100644 --- a/tsfc/kernel_interface/common.py +++ b/tsfc/kernel_interface/common.py @@ -5,8 +5,12 @@ from itertools import chain, product import copy +from ufl.algorithms.analysis import has_type +from ufl.classes import Cofunction, GeometricQuantity, Interpolate +from ufl.finiteelement import AbstractFiniteElement from ufl.utils.sequences import max_degree -from ufl.domain import extract_unique_domain +from ufl.domain import MeshSequence, extract_unique_domain +from ufl.algorithms.apply_coefficient_split import CoefficientSplitter import gem import gem.impero_utils as impero_utils @@ -18,8 +22,10 @@ from gem.node import traversal from gem.optimise import constant_fold_zero from gem.optimise import remove_componenttensors as prune +from gem.unconcatenate import unconcatenate from numpy import asarray -from tsfc import fem +from tsfc.parameters import is_complex +from tsfc.ufl_utils import preprocess_interpolate from finat.element_factory import as_fiat_cell, create_element from finat.ufl import MixedElement from tsfc.kernel_interface import KernelInterface @@ -38,9 +44,9 @@ def __init__(self, scalar_type): # Coordinates self.domain_coordinate = {} - # Coefficients self.coefficient_map = collections.OrderedDict() + self.coefficient_split = {} # Constants self.constant_map = collections.OrderedDict() @@ -63,6 +69,11 @@ def coefficient(self, ufl_coefficient, restriction): else: return kernel_arg[{'+': 0, '-': 1}[restriction]] + def coefficient_components(self, ufl_coefficient, restriction): + """Return GEM expressions for a coefficient's stored components.""" + coefficients = self.coefficient_split.get(ufl_coefficient, (ufl_coefficient,)) + return tuple(self.coefficient(coefficient, restriction) for coefficient in coefficients) + def constant(self, const): return self.constant_map[const] @@ -133,9 +144,94 @@ def domain_integral_type_map(self): return self._domain_integral_type_map -class KernelBuilderMixin(object): +class KernelBuilderMixin: """Mixin for KernelBuilder classes.""" + def compile_interpolate(self, ufl_interpolate: Interpolate, + target_element: AbstractFiniteElement, + params: dict, ctx: dict) -> list: + """Compile UFL interpolate. + + Parameters + ---------- + ufl_interpolate + Unprocessed UFL interpolation. + target_element + UFL element of the interpolation target. This is not the dual + argument's own element when the target is a point cloud: the + points are then only known at run time, so the target is a + quadrature element on the source cell. + params + Parameter dictionary containing ``"mode"``. + ctx + Context created with :meth:`create_context`. + + Returns + ------- + list + Integral representations for the interpolation. + """ + preprocessed_interpolate = preprocess_interpolate( + ufl_interpolate, + target_element, + self.integral_data_info.domain, + complex_mode=is_complex(params["scalar_type"]), + ) + dual_arg, operand = preprocessed_interpolate.argument_slots() + operand = CoefficientSplitter(self.coefficient_split)(operand) + from tsfc import fem + + elements = [f.ufl_element() for f in (*self.integral_data_info.coefficients, + *self.integral_data_info.arguments)] + needs_external_coords = bool( + has_type(operand, GeometricQuantity) + or any(map(fem.needs_coordinate_mapping, elements)) + ) + if needs_external_coords: + self.set_coordinates(tuple(self.integral_data_info.domain_integral_type_map)) + try: + target_element = self.create_element(target_element) + except KeyError: + # FInAT only elements + raise NotImplementedError(f"Don't know how to create FIAT element for {target_element}") + config = self.fem_config() + config.update(argument_multiindices=self.argument_multiindices, + index_cache=ctx["index_cache"]) + evaluations = fem.dual_evaluate(operand, dual_arg, target_element, config) + if not isinstance(dual_arg, Cofunction): + evaluation, quadrature_indices, basis_indices = evaluations[0] + # A dual Argument indexes the return value, so the dual basis must + # tabulate onto the indices the output tensor was built with. + output_indices = self.argument_multiindices[self.integral_data_info.arguments.index(dual_arg)] + if tuple(i.extent for i in basis_indices) != tuple(i.extent for i in output_indices): + raise ValueError("Interpolation output index shape mismatch") + evaluation, = gem.optimise.remove_componenttensors([evaluation], tuple(zip(basis_indices, output_indices))) + evaluations = [(evaluation, quadrature_indices, output_indices)] + + return_variables = [] + reps = [] + for evaluation, quadrature_indices, _ in evaluations: + for variable, expr in unconcatenate( + [(self.return_variables[0], evaluation)], ctx["index_cache"]): + summed_indices = tuple( + dict.fromkeys(chain( + (index for index in quadrature_indices + if index in expr.free_indices), + (index for index in expr.free_indices + if index not in variable.free_indices), + )) + ) + ctx["quadrature_indices"].extend(summed_indices) + return_variables.append(variable) + reps.extend(self.construct_integrals( + [expr], params, summed_indices, + (variable.index_ordering(),) + )) + self.return_variables = tuple(return_variables) + # Argument factorisation does not cancel every Delta here, so lower them. + ctx["finalise_options"]["replace_delta"] = True + return reps + def compile_integrand(self, integrand, params, ctx): """Compile UFL integrand. @@ -154,12 +250,16 @@ def compile_integrand(self, integrand, params, ctx): config['argument_multiindices'] = self.argument_multiindices config['quadrature_rule'] = quad_rule config['index_cache'] = ctx['index_cache'] + from tsfc import fem + expressions = fem.compile_ufl(integrand, fem.PointSetContext(**config)) ctx['quadrature_indices'].extend(quad_rule.point_set.indices) return expressions - def construct_integrals(self, integrand_expressions, params): + def construct_integrals(self, integrand_expressions, params, + quadrature_multiindex=None, + argument_multiindices=None): """Construct integrals from integrand expressions. :arg integrand_expressions: gem expressions for integrands. @@ -170,12 +270,22 @@ def construct_integrals(self, integrand_expressions, params): method or by modifying the gem expressions returned by :meth:`compile_integrand`. + quadrature_multiindex is the sequence of indices to contract. When it + is not given, use the point indices from the quadrature rule. + + argument_multiindices are the free indices of the return variables. + When they are not given, use the builder's argument multiindices. + See :meth:`create_context` for typical calling sequence. """ mode = pick_mode(params["mode"]) + if quadrature_multiindex is None: + quadrature_multiindex = params["quadrature_rule"].point_set.indices + if argument_multiindices is None: + argument_multiindices = self.argument_multiindices return mode.Integrals(integrand_expressions, - params["quadrature_rule"].point_set.indices, - self.argument_multiindices, + quadrature_multiindex, + argument_multiindices, params) def stash_integrals(self, reps, params, ctx): @@ -220,6 +330,7 @@ def compile_gem(self, ctx): options = dict(reduce(operator.and_, [mode.finalise_options.items() for mode in mode_irs.keys()])) + options.update(ctx['finalise_options']) expressions = impero_utils.preprocess_gem(expressions, **options) # Let the kernel interface inspect the optimised IR to register @@ -277,6 +388,11 @@ def create_context(self): Dict for mode representations. + *finalise_options* + + Options overriding the modes' own :func:`impero_utils.preprocess_gem` + options. + For each set of integrals to make a kernel for (i,e., `integral_data.integrals`), one must first create a ctx object by calling :meth:`create_context` method. @@ -299,12 +415,15 @@ def create_context(self): """ return {'index_cache': {}, 'quadrature_indices': [], + 'finalise_options': {}, 'mode_irs': collections.OrderedDict()} def set_quad_rule(params, cell, integral_type, functions): # Check if the integral has a quad degree or quad element attached, # otherwise use the estimated polynomial degree attached by compute_form_data + from tsfc import fem + quad_rule = params.get("quadrature_rule", "default") elements = [] for f in functions: @@ -577,7 +696,16 @@ def expression(restricted): c_shape = copy.deepcopy(u_shape) rs_tuples = [] for arg_num, arg in enumerate(arguments): - integral_type = domain_integral_type_map[extract_unique_domain(arg)] + domain = arg.ufl_function_space().ufl_domain() + try: + integral_type = domain_integral_type_map[domain] + except KeyError: + # An unsplit argument (e.g. a mixed-space patch argument) reports + # its domain as a MeshSequence rather than a single mesh: every + # mesh it sequences is the same iteration, so they must agree. + if not isinstance(domain, MeshSequence): + raise + integral_type, = {domain_integral_type_map[m] for m in domain.meshes} if integral_type is None: raise RuntimeError(f"Can not determine integral_type on {arg}") if integral_type.startswith("interior_facet"): diff --git a/tsfc/kernel_interface/firedrake_loopy.py b/tsfc/kernel_interface/firedrake_loopy.py index cc8fd7a61e..2f7d2909de 100644 --- a/tsfc/kernel_interface/firedrake_loopy.py +++ b/tsfc/kernel_interface/firedrake_loopy.py @@ -4,7 +4,7 @@ from ufl import Coefficient, FunctionSpace from ufl.domain import MeshSequence -from finat.ufl import MixedElement as ufl_MixedElement, FiniteElement +from finat.ufl import FiniteElement import gem from gem.flop_count import count_flops @@ -17,14 +17,6 @@ from tsfc.loopy import generate as generate_loopy -# Expression kernel description type -ExpressionKernel = namedtuple('ExpressionKernel', ['ast', 'oriented', 'needs_cell_sizes', - 'coefficient_numbers', - 'needs_external_coords', - 'tabulations', 'name', 'arguments', - 'flop_count', 'event']) - - ActiveDomainNumbers = namedtuple('ActiveDomainNumbers', ['coordinates', 'cell_orientations', 'cell_sizes', @@ -79,6 +71,24 @@ def __init__(self, ast=None, arguments=None, integral_type=None, self.name = name self.event = event + def _has_argument(self, argument_type) -> bool: + return any(isinstance(arg, argument_type) for arg in self.arguments or ()) + + @property + def needs_external_coords(self) -> bool: + """Whether the kernel expects coordinates from the caller.""" + return self._has_argument(kernel_args.CoordinatesKernelArg) + + @property + def oriented(self) -> bool: + """Whether the kernel expects cell orientations from the caller.""" + return self._has_argument(kernel_args.CellOrientationsKernelArg) + + @property + def needs_cell_sizes(self) -> bool: + """Whether the kernel expects cell sizes from the caller.""" + return self._has_argument(kernel_args.CellSizesKernelArg) + class KernelBuilderBase(_KernelBuilderBase): @@ -152,18 +162,20 @@ def set_cell_sizes(self, domains): measure of the mesh size around each vertex (hence this lives in P1). - Should the domain have topological dimension 0 this does - nothing. + A domain of topological dimension 0 gets a ``None`` entry: every + domain must keep its slot, since the active domain numbers index + this dict positionally. """ self._cell_sizes = {} for i, domain in enumerate(domains): if domain.ufl_cell().topological_dimension > 0: - # Can't create P1 since only P0 is a valid finite element if - # topological_dimension is 0 and the concept of "cell size" - # is not useful for a vertex. f = Coefficient(FunctionSpace(domain, FiniteElement("P", domain.ufl_cell(), 1))) expr = prepare_coefficient(f, f"cell_sizes_{i}", self._domain_integral_type_map) - self._cell_sizes[domain] = expr + else: + # Only P0 is a valid finite element on a vertex, and the + # concept of "cell size" is not useful there. + expr = None + self._cell_sizes[domain] = expr def create_element(self, element, **kwargs): """Create a FInAT element (suitable for tabulating with) given @@ -190,97 +202,6 @@ def generate_arg_from_expression(self, expr, dtype=None): return self.generate_arg_from_variable(var, dtype=dtype or self.scalar_type) -class ExpressionKernelBuilder(KernelBuilderBase): - """Builds expression kernels for UFL interpolation in Firedrake.""" - - def __init__(self, scalar_type): - super(ExpressionKernelBuilder, self).__init__(scalar_type=scalar_type) - self.oriented = False - self.cell_sizes = False - - def set_coefficients(self, coefficients): - """Prepare the coefficients of the expression. - - :arg coefficients: UFL coefficients from Firedrake - """ - self.coefficient_split = {} - - for i, coefficient in enumerate(coefficients): - if type(coefficient.ufl_element()) == ufl_MixedElement: - subcoeffs = coefficient.subfunctions # Firedrake-specific - self.coefficient_split[coefficient] = subcoeffs - for j, subcoeff in enumerate(subcoeffs): - self._coefficient(subcoeff, f"w_{i}_{j}") - else: - self._coefficient(coefficient, f"w_{i}") - - def set_constants(self, constants): - for i, const in enumerate(constants): - gemexpr = prepare_constant(const, i) - self.constant_map[const] = gemexpr - - def set_coefficient_numbers(self, coefficient_numbers): - """Store the coefficient indices of the original form. - - :arg coefficient_numbers: Iterable of indices describing which coefficients - from the input expression need to be passed in to the kernel. - """ - self.coefficient_numbers = coefficient_numbers - - def register_requirements(self, ir): - """Inspect what is referenced by the IR that needs to be - provided by the kernel interface.""" - self.oriented, self.cell_sizes, self.tabulations = check_requirements(ir) - - def set_output(self, o): - """Produce the kernel return argument""" - loopy_arg = lp.GlobalArg(o.name, dtype=self.scalar_type, shape=o.shape) - self.output_arg = kernel_args.OutputKernelArg(loopy_arg) - - def construct_kernel(self, impero_c, index_names, needs_external_coords, log=False, name=None): - """Constructs an :class:`ExpressionKernel`. - - :arg impero_c: gem.ImperoC object that represents the kernel - :arg index_names: pre-assigned index names - :arg needs_external_coords: If ``True``, the first argument to - the kernel is an externally provided coordinate field. - :arg log: bool if the Kernel should be profiled with Log events - - :returns: :class:`ExpressionKernel` object - """ - args = [self.output_arg] - if self.oriented: - cell_orientations, = tuple(self._cell_orientations.values()) - funarg = self.generate_arg_from_expression(cell_orientations, dtype=numpy.int32) - args.append(kernel_args.CellOrientationsKernelArg(funarg)) - if self.cell_sizes: - cell_sizes, = tuple(self._cell_sizes.values()) - funarg = self.generate_arg_from_expression(cell_sizes) - args.append(kernel_args.CellSizesKernelArg(funarg)) - for _, expr in self.coefficient_map.items(): - # coefficient_map is OrderedDict. - funarg = self.generate_arg_from_expression(expr) - args.append(kernel_args.CoefficientKernelArg(funarg)) - - # now constants - for gemexpr in self.constant_map.values(): - funarg = self.generate_arg_from_expression(gemexpr) - args.append(kernel_args.ConstantKernelArg(funarg)) - - for name_, shape in self.tabulations: - tab_loopy_arg = lp.GlobalArg(name_, dtype=self.scalar_type, shape=shape) - args.append(kernel_args.TabulationKernelArg(tab_loopy_arg)) - - loopy_args = [arg.loopy_arg for arg in args] - - name = name or "expression_kernel" - loopy_kernel, event = generate_loopy(impero_c, loopy_args, self.scalar_type, - name, index_names, log=log) - return ExpressionKernel(loopy_kernel, self.oriented, self.cell_sizes, - self.coefficient_numbers, needs_external_coords, - self.tabulations, name, args, count_flops(impero_c), event) - - class KernelBuilder(KernelBuilderBase, KernelBuilderMixin): """Helper class for building a :class:`Kernel` object.""" @@ -293,8 +214,12 @@ def __init__(self, integral_data_info, scalar_type, self.local_tensor = None self.coefficient_number_index_map = OrderedDict() self.integral_data_info = integral_data_info - self._domain_integral_type_map = integral_data_info.domain_integral_type_map # For consistency with ExpressionKernelBuilder. + self.coefficient_split = integral_data_info.coefficient_split + self._domain_integral_type_map = integral_data_info.domain_integral_type_map self.set_arguments() + domains = tuple(integral_data_info.domain_integral_type_map) + self.set_entity_numbers(domains) + self.set_entity_orientations(domains) def set_arguments(self): """Process arguments.""" diff --git a/tsfc/modified_terminals.py b/tsfc/modified_terminals.py index a26e5c2980..2b63022188 100644 --- a/tsfc/modified_terminals.py +++ b/tsfc/modified_terminals.py @@ -23,7 +23,7 @@ from ufl.classes import (ReferenceValue, ReferenceGrad, NegativeRestricted, PositiveRestricted, Restricted, ConstantValue, - Jacobian, SpatialCoordinate, Zero) + Interpolate, Jacobian, SpatialCoordinate, Zero) from ufl.checks import is_cellwise_constant from ufl.domain import extract_unique_domain @@ -82,7 +82,7 @@ def __str__(self): def is_modified_terminal(v): "Check if v is a terminal or a terminal wrapped in terminal modifier types." - while not v._ufl_is_terminal_: + while not (v._ufl_is_terminal_ or isinstance(v, Interpolate)): if v._ufl_is_terminal_modifier_: v = v.ufl_operands[0] else: @@ -92,7 +92,7 @@ def is_modified_terminal(v): def strip_modified_terminal(v): "Extract core Terminal from a modified terminal or return None." - while not v._ufl_is_terminal_: + while not (v._ufl_is_terminal_ or isinstance(v, Interpolate)): if v._ufl_is_terminal_modifier_: v = v.ufl_operands[0] else: @@ -115,7 +115,7 @@ def analyse_modified_terminal(expr): # Start with expr and strip away layers of modifiers t = expr - while not t._ufl_is_terminal_: + while not (t._ufl_is_terminal_ or isinstance(t, Interpolate)): if isinstance(t, ReferenceValue): assert reference_value is None, "Got twice pulled back terminal!" reference_value = True diff --git a/tsfc/spectral.py b/tsfc/spectral.py index a521fdb2fd..9d9496e7fb 100644 --- a/tsfc/spectral.py +++ b/tsfc/spectral.py @@ -2,7 +2,7 @@ from functools import partial from itertools import chain, zip_longest -from gem.gem import Delta, Indexed, Sum, index_sum, one +from gem.gem import Delta, FlexiblyIndexed, Indexed, Sum, index_sum, one from gem.node import Memoizer, MemoizerArg from gem.optimise import filtered_replace_indices from gem.optimise import delta_elimination as _delta_elimination @@ -107,8 +107,10 @@ def group_key(pair): monomial_sum = delta_simplified[variable] # Collect sum indices applicable to the current MonomialSum sum_indices = set(chain.from_iterable(m.sum_indices for m in monomial_sum)) - # Put them in a deterministic order - sum_indices = [i for i in quadrature_indices if i in sum_indices] + # Put them in a deterministic order, quadrature indices first + sum_indices = ([i for i in quadrature_indices if i in sum_indices] + + sorted(sum_indices.difference(quadrature_indices), + key=lambda index: index.count)) # Apply sum factorisation combined with COFFEE technology expression = sum_factorise(variable, sum_indices, monomial_sum) yield (variable, expression) @@ -123,7 +125,7 @@ def classify(argument_indices, expression, delta_inside): if n == 0: return OTHER elif n == 1: - if isinstance(expression, (Delta, Indexed)) and not delta_inside(expression): + if isinstance(expression, (Delta, FlexiblyIndexed, Indexed)) and not delta_inside(expression): return ATOMIC else: return COMPOUND diff --git a/tsfc/ufl_utils.py b/tsfc/ufl_utils.py index ce25dc3087..e9e0739701 100644 --- a/tsfc/ufl_utils.py +++ b/tsfc/ufl_utils.py @@ -2,14 +2,16 @@ from functools import singledispatch -import numpy - import ufl -from ufl import as_tensor, indices, replace +from ufl import replace from ufl.algorithms import compute_form_data as ufl_compute_form_data from ufl.algorithms import estimate_total_polynomial_degree from ufl.algorithms.analysis import extract_arguments, extract_coefficients, extract_type -from ufl.algorithms.apply_function_pullbacks import apply_function_pullbacks +from ufl.algorithms.apply_function_pullbacks import ( + apply_function_pullbacks, + apply_interpolate_pullbacks, + apply_inverse_pullback, +) from ufl.algorithms.apply_algebra_lowering import apply_algebra_lowering from ufl.algorithms.apply_derivatives import apply_derivatives from ufl.algorithms.apply_geometry_lowering import apply_geometry_lowering @@ -27,15 +29,97 @@ Product, ScalarValue, Sqrt, Zero, CellVolume, FacetArea) from ufl.utils.sorting import sorted_by_count -from ufl.domain import extract_domains, extract_unique_domain +from ufl.domain import extract_domains +import gem from gem.node import MemoizerArg +from finat.element_factory import as_fiat_cell +from finat.point_set import UnknownPointSet +from finat.quadrature import QuadratureRule +from finat.ufl import FiniteElement, TensorElement + from tsfc.modified_terminals import is_modified_terminal, analyse_modified_terminal preserve_geometry_types = (CellVolume, FacetArea) +# The gem.Variable that holds a point which is only known at run time. TSFC +# tabulates at any variable whose name starts with "rt_", and the assembler +# passes the reference coordinates of the point cloud in its place. +RUNTIME_POINT_VARIABLE = "rt_X" + + +def runtime_quadrature_element(domain: ufl.AbstractDomain, + ufl_element: ufl.AbstractFiniteElement) -> ufl.AbstractFiniteElement: + """Construct a Quadrature element for interpolation onto a runtime point. + + The point is one that is only known at run time, for example a point of a + `firedrake.mesh.VertexOnlyMesh`. + + Parameters + ---------- + domain + The source domain, whose cells the point is located in. + ufl_element + The UFL element of the target FunctionSpace. + + Returns + ------- + ufl.AbstractFiniteElement + A Quadrature element on the source cell, with one point that TSFC reads + from `RUNTIME_POINT_VARIABLE`, and the value shape of ``ufl_element``. + + """ + cell = domain.ufl_cell() + point_expr = gem.Variable(RUNTIME_POINT_VARIABLE, (1, cell.topological_dimension)) + point_set = UnknownPointSet(point_expr) + rule = QuadratureRule(point_set, weights=[1.0], ref_el=as_fiat_cell(cell)) + + shape = ufl_element.pullback.physical_value_shape(ufl_element, domain) + rt_element = FiniteElement("Quadrature", cell=cell, degree=0, quad_scheme=rule) + if shape: + symmetry = None if len(shape) < 2 else ufl_element.symmetry() + rt_element = TensorElement(rt_element, shape=shape, symmetry=symmetry) + return rt_element + + +def preprocess_interpolate(ufl_interpolate: ufl.Interpolate, + element: ufl.AbstractFiniteElement, + domain: ufl.AbstractDomain, + complex_mode: bool = False) -> ufl.Interpolate: + """Prepare a standalone interpolation for TSFC. + + Parameters + ---------- + ufl_interpolate + The interpolation to preprocess. + element + The UFL element of the interpolation target. + domain + The domain the operand is evaluated on. + complex_mode + Is the scalar type complex? + + Returns + ------- + ufl.Interpolate + The interpolation with its operand in ``element``'s reference frame. + + Notes + ----- + A standalone interpolation never reaches `compute_form_data`, so the operand + gets the scalar preprocessing here. Interpolations inside a form are lowered + by `ufl.algorithms.apply_interpolate_pullbacks`. + """ + dual_arg, operand = ufl_interpolate.argument_slots() + operand = apply_inverse_pullback(operand, element, domain) + operand = preprocess_expression(operand, complex_mode=complex_mode) + operand = simplify_abs(operand, complex_mode) + # Build the UFL node directly: the operand is now in the reference frame, + # so it no longer matches the physical shape a Firedrake Interpolate checks. + return ufl.Interpolate(operand, dual_arg) + def compute_form_data(form, do_apply_function_pullbacks=True, @@ -136,6 +220,7 @@ def preprocess_expression(expression, complex_mode=False, Useful, for example, to preprocess non-scalar expressions, which are not and cannot be forms. """ + expression = apply_interpolate_pullbacks(expression) if complex_mode: expression = do_comparison_check(expression) else: @@ -328,141 +413,6 @@ def simplify_abs(expression, complex_mode): return mapper(expression, False) -def apply_mapping(expression, element, domain): - """Apply the inverse of the pullback for element to an expression. - - :arg expression: An expression in physical space - :arg element: The element we're going to interpolate into, whose - value_shape must match the shape of the expression, and will - advertise the pullback to apply. - :arg domain: Optional domain to provide in case expression does - not contain a domain (used for constructing geometric quantities). - :returns: A new UFL expression with shape element.reference_value_shape - :raises NotImplementedError: If we don't know how to apply the - inverse of the pullback. - :raises ValueError: If we get shape mismatches. - - The following is borrowed from the UFC documentation: - - Let g be a field defined on a physical domain T with physical - coordinates x. Let T_0 be a reference domain with coordinates - X. Assume that F: T_0 -> T such that - - x = F(X) - - Let J be the Jacobian of F, i.e J = dx/dX and let K denote the - inverse of the Jacobian K = J^{-1}. Then we (currently) have the - following four types of mappings: - - 'identity' mapping for g: - - G(X) = g(x) - - For vector fields g: - - 'contravariant piola' mapping for g: - - G(X) = det(J) K g(x) i.e G_i(X) = det(J) K_ij g_j(x) - - 'covariant piola' mapping for g: - - G(X) = J^T g(x) i.e G_i(X) = J^T_ij g(x) = J_ji g_j(x) - - 'double covariant piola' mapping for g: - - G(X) = J^T g(x) J i.e. G_il(X) = J_ji g_jk(x) J_kl - - 'double contravariant piola' mapping for g: - - G(X) = det(J)^2 K g(x) K^T i.e. G_il(X)=(detJ)^2 K_ij g_jk K_lk - - 'covariant contravariant piola' mapping for g: - - G(X) = det(J) J^T g(x) K^T i.e. G_il(X) = det(J) J_ji g_jk(x) K_lk - - If 'contravariant piola' or 'covariant piola' (or their double - variants) are applied to a matrix-valued function, the appropriate - mappings are applied row-by-row. - """ - mesh = extract_unique_domain(expression) - if mesh is None: - mesh = domain - if domain is not None and mesh != domain: - raise NotImplementedError("Multiple domains not supported") - pvs = element.pullback.physical_value_shape(element, mesh) - if expression.ufl_shape != pvs: - raise ValueError(f"Mismatching shapes, got {expression.ufl_shape}, expected {pvs}") - mapping = element.mapping().lower() - if mapping == "identity": - rexpression = expression - elif mapping == "covariant piola": - J = Jacobian(mesh) - *k, i, j = indices(len(expression.ufl_shape) + 1) - kj = (*k, j) - rexpression = as_tensor(J[j, i] * expression[kj], (*k, i)) - elif mapping == "l2 piola": - detJ = JacobianDeterminant(mesh) - rexpression = expression * detJ - elif mapping == "contravariant piola": - K = JacobianInverse(mesh) - detJ = JacobianDeterminant(mesh) - *k, i, j = indices(len(expression.ufl_shape) + 1) - kj = (*k, j) - rexpression = as_tensor(detJ * K[i, j] * expression[kj], (*k, i)) - elif mapping == "double covariant piola": - J = Jacobian(mesh) - *k, i, j, m, n = indices(len(expression.ufl_shape) + 2) - kmn = (*k, m, n) - rexpression = as_tensor(J[m, i] * expression[kmn] * J[n, j], (*k, i, j)) - elif mapping == "double contravariant piola": - K = JacobianInverse(mesh) - detJ = JacobianDeterminant(mesh) - *k, i, j, m, n = indices(len(expression.ufl_shape) + 2) - kmn = (*k, m, n) - rexpression = as_tensor(detJ**2 * K[i, m] * expression[kmn] * K[j, n], (*k, i, j)) - elif mapping == "covariant contravariant piola": - J = Jacobian(mesh) - K = JacobianInverse(mesh) - detJ = JacobianDeterminant(mesh) - *k, i, j, m, n = indices(len(expression.ufl_shape) + 2) - kmn = (*k, m, n) - rexpression = as_tensor(detJ * J[m, i] * expression[kmn] * K[j, n], (*k, i, j)) - elif mapping == "symmetries": - # This tells us how to get from the pieces of the reference - # space expression to the physical space one. - # We're going to apply the inverse of the physical to - # reference space mapping. - fcm = element.flattened_sub_element_mapping() - sub_elem = element.sub_elements[0] - shape = expression.ufl_shape - flat = ufl.as_vector([expression[i] for i in numpy.ndindex(shape)]) - vs = sub_elem.pullback.physical_value_shape(sub_elem, mesh) - rvs = sub_elem.reference_value_shape - seen = set() - rpieces = [] - gm = int(numpy.prod(vs, dtype=int)) - for gi, ri in enumerate(fcm): - # For each unique piece in reference space - if ri in seen: - continue - seen.add(ri) - # Get the physical space piece - piece = [flat[gm*gi + j] for j in range(gm)] - piece = as_tensor(numpy.asarray(piece).reshape(vs)) - # get into reference space - piece = apply_mapping(piece, sub_elem, mesh) - assert piece.ufl_shape == rvs - # Concatenate with the other pieces - rpieces.extend([piece[idx] for idx in numpy.ndindex(rvs)]) - # And reshape - rexpression = as_tensor(numpy.asarray(rpieces).reshape(element.reference_value_shape)) - else: - raise NotImplementedError(f"Don't know how to handle mapping type {mapping} for expression of rank {ufl.FunctionSpace(mesh, element).value_shape}") - if rexpression.ufl_shape != element.reference_value_shape: - raise ValueError(f"Mismatching reference shapes, got {rexpression.ufl_shape} expected {element.reference_value_shape}") - return rexpression - - class TSFCConstantMixin: """ Mixin class to identify Constants """