Skip to content

Commit 8711da9

Browse files
authored
Merge pull request #261 from firedrakeproject/ksagiyam/tsfc_refactor_1
Ksagiyam/tsfc refactor 1
2 parents 4cdf8f4 + 4a1cd20 commit 8711da9

5 files changed

Lines changed: 534 additions & 341 deletions

File tree

tsfc/driver.py

Lines changed: 59 additions & 239 deletions
Original file line numberDiff line numberDiff line change
@@ -1,31 +1,21 @@
11
import collections
2-
import operator
3-
import string
42
import time
53
import sys
6-
from functools import reduce
74
from itertools import chain
85
from finat.physically_mapped import DirectlyDefinedElement, PhysicallyMappedElement
96

10-
from numpy import asarray
11-
127
import ufl
138
from ufl.algorithms import extract_arguments, extract_coefficients
149
from ufl.algorithms.analysis import has_type
1510
from ufl.classes import Form, GeometricQuantity
1611
from ufl.log import GREEN
17-
from ufl.utils.sequences import max_degree
1812

1913
import gem
2014
import gem.impero_utils as impero_utils
2115

22-
from FIAT.reference_element import TensorProductCell
23-
2416
import finat
25-
from finat.quadrature import AbstractQuadratureRule, make_quadrature
2617

2718
from tsfc import fem, ufl_utils
28-
from tsfc.finatinterface import as_fiat_cell
2919
from tsfc.logging import logger
3020
from tsfc.parameters import default_parameters, is_complex
3121
from tsfc.ufl_utils import apply_mapping
@@ -34,6 +24,28 @@
3424
sys.setrecursionlimit(3000)
3525

3626

