Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 19 additions & 16 deletions tsfc/driver.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@

import gem
import gem.impero_utils as impero_utils
from gem.unconcatenate import unconcatenate

import finat
from finat.element_factory import as_fiat_cell
Expand All @@ -25,6 +24,7 @@
from tsfc.parameters import default_parameters, is_complex
from tsfc.ufl_utils import apply_mapping, extract_firedrake_constants, simplify_abs
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


Expand Down Expand Up @@ -345,42 +345,45 @@ def compile_expression_dual_evaluation(expression, ufl_element, *,
coordinate_mapping = fem.CoordinateMapping(mt, ctx)
else:
coordinate_mapping = None
evaluation, basis_indices = to_element.dual_evaluation(fn, coordinate_mapping)
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)
evaluation = gem.IndexSum(evaluation * gem_dual[basis_indices], basis_indices)
# The dual argument contracts over the nodes, so the basis indices are
# reduction indices like the points, not indices of the return value.
evaluation = evaluation * gem_dual[basis_indices]
quadrature_multiindex += tuple(basis_indices)
basis_indices = ()
else:
argument_multiindices[dual_arg.number()] = basis_indices

argument_multiindices = dict(sorted(argument_multiindices.items()))

# Unroll
max_extent = parameters["unroll_indexsum"]
if max_extent:
def predicate(index):
return index.extent <= max_extent
evaluation, = gem.optimise.unroll_indexsum([evaluation], predicate=predicate)

# 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])

# TODO: one should apply some GEM optimisations as in assembly,
# but we don't for now.
evaluation, = impero_utils.preprocess_gem([evaluation])
pairs = unconcatenate([(return_expr, evaluation)])
impero_c = impero_utils.compile_gem(pairs, return_indices)
# 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([evaluation])
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)
Expand Down
Loading