Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 38 additions & 31 deletions finat/enriched.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,7 @@ def merge(tables):
tables = tuple(tables)
zeta = self.get_value_indices()
tensors = []
for elem, table in zip(self.elements, tables):
for elem, table in zip(self.summands, tables):
beta_i = elem.get_indices()
tensors.append(gem.ComponentTensor(
gem.Indexed(table, beta_i + zeta),
Expand All @@ -138,7 +138,7 @@ def basis_evaluation(self, order, ps, entity=None, coordinate_mapping=None):
:param entity: the cell entity on which to tabulate.
'''
results = [element.basis_evaluation(order, ps, entity, coordinate_mapping=coordinate_mapping)
for element in self.elements]
for element in self.summands]
return self._compose_evaluations(results)

def point_evaluation(self, order, refcoords, entity=None, coordinate_mapping=None):
Expand All @@ -153,7 +153,7 @@ def point_evaluation(self, order, refcoords, entity=None, coordinate_mapping=Non
:param entity: the cell entity on which to tabulate.
'''
results = [element.point_evaluation(order, refcoords, entity, coordinate_mapping)
for element in self.elements]
for element in self.summands]
return self._compose_evaluations(results)

@property
Expand All @@ -166,20 +166,17 @@ def mapping(self):
return result

@cached_property
def _summands(self):
def summands(self):
"""The summands that are not themselves direct sums, in basis order.

An element is brought out as a direct sum one level at a time, so a
summand may be a direct sum in turn. These are the elements that
evaluate their dual basis on their own points, and whose points make
up the union that :attr:`dual_basis` works against.
An element is brought out as a direct sum one level at a time. A summand
of :attr:`elements` may therefore be a direct sum in turn, and these are
the elements that remain once every level is brought out. They are the
elements that evaluate their dual basis on their own points, and whose
points make up the union that :attr:`dual_basis` works against.
"""
summands = []
for element in self.elements:
expanded = as_enriched(element)
summands.extend(expanded._summands if expanded is not None
else [element])
return tuple(summands)
return tuple(chain.from_iterable(element.summands
for element in self.elements))

@property
def dual_basis(self):
Expand All @@ -205,12 +202,12 @@ def dual_basis(self):
f"Dual basis not defined for non-nodal {type(self).__name__}"
)
if any(type(e).dual_transformation is not FiniteElementBase.dual_transformation
for e in self._summands):
for e in self.summands):
raise NotImplementedError(
f"dual_basis not defined for {type(self).__name__} with a summand"
" that has its own dual_transformation; use dual_evaluation instead"
)
duals = [element.dual_basis for element in self._summands]
duals = [element.dual_basis for element in self.summands]
x = UnionPointSet([xk for _, xk in duals])
p, = x.indices
zeta = self.get_value_indices()
Expand All @@ -219,7 +216,7 @@ def dual_basis(self):
shapes = [tuple(i.extent for i in xk.indices) for _, xk in duals]

blocks = []
for k, (element, (Q, xk)) in enumerate(zip(self._summands, duals)):
for k, (element, (Q, xk)) in enumerate(zip(self.summands, duals)):
alpha = element.get_indices()
# Turn this summand's point indices into a shape, so that its
# weights can be embedded at its own offset in the union.
Expand All @@ -245,29 +242,31 @@ def _dual_evaluation(self, fn, coordinate_mapping=None):
provides physical geometry callbacks (may be None).
:returns: an ``(evaluation, point_indices, basis_indices)`` triple, as
:meth:`~finat.finiteelementbase.FiniteElementBase.dual_evaluation`
returns. The points are contracted here, so ``point_indices`` is
empty.
returns. The summand point indices remain free, so the caller can
choose how to contract each direct-sum component.

The summands do not share their points, so each one contracts on its
own, and the results stack along the basis index. Concatenating over
a free index is what :func:`~gem.unconcatenate.unconcatenate` splits
downstream; a concatenation over the contracted points could not be.
The summands do not share their points, so their evaluations stack
along the basis index while retaining their own point indices.
Concatenating over a free basis index is what
:func:`~gem.unconcatenate.unconcatenate` splits downstream.
"""
if not self.is_nodal_enriched:
raise NotImplementedError(
f"Dual evaluation not defined for non-nodal {type(self).__name__}"
)
# Each summand contracts through its own dual_basis, so a non-nodal
# sum has to be refused here as well as in dual_basis: this path
# never asks self for one.
# Each summand uses its own dual_basis, so a non-nodal sum has to be
# refused here as well as in dual_basis: this path never asks self
# for one.
evals = []
for element in self.elements:
expr, point_indices, indices = element.dual_evaluation(
point_indices = []
for element in self.summands:
expr, element_points, indices = element.dual_evaluation(
fn, coordinate_mapping=coordinate_mapping)
evals.append(broadcast_tensor(gem.IndexSum(expr, point_indices), indices))
evals.append(broadcast_tensor(expr, indices))
point_indices.extend(element_points)

beta = self.get_indices()
return gem.Indexed(gem.Concatenate(*evals), beta), (), beta
return gem.Indexed(gem.Concatenate(*evals), beta), tuple(dict.fromkeys(point_indices)), beta


@singledispatch
Expand All @@ -290,7 +289,15 @@ def as_enriched_enriched(element):

@as_enriched.register(FlattenedDimensions)
def as_enriched_flattened(element):
return as_enriched(element.product)
"""Distribute the flattening over the sum the product is.