27+
TSFCIntegralDataInfo = collections.namedtuple("TSFCIntegralDataInfo",
28+
["domain", "integral_type", "subdomain_id", "domain_number",
29+
"arguments",
30+
"coefficients", "coefficient_numbers"])
31+
TSFCIntegralDataInfo.__doc__ = """
32+
Minimal set of objects for kernel builders.
33+
34+
domain - The mesh.
35+
integral_type - The type of integral.
36+
subdomain_id - What is the subdomain id for this kernel.
37+
domain_number - Which domain number in the original form
38+
does this kernel correspond to (can be used to index into
39+
original_form.ufl_domains() to get the correct domain).
40+
coefficients - A list of coefficients.
41+
coefficient_numbers - A list of which coefficients from the
42+
form the kernel needs.
43+
44+
This is a minimal set of objects that kernel builders need to
45+
construct a kernel from :attr:`integrals` of :class:`~ufl.IntegralData`.
46+
"""
47+
48+
3749
def compile_form(form, prefix="form", parameters=None, interface=None, coffee=True, diagonal=False):
3850
"""Compiles a UFL form into a set of assembly kernels.
3951
@@ -76,12 +88,7 @@ def compile_integral(integral_data, form_data, prefix, parameters, interface, co
7688
:arg diagonal: Are we building a kernel for the diagonal of a rank-2 element tensor?
7789
:returns: a kernel constructed by the kernel interface
7890
"""
79-
if parameters is None:
80-
parameters = default_parameters()
81-
else:
82-
_ = default_parameters()
83-
_.update(parameters)
84-
parameters = _
91+
parameters = preprocess_parameters(parameters)
8592
if interface is None:
8693
if coffee:
8794
import tsfc.kernel_interface.firedrake as firedrake_interface_coffee
@@ -90,180 +97,61 @@ def compile_integral(integral_data, form_data, prefix, parameters, interface, co
9097
# Delayed import, loopy is a runtime dependency
9198
import tsfc.kernel_interface.firedrake_loopy as firedrake_interface_loopy
9299
interface = firedrake_interface_loopy.KernelBuilder
93-
if coffee:
94-
scalar_type = parameters["scalar_type_c"]
95-
else:
96-
scalar_type = parameters["scalar_type"]
97-
98-
# Remove these here, they're handled below.
99-
if parameters.get("quadrature_degree") in ["auto", "default", None, -1, "-1"]:
100-
del parameters["quadrature_degree"]
101-
if parameters.get("quadrature_rule") in ["auto", "default", None]:
102-
del parameters["quadrature_rule"]
103-
100+
scalar_type = parameters["scalar_type"]
104101
integral_type = integral_data.integral_type
105-
interior_facet = integral_type.startswith("interior_facet")
106102
mesh = integral_data.domain
107-
cell = integral_data.domain.ufl_cell()
108103
arguments = form_data.preprocessed_form.arguments()
109104
kernel_name = "%s_%s_integral_%s" % (prefix, integral_type, integral_data.subdomain_id)
110105
# Handle negative subdomain_id
111106
kernel_name = kernel_name.replace("-", "_")
112-
113-
fiat_cell = as_fiat_cell(cell)
114-
integration_dim, entity_ids = lower_integral_type(fiat_cell, integral_type)
115-
116-
quadrature_indices = []
117-
118107
# Dict mapping domains to index in original_form.ufl_domains()
119108
domain_numbering = form_data.original_form.domain_numbering()
120-
builder = interface(integral_type, integral_data.subdomain_id,
121-
domain_numbering[integral_data.domain],
109+
domain_number = domain_numbering[integral_data.domain]
110+
coefficients = [form_data.function_replace_map[c] for c in integral_data.integral_coefficients]
111+
# This is which coefficient in the original form the
112+
# current coefficient is.
113+
# Consider f*v*dx + g*v*ds, the full form contains two
114+
# coefficients, but each integral only requires one.
115+
coefficient_numbers = tuple(form_data.original_coefficient_positions[i]
116+
for i, (_, enabled) in enumerate(zip(form_data.reduced_coefficients, integral_data.enabled_coefficients))
117+
if enabled)
118+
integral_data_info = TSFCIntegralDataInfo(domain=integral_data.domain,
119+
integral_type=integral_data.integral_type,
120+
subdomain_id=integral_data.subdomain_id,
121+
domain_number=domain_number,
122+
arguments=arguments,
123+
coefficients=coefficients,
124+
coefficient_numbers=coefficient_numbers)
125+
builder = interface(integral_data_info,
122126
scalar_type,
123127
diagonal=diagonal)
124-
argument_multiindices = tuple(builder.create_element(arg.ufl_element()).get_indices()
125-
for arg in arguments)
126-
if diagonal:
127-
# Error checking occurs in the builder constructor.
128-
# Diagonal assembly is obtained by using the test indices for
129-
# the trial space as well.
130-
a, _ = argument_multiindices
131-
argument_multiindices = (a, a)
132-
133-
return_variables = builder.set_arguments(arguments, argument_multiindices)
134-
135128
builder.set_coordinates(mesh)
136129
builder.set_cell_sizes(mesh)
137-
138130
builder.set_coefficients(integral_data, form_data)
139-
140-
# Map from UFL FiniteElement objects to multiindices. This is
141-
# so we reuse Index instances when evaluating the same coefficient
142-
# multiple times with the same table.
143-
#
144-
# We also use the same dict for the unconcatenate index cache,
145-
# which maps index objects to tuples of multiindices. These two
146-
# caches shall never conflict as their keys have different types
147-
# (UFL finite elements vs. GEM index objects).
148-
index_cache = {}
149-
150-
kernel_cfg = dict(interface=builder,
151-
ufl_cell=cell,
152-
integral_type=integral_type,
153-
integration_dim=integration_dim,
154-
entity_ids=entity_ids,
155-
argument_multiindices=argument_multiindices,
156-
index_cache=index_cache,
157-
scalar_type=parameters["scalar_type"])
158-
159-
mode_irs = collections.OrderedDict()
131+
ctx = builder.create_context()
160132
for integral in integral_data.integrals:
161133
params = parameters.copy()
162134
params.update(integral.metadata()) # integral metadata overrides
163-
if params.get("quadrature_rule") == "default":
164-
del params["quadrature_rule"]
165-
166-
mode = pick_mode(params["mode"])
167-
mode_irs.setdefault(mode, collections.OrderedDict())
168-
169135
integrand = ufl.replace(integral.integrand(), form_data.function_replace_map)
170-
integrand = ufl_utils.split_coefficients(integrand, builder.coefficient_split)
171-
172-
# Check if the integral has a quad degree attached, otherwise use
173-
# the estimated polynomial degree attached by compute_form_data
174-
quadrature_degree = params.get("quadrature_degree",
175-
params["estimated_polynomial_degree"])
176-
try:
177-
quadrature_degree = params["quadrature_degree"]
178-
except KeyError:
179-
quadrature_degree = params["estimated_polynomial_degree"]
180-
functions = list(arguments) + [builder.coordinate(mesh)] + list(integral_data.integral_coefficients)
181-
function_degrees = [f.ufl_function_space().ufl_element().degree() for f in functions]
182-
if all((asarray(quadrature_degree) > 10 * asarray(degree)).all()
183-
for degree in function_degrees):
184-
logger.warning("Estimated quadrature degree %s more "
185-
"than tenfold greater than any "
186-
"argument/coefficient degree (max %s)",
187-
quadrature_degree, max_degree(function_degrees))
188-
189-
try:
190-
quad_rule = params["quadrature_rule"]
191-
except KeyError:
192-
integration_cell = fiat_cell.construct_subelement(integration_dim)
193-
quad_rule = make_quadrature(integration_cell, quadrature_degree)
194-
195-
if not isinstance(quad_rule, AbstractQuadratureRule):
196-
raise ValueError("Expected to find a QuadratureRule object, not a %s" %
197-
type(quad_rule))
198-
199-
quadrature_multiindex = quad_rule.point_set.indices
200-
quadrature_indices.extend(quadrature_multiindex)
201-
202-
config = kernel_cfg.copy()
203-
config.update(quadrature_rule=quad_rule)
204-
expressions = fem.compile_ufl(integrand,
205-
fem.PointSetContext(**config),
206-
interior_facet=interior_facet)
207-
reps = mode.Integrals(expressions, quadrature_multiindex,
208-
argument_multiindices, params)
209-
for var, rep in zip(return_variables, reps):
210-
mode_irs[mode].setdefault(var, []).append(rep)
211-
212-
# Finalise mode representations into a set of assignments
213-
assignments = []
214-
for mode, var_reps in mode_irs.items():
215-
assignments.extend(mode.flatten(var_reps.items(), index_cache))
216-
217-
if assignments:
218-
return_variables, expressions = zip(*assignments)
219-
else:
220-
return_variables = []
221-
expressions = []
222-
223-
# Need optimised roots
224-
options = dict(reduce(operator.and_,
225-
[mode.finalise_options.items()
226-
for mode in mode_irs.keys()]))
227-
expressions = impero_utils.preprocess_gem(expressions, **options)
228-
assignments = list(zip(return_variables, expressions))
229-
230-
# Let the kernel interface inspect the optimised IR to register
231-
# what kind of external data is required (e.g., cell orientations,
232-
# cell sizes, etc.).
233-
builder.register_requirements(expressions)
234-
235-
# Construct ImperoC
236-
split_argument_indices = tuple(chain(*[var.index_ordering()
237-
for var in return_variables]))
238-
index_ordering = tuple(quadrature_indices) + split_argument_indices
239-
try:
240-
impero_c = impero_utils.compile_gem(assignments, index_ordering, remove_zeros=True)
241-
except impero_utils.NoopError:
242-
# No operations, construct empty kernel
243-
return builder.construct_empty_kernel(kernel_name)
244-
245-
# Generate COFFEE
246-
index_names = []
247-
248-
def name_index(index, name):
249-
index_names.append((index, name))
250-
if index in index_cache:
251-
for multiindex, suffix in zip(index_cache[index],
252-
string.ascii_lowercase):
253-
name_multiindex(multiindex, name + suffix)
254-
255-
def name_multiindex(multiindex, name):
256-
if len(multiindex) == 1:
257-
name_index(multiindex[0], name)
258-
else:
259-
for i, index in enumerate(multiindex):
260-
name_index(index, name + str(i))
136+
integrand_exprs = builder.compile_integrand(integrand, params, ctx)
137+
integral_exprs = builder.construct_integrals(integrand_exprs, params)
138+
builder.stash_integrals(integral_exprs, params, ctx)
139+
return builder.construct_kernel(kernel_name, ctx)
261140

262-
name_multiindex(quadrature_indices, 'ip')
263-
for multiindex, name in zip(argument_multiindices, ['j', 'k']):
264-
name_multiindex(multiindex, name)
265141

266-
return builder.construct_kernel(kernel_name, impero_c, index_names, quad_rule)
142+
def preprocess_parameters(parameters):
143+
if parameters is None:
144+
parameters = default_parameters()
145+
else:
146+
_ = default_parameters()
147+
_.update(parameters)
148+
parameters = _
149+
# Remove these here, they're handled later on.
150+
if parameters.get("quadrature_degree") in ["auto", "default", None, -1, "-1"]:
151+
del parameters["quadrature_degree"]
152+
if parameters.get("quadrature_rule") in ["auto", "default", None]:
153+
del parameters["quadrature_rule"]
154+
return parameters
267155

268156

269157
def compile_expression_dual_evaluation(expression, to_element, ufl_element, *,
@@ -429,71 +317,3 @@ def __call__(self, ps):
429317
assert set(gem_expr.free_indices) <= set(chain(ps.indices, *argument_multiindices))
430318

431319
return gem_expr
432-
433-
434-
def lower_integral_type(fiat_cell, integral_type):
435-
"""Lower integral type into the dimension of the integration
436-
subentity and a list of entity numbers for that dimension.
437-
438-
:arg fiat_cell: FIAT reference cell
439-
:arg integral_type: integral type (string)
440-
"""
441-
vert_facet_types = ['exterior_facet_vert', 'interior_facet_vert']
442-
horiz_facet_types = ['exterior_facet_bottom', 'exterior_facet_top', 'interior_facet_horiz']
443-
444-
dim = fiat_cell.get_dimension()
445-
if integral_type == 'cell':
446-
integration_dim = dim
447-
elif integral_type in ['exterior_facet', 'interior_facet']:
448-
if isinstance(fiat_cell, TensorProductCell):
449-
raise ValueError("{} integral cannot be used with a TensorProductCell; need to distinguish between vertical and horizontal contributions.".format(integral_type))
450-
integration_dim = dim - 1
451-
elif integral_type == 'vertex':
452-
integration_dim = 0
453-
elif integral_type in vert_facet_types + horiz_facet_types:
454-
# Extrusion case
455-
if not isinstance(fiat_cell, TensorProductCell):
456-
raise ValueError("{} integral requires a TensorProductCell.".format(integral_type))
457-
basedim, extrdim = dim
458-
assert extrdim == 1
459-
460-
if integral_type in vert_facet_types:
461-
integration_dim = (basedim - 1, 1)
462-
elif integral_type in horiz_facet_types:
463-
integration_dim = (basedim, 0)
464-
else:
465-
raise NotImplementedError("integral type %s not supported" % integral_type)
466-
467-
if integral_type == 'exterior_facet_bottom':
468-
entity_ids = [0]
469-
elif integral_type == 'exterior_facet_top':
470-
entity_ids = [1]
471-
else:
472-
entity_ids = list(range(len(fiat_cell.get_topology()[integration_dim])))
473-
474-
return integration_dim, entity_ids
475-
476-
477-
def pick_mode(mode):
478-
"Return one of the specialized optimisation modules from a mode string."
479-
try:
480-
from firedrake_citations import Citations
481-
cites = {"vanilla": ("Homolya2017", ),
482-
"coffee": ("Luporini2016", "Homolya2017", ),
483-
"spectral": ("Luporini2016", "Homolya2017", "Homolya2017a"),
484-
"tensor": ("Kirby2006", "Homolya2017", )}
485-
for c in cites[mode]:
486-
Citations().register(c)
487-
except ImportError:
488-
pass
489-
if mode == "vanilla":
490-
import tsfc.vanilla as m
491-
elif mode == "coffee":
492-
import tsfc.coffee_mode as m
493-
elif mode == "spectral":
494-
import tsfc.spectral as m
495-
elif mode == "tensor":
496-
import tsfc.tensor as m
497-
else:
498-
raise ValueError("Unknown mode: {}".format(mode))
499-
return m

0 commit comments

Comments
 (0)