Skip to content
Merged
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
16 changes: 16 additions & 0 deletions test/test_duals.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,6 +188,22 @@ def test_zero_base_form_list_arguments():
assert hash(f) == hash(ZeroBaseForm((v,)))


def test_zero_base_form_signature():
domain_2d = Mesh(LagrangeElement(triangle, 1, (2,)))
f_2d = LagrangeElement(triangle, 1)
V = FunctionSpace(domain_2d, f_2d)

v = TestFunction(V)
f = ZeroBaseForm((v,))

assert f.signature() == ZeroBaseForm((v,)).signature()
assert isinstance(f.signature(), str)

W = FunctionSpace(domain_2d, LagrangeElement(triangle, 2))
w = TestFunction(W)
assert f.signature() != ZeroBaseForm((w,)).signature()


def test_zero_base_form_reconstruct():
# ZeroBaseForm._ufl_expr_reconstruct_ inherited BaseForm's default,
# `type(self)(*operands)`, which unpacks `ufl_operands` into positional
Expand Down
13 changes: 13 additions & 0 deletions ufl/form.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
# Modified by Nacime Bouziani, 2020.
# Modified by Jørgen S. Dokken 2023.

import hashlib
import numbers
import typing
import warnings
Expand Down Expand Up @@ -884,6 +885,7 @@ class ZeroBaseForm(BaseForm):
"_coefficients",
"_domains",
"_hash",
"_signature",
# Pyadjoint compatibility
"form",
"ufl_operands",
Expand All @@ -896,6 +898,7 @@ def __init__(self, arguments):
self._arguments = arguments
self.ufl_operands = arguments
self._hash = None
self._signature = None
self._domains = None
self.form = None

Expand Down Expand Up @@ -927,6 +930,16 @@ def empty(self):
"""Returns whether the ZeroBaseForm has no components, which is always true."""
return True

def signature(self):
"""Return a signature for use with JIT caches."""
if self._signature is None:
renumbering = {domain: i for i, domain in enumerate(self.ufl_domains())}
data = tuple(
argument._ufl_signature_data_(renumbering) for argument in self.arguments()
)
self._signature = hashlib.sha512(str(data).encode("utf-8")).hexdigest()
return self._signature

def __ne__(self, other):
"""Overwrite BaseForm.__neq__ which relies on `equals`."""
return not self == other
Expand Down