Each summand keeps the cell of the element that it came out of. The
summands therefore tabulate against the same entities as that element.
"""
summands = as_enriched(element.product)
if summands is None:
return None
return distribute_over_sum(FlattenedDimensions, summands)


@as_enriched.register(DiscontinuousElement)
Expand Down
22 changes: 22 additions & 0 deletions finat/finiteelementbase.py
Original file line number Diff line number Diff line change
Expand Up @@ -291,6 +291,28 @@ def dual_basis(self):
f"Dual basis not defined for element {type(self).__name__}"
)

@cached_property
def summands(self):
"""The direct summands whose bases stack into this element's basis.

A direct sum blocks its tabulation and its dual basis along these
summands alike. A contraction of the one against the other therefore
splits into a sum over them.

Returns
-------
tuple
The elements, on this element's own cell and in basis order, whose
tabulations concatenate into this element's tabulation and whose
dual bases stack into its dual basis. An element that is not a
direct sum is its own only summand.
"""
from finat.enriched import as_enriched # Avoid circular import
summands = as_enriched(self)
if summands is None:
return (self,)
return summands.summands

def dual_evaluation(self, fn, coordinate_mapping=None):
'''Get a GEM expression for performing the dual basis evaluation at
the nodes of the reference element. Currently only works for flat
Expand Down
6 changes: 3 additions & 3 deletions gem/optimise.py
Original file line number Diff line number Diff line change
Expand Up @@ -194,9 +194,9 @@ def _constant_fold_zero(node, self):

@_constant_fold_zero.register(Literal)
def _constant_fold_zero_literal(node, self):
if numpy.array_equal(node.array, 0):
# All zeros, make symbolic zero
return Zero(node.shape)
if not node.array.any():
# A table of any shape that holds only zeros is a symbolic zero.
return Zero(node.shape, dtype=node.dtype)
else:
return node

Expand Down
114 changes: 96 additions & 18 deletions gem/unconcatenate.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@
from gem.interpreter import evaluate


__all__ = ['flatten', 'unconcatenate']
__all__ = ['flatten', 'split_contraction', 'unconcatenate']


def find_group(expressions, splittable_indices):
Expand Down Expand Up @@ -175,16 +175,24 @@ def replace_node(expression, mapping, cut=None):
return mapper(expression)


def _unconcatenate(cache, pairs):
# Tail-call recursive core of unconcatenate.
# Assumes that input has already been sanitised.
# Only an index carried by an assignment variable can be split against it.
splittable = set().union(chain(*[v.free_indices for v, e in pairs]))
concat_group = find_group([e for v, e in pairs], splittable)
if concat_group is None:
return pairs
def split_group(cache, concat_group):
"""Splits a group of indexed Concatenate nodes into their blocks.

Parameters
----------
cache
Index splitting cache :py:class:`dict`.
concat_group
A group of indexed :py:class:`Concatenate` nodes, as
:py:func:`find_group` returns.

# Get the index split
Returns
-------
tuple
The index that the group shares, one multiindex for each block, and
one substitution for each block. A substitution replaces every node
of the group by that block of it.
"""
concat_ref = next(iter(concat_group))
assert isinstance(concat_ref, Indexed)
concat_expr, = concat_ref.children
Expand All @@ -197,19 +205,31 @@ def _unconcatenate(cache, pairs):
for child in concat_expr.children)
cache[index] = multiindices

def cut(node):
"""No need to rebuild expression of independent of the
relevant concatenation index."""
return index not in node.free_indices

# Build Concatenate node replacement mappings
mappings = [{} for i in range(len(multiindices))]
for concat_ref in concat_group:
concat_expr, = concat_ref.children
for i in range(len(multiindices)):
sub_ref = Indexed(concat_expr.children[i], multiindices[i])
for i, multiindex in enumerate(multiindices):
sub_ref = Indexed(concat_expr.children[i], multiindex)
sub_ref, = remove_componenttensors((sub_ref,))
mappings[i][concat_ref] = sub_ref
return index, multiindices, mappings


def _unconcatenate(cache, pairs):
# Tail-call recursive core of unconcatenate.
# Assumes that input has already been sanitised.
# Only an index carried by an assignment variable can be split against it.
splittable = set().union(chain(*[v.free_indices for v, e in pairs]))
concat_group = find_group([e for v, e in pairs], splittable)
if concat_group is None:
return pairs

index, multiindices, mappings = split_group(cache, concat_group)

def cut(node):
"""No need to rebuild expression of independent of the
relevant concatenation index."""
return index not in node.free_indices

# Finally, split assignment pairs
split_pairs = []
Expand All @@ -224,6 +244,64 @@ def cut(node):
return _unconcatenate(cache, split_pairs)


def _split_contraction(cache, expression, indices):
# Tail-call recursive core of split_contraction.
# Assumes that input has already been sanitised.
concat_group = find_group([expression], set(indices))
if concat_group is None:
return [(expression, indices)]

index, multiindices, mappings = split_group(cache, concat_group)

def cut(node):
"""No need to rebuild expression of independent of the
relevant concatenation index."""
return index not in node.free_indices

# Split the contraction, one block at a time
rest = tuple(i for i in indices if i != index)
terms = []
for multiindex, mapping in zip(multiindices, mappings):
terms.extend(_split_contraction(cache, replace_node(expression, mapping, cut),
rest + multiindex))
return terms


def split_contraction(expression, indices, cache=None):
"""Splits a contraction along the :py:class:`Concatenate` nodes it sums over.

No assignment variable need carry the concatenation index here. The sum
is what the Concatenate splits against. A sum over a whole direct sum is
the sum of the sums over its blocks:

sum_j Indexed(Concatenate(A, B), (j,)) * Indexed(Concatenate(C, D), (j,))
= sum_{ja} A_ja * C_ja + sum_{jb} B_jb * D_jb.

Every Concatenate that one index indexes must concatenate the same blocks.
A FInAT element gives that guarantee: it blocks its tabulation and its dual
basis along the same summands.

Parameters
----------
expression
A scalar GEM expression.
indices
The multiindex that ``expression`` is summed over.
cache
Index splitting cache :py:class:`dict` (optional).

Returns
-------
list
The (expression, multiindex) pairs whose index sums add up to the
index sum of ``expression`` over ``indices``.
"""
if cache is None:
cache = {}
expression, = remove_componenttensors([expression])
return _split_contraction(cache, expression, tuple(indices))


def unconcatenate(pairs, cache=None):
"""Splits a list of (indexed variable, expression) pairs along
:py:class:`Concatenate` nodes embedded in the expressions.
Expand Down
32 changes: 24 additions & 8 deletions test/finat/test_dual_basis.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from finat.quadrature import QuadratureRule
from finat.quadrature_element import QuadratureElement
from gem.interpreter import evaluate
from gem.unconcatenate import unconcatenate
from FIAT import ufc_simplex


Expand Down Expand Up @@ -50,11 +51,28 @@ def tabulate(ps):
table = element.basis_evaluation(0, ps)[(0,) * dim]
return gem.ComponentTensor(gem.Indexed(table, j + zeta), zeta)

expr, point_indices, indices = element.dual_evaluation(tabulate)
if point_indices:
expr = gem.IndexSum(expr, point_indices)
result, = evaluate([gem.ComponentTensor(expr, indices + j)])
expr, _, indices = element.dual_evaluation(tabulate)
n = element.space_dimension()
strides = tuple(
numpy.prod(tuple(index.extent for index in indices[offset + 1:]), dtype=int)
for offset in range(len(indices))
)
variable = gem.FlexiblyIndexed(
gem.Variable("A", (n,)), ((0, tuple(zip(indices, strides))),)
)
blocks = []
for variable, evaluation in unconcatenate([(variable, expr)]):
point_indices = tuple(
index for index in evaluation.free_indices
if index not in variable.free_indices and index not in j
)
blocks.append(gem.ComponentTensor(
gem.IndexSum(evaluation, point_indices),
variable.index_ordering(),
))
i = gem.Index(extent=n)
evaluation = gem.Indexed(gem.Concatenate(*blocks), (i,))
result, = evaluate([gem.ComponentTensor(evaluation, (i,) + j)])
assert numpy.allclose(result.arr.reshape(n, n), numpy.eye(n))


Expand All @@ -63,10 +81,8 @@ def check_dual_basis(element):
Q, x = element.dual_basis
assert Q.shape == element.index_shape + element.value_shape
assert set(Q.free_indices) == set(x.indices)
summands = as_enriched(element)
if summands is not None:
assert len(x.points) == sum(len(e.dual_basis[1].points)
for e in summands._summands)
assert len(x.points) == sum(len(e.dual_basis[1].points)
for e in element.summands)

i = element.get_indices()
j = element.get_indices()
Expand Down
Loading
Loading