Skip to content
172 changes: 85 additions & 87 deletions src/dolfinx_adjoint/blocks/function_assigner.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,19 +2,26 @@

import dolfinx
import numpy as np
import numpy.typing as npt
import ufl
from pyadjoint import AdjFloat, Block, OverloadedType
from ufl.corealg.traversal import traverse_unique_terminals
from ufl.formatting.ufl2unicode import ufl2unicode

from ..utils import function_from_vector
from ..utils import assign_linear_combination, extract_linear_combination, function_from_vector
from ._vector import _vector


class FunctionAssignBlock(Block):
"""Block for assigning data directly to a `Function` on the tape.

This block handles the assignment of a linear combination of `Function`s
or constants to a target `Function`.
"""
Comment thread
jorgensd marked this conversation as resolved.
Outdated

def __init__(
self,
other: typing.Union[np.inexact, int, float],
func: dolfinx.fem.Function,
Comment thread
jorgensd marked this conversation as resolved.
ad_block_tag: typing.Optional[str] = None,
):
super().__init__(ad_block_tag=ad_block_tag)
Expand All @@ -25,15 +32,15 @@ def __init__(
elif isinstance(other, float) or isinstance(other, int):
other = AdjFloat(other)
self.add_dependency(other, no_duplicates=True)
elif not (isinstance(other, float) or isinstance(other, int)):
raise NotImplementedError("This should eventually be supported")
# # Assume that this is a point-wise evaluated UFL expression (firedrake only)
# for op in traverse_unique_terminals(other):
# if isinstance(op, OverloadedType):
# self.add_dependency(op, no_duplicates=True)
# self.expr = other
else:
raise NotImplementedError("We should not get here!")
# Extract linear combination
lin_comb = extract_linear_combination(other)
if len(lin_comb) == 0:
raise ValueError("No linear combination found in the expression.")
for op in traverse_unique_terminals(other):
if isinstance(op, OverloadedType):
self.add_dependency(op, no_duplicates=True)
self.expr = other

def _replace_with_saved_output(self):
if self.expr is None:
Expand All @@ -53,85 +60,76 @@ def prepare_evaluate_adj(self, inputs, adj_inputs, relevant_dependencies):
expr = self._replace_with_saved_output()
return expr, adj_input_func

def _compute_adjoint_of_broadcast(self, input: dolfinx.la.Vector | npt.NDArray | float | int) -> float | int:
Comment thread
jorgensd marked this conversation as resolved.
Outdated
# Adjoint of a broadcast is just a sum
if isinstance(input, dolfinx.la.Vector):
one = dolfinx.la.vector(input.index_map, input.block_size, input.array.dtype)
one.array[:] = 1
return dolfinx.cpp.la.inner_product(input._cpp_object, one._cpp_object) # type: ignore[arg-type]
else:
if hasattr(input, "sum"):
return input.sum()
else:
# Catch the case where input is just a float
return input

def evaluate_adj_component(self, inputs, adj_inputs, block_variable, idx, prepared=None):
bo = block_variable.output
if self.expr is None:
if isinstance(block_variable.output, AdjFloat):
# Adjoint of a broadcast is just a sum
if isinstance(adj_inputs[0], dolfinx.la.Vector):
vec = adj_inputs[0]
one = dolfinx.la.vector(
adj_inputs[0].index_map, adj_inputs[0].block_size, adj_inputs[0].array.dtype
)
one.array[:] = 1
return dolfinx.cpp.la.inner_product(vec._cpp_object, one._cpp_object)
else:
try:
return adj_inputs[0].sum()
except AttributeError:
# Catch the case where adj_inputs[0] is just a float
return adj_inputs[0]
elif isinstance(func := block_variable.output, dolfinx.fem.Function):
assert len(adj_inputs) == 1
if isinstance(func := bo, AdjFloat):
return self._compute_adjoint_of_broadcast(adj_inputs[0])
Comment thread
jorgensd marked this conversation as resolved.
Outdated
elif isinstance(func, dolfinx.fem.Function):
assert func.function_space == prepared.function_space
vec = _vector(
prepared.x.index_map, prepared.x.block_size, func.function_space, dtype=prepared.x.array.dtype
)
vec.array[:] = prepared.x.array[:]
return vec
elif isinstance(bo, dolfinx.fem.Constant):
raise NotImplementedError(
"Adjoint for Constant assignment not implemented, use dolfinx_adjoint.Constant instead."
)
else:
raise NotImplementedError(f"Adjoint for {block_variable=} not implemented.")
else:
# Linear combination
expr, adj_input_func = prepared
vec = _vector(
bo.x.index_map,
bo.x.block_size,
bo.function_space,
dtype=bo.x.array.dtype,
)
if isinstance(bo, dolfinx.fem.Function) and bo.function_space == adj_input_func.function_space:
# Differentiate with respect to one of the input functions
diff_expr = ufl.algorithms.expand_derivatives(
ufl.derivative(expr, block_variable.saved_output, adj_input_func)
)
temp_func = dolfinx.fem.Function(bo.function_space)
assign_linear_combination(diff_expr, temp_func)
vec.array[:] = temp_func.x.array[:]
return vec
elif isinstance(bo, dolfinx.fem.Function) and bo.ufl_element().is_real:
# Differentiate with respect to a real function (constant stored as Function)
# Create a perturbation direction in the Real space (value = 1.0)
direction = dolfinx.fem.Function(bo.function_space)
direction.x.array[0] = 1.0

# Differentiate expr w.r.t 'bo' in that direction
diff_expr = ufl.algorithms.expand_derivatives(
ufl.derivative(expr, block_variable.saved_output, direction)
)

# Evaluate the derivative at the DOFs of the target space V
diff_eval = dolfinx.fem.Function(adj_input_func.function_space)
assign_linear_combination(diff_expr, diff_eval)
Comment thread
jorgensd marked this conversation as resolved.
Outdated

# Chain rule: dot product of (dz/dr) and adjoint inputs (bar_u)
vec.array[0] = dolfinx.cpp.la.inner_product(diff_eval.x._cpp_object, adj_input_func.x._cpp_object)
return vec
else:
raise NotImplementedError(f"Adjoint for {block_variable=} not implemented.")
# elif isinstance(block_variable.output, dolfinx.fem.Constant):
# R = block_variable.output._ad_function_space(prepared.function_space.mesh)
# return self._adj_assign_constant(prepared, R)
# else:
# adj_output = dolfinx.fem.Function(
# block_variable.output.function_space())
# adj_output.assign(prepared)
# return adj_output.vector()
# else:
# # Linear combination
# expr, adj_input_func = prepared
# adj_output = dolfinx.fem.Function(adj_input_func.function_space)
# if not isinstance(block_variable.output, dolfinx.fem.Constant):
# diff_expr = ufl.algorithms.expand_derivatives(
# ufl.derivative(expr, block_variable.saved_output, adj_input_func)
# )
# adj_output.assign(diff_expr)
# else:
# mesh = adj_output.function_space().mesh()
# diff_expr = ufl.algorithms.expand_derivatives(
# ufl.derivative(
# expr,
# block_variable.saved_output,
# create_constant(1., domain=mesh)
# )
# )
# adj_output.assign(diff_expr)
# return adj_output.vector().inner(adj_input_func.vector())

# if isinstance(block_variable.output, dolfin.Constant):
# R = block_variable.output._ad_function_space(adj_output.function_space().mesh())
# return self._adj_assign_constant(adj_output, R)
# else:
# return adj_output.vector()

def _adj_assign_constant(self, adj_output, constant_fs):
r = dolfinx.fem.Function(constant_fs)
shape = r.ufl_shape
raise NotImplementedError("Not implemented for constants.")

if shape == () or shape[0] == 1:
# Scalar Constant
raise NotImplementedError("Not implemented for scalar constants yet.")
# r.vector()[:] = adj_output.vector().sum()
# else:
# # We assume the shape of the constant == shape of the output function if not scalar.
# # This assumption is due to FEniCS not supporting products with non-scalar constants in assign.
# values = []
# for i in range(shape[0]):
# values.append(adj_output.sub(i, deepcopy=True).vector().sum())
# r.assign(dolfin.Constant(values))
return r.vector()

def prepare_evaluate_tlm(self, inputs, tlm_inputs, relevant_outputs):
if self.expr is None:
Expand All @@ -144,11 +142,13 @@ def evaluate_tlm_component(self, inputs, tlm_inputs, block_variable, idx, prepar
return tlm_inputs[0]
expr = prepared
dudm = dolfinx.fem.Function(block_variable.output.function_space)
dudm.x.array[:] = 0.0
dudmi = dolfinx.fem.Function(block_variable.output.function_space)
for dep in self.get_dependencies():
if dep.tlm_value:
dudmi.assign(ufl.algorithms.expand_derivatives(ufl.derivative(expr, dep.saved_output, dep.tlm_value)))
dudm.vector().axpy(1.0, dudmi.vector())
diff_expr = ufl.algorithms.expand_derivatives(ufl.derivative(expr, dep.saved_output, dep.tlm_value))
assign_linear_combination(diff_expr, dudmi)
dudm.x.array[:] += dudmi.x.array[:]

return dudm

Expand Down Expand Up @@ -181,14 +181,12 @@ def recompute_component(self, inputs, block_variable, idx, prepared):
# We should return the exact object instance to maintain C++ memory bindings
# (especially for DirichletBCs), updating it in-place.
output = block_variable.saved_output

try:
if output.function_space == prepared.function_space:
output.x.array[:] = prepared.x.array[:]
except AttributeError:
# Handling float value
if isinstance(prepared, dolfinx.fem.Function):
output.x.array[:] = prepared.x.array[:]
elif isinstance(prepared, (float, int)):
output.x.array[:] = prepared

else:
assign_linear_combination(prepared, output)
return output

def __str__(self):
Expand Down
7 changes: 5 additions & 2 deletions src/dolfinx_adjoint/types/function.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@

from ..blocks.assembly import assemble_compiled_form
from ..blocks.function_assigner import FunctionAssignBlock
from ..utils import ad_kwargs, function_from_vector, gather
from ..utils import ad_kwargs, assign_linear_combination, function_from_vector, gather


class Function(dolfinx.fem.Function, FloatingType):
Expand Down Expand Up @@ -273,7 +273,7 @@ def assign(value: typing.Union[numpy.inexact, float, int], function: Function, *
if annotate:
if not isinstance(value, ufl.core.operator.Operator):
value = create_overloaded_object(value)
block = FunctionAssignBlock(value, function, ad_block_tag=ad_block_tag)
block = FunctionAssignBlock(value, ad_block_tag=ad_block_tag)
tape = get_working_tape()
tape.add_block(block)

Expand All @@ -285,6 +285,9 @@ def assign(value: typing.Union[numpy.inexact, float, int], function: Function, *
"Function spaces of the value and function must match for assignment."
)
function.x.array[:] = value.x.array[:]
elif isinstance(value, ufl.core.expr.Expr):
# Linear combination of functions, e.g., 2*u + 3*v
assign_linear_combination(value, function)
else:
raise ValueError(f"Unsupported value type for assignment: {type(value)})")
if annotate:
Expand Down
126 changes: 126 additions & 0 deletions src/dolfinx_adjoint/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import dolfinx
import numpy
import numpy.typing as npt
import ufl


def function_from_vector(
Expand Down Expand Up @@ -41,3 +42,128 @@ class ad_kwargs(typing.TypedDict):
"""Tag for the block in the adjoint tape."""
annotate: typing.NotRequired[bool]
"""Whether to annotate the assignment in the adjoint tape."""


def extract_scalar_value(scalar_expr):
Comment thread
jorgensd marked this conversation as resolved.
Outdated
"""Extract float from a scalar UFL expression."""
if isinstance(scalar_expr, (ufl.classes.IntValue, ufl.classes.FloatValue)):
return float(scalar_expr)
elif isinstance(scalar_expr, dolfinx.fem.Function):
# Check if it's a RealElement (constant stored as Function)
if scalar_expr.function_space.ufl_element().is_real and scalar_expr.ufl_shape == ():
return float(scalar_expr.x.array[0])
else:
raise ValueError(f"Cannot extract scalar from spatial Function: {scalar_expr}")
elif isinstance(scalar_expr, dolfinx.fem.Constant) and scalar_expr.ufl_shape == ():
val = scalar_expr.value
return float(val) if hasattr(val, "__float__") else float(val.item())
elif isinstance(scalar_expr, ufl.classes.ScalarValue):
return float(scalar_expr._value)
elif isinstance(scalar_expr, ufl.classes.Product):
result = 1.0
for op in scalar_expr.ufl_operands:
result *= extract_scalar_value(op)
return result
elif isinstance(scalar_expr, ufl.classes.Division):
num, den = scalar_expr.ufl_operands
return extract_scalar_value(num) / extract_scalar_value(den)
else:
raise ValueError(f"Cannot extract scalar from {type(scalar_expr)}: {scalar_expr}")


def extract_function(expr) -> tuple[bool, dolfinx.fem.Function | None]:
"""Recursively extract a Function from nested UFL expressions."""
Comment thread
jorgensd marked this conversation as resolved.
Outdated
if isinstance(expr, dolfinx.fem.Function):
is_real = expr.function_space.ufl_element().is_real
if is_real:
return (False, None)
return (False, expr)
elif isinstance(expr, (ufl.classes.Indexed, ufl.classes.ComponentTensor)):
return extract_function(expr.ufl_operands[0])
elif hasattr(expr, "ufl_operands"):
found_func = None
for op in expr.ufl_operands:
is_real, func = extract_function(op)
if func is not None:
if found_func is not None:
raise ValueError(f"Non-linear expression detected: multiple spatial functions in {expr}")
found_func = func
return (False, found_func)
return (False, None)


def extract_term(term):
Comment thread
jorgensd marked this conversation as resolved.
Outdated
"""Extract (weight, function) from a single term."""
if isinstance(term, dolfinx.fem.Function):
is_real = term.function_space.ufl_element().is_real
if is_real:
return None
return (1.0, term)
elif isinstance(term, ufl.classes.ComponentTensor):
return extract_term(term.ufl_operands[0])
elif isinstance(term, ufl.classes.Indexed):
is_real, func = extract_function(term)
if func is None:
return None
return (1.0, func)
elif isinstance(term, ufl.classes.Product):
weight = 1.0
func = None
for op in term.ufl_operands:
is_real, extracted_func = extract_function(op)
if extracted_func is not None:
if func is not None:
raise ValueError(f"Non-linear term detected: multiple spatial functions in {term}")
func = extracted_func
else:
weight *= extract_scalar_value(op)
return (weight, func) if func is not None else None
elif isinstance(term, ufl.classes.Division):
num, den = term.ufl_operands
denom_val = extract_scalar_value(den)
if isinstance(num, dolfinx.fem.Function):
is_real = num.function_space.ufl_element().is_real
if is_real:
return None
return (1.0 / denom_val, num)
elif isinstance(num, ufl.classes.Product):
result = extract_term(num)
return (result[0] / denom_val, result[1]) if result else None
return None


def extract_linear_combination(expr: ufl.core.expr.Expr) -> list[tuple[float, dolfinx.fem.Function]]:
"""Extract (weight, function) pairs from a UFL linear combination.

Analyzes expressions like: 0.5*u + 0.3*v + 0.2*w
Returns: [(0.5, u), (0.3, v), (0.2, w)]

:param expr: UFL expression (Sum, Product, or single Function)
:returns: List of (weight, function) tuples
Comment thread
jorgensd marked this conversation as resolved.
Outdated
"""

# Parse the expression, flattening nested Sums recursively
if isinstance(expr, ufl.classes.Sum):
summands = expr.ufl_operands
else:
summands = [expr]
terms = []
for summand in summands:
if isinstance(summand, ufl.classes.Sum):
# Recursively flatten nested Sum structures
terms.extend(extract_linear_combination(summand))
else:
result = extract_term(summand)
if result is not None:
terms.append(result)
return terms


def assign_linear_combination(value: ufl.core.expr.Expr, function: dolfinx.fem.Function):
Comment thread
jorgensd marked this conversation as resolved.
Outdated
pairs = extract_linear_combination(value)
function.x.array[:] = 0.0
for weight, func in pairs:
if not func.function_space == function.function_space:
raise ValueError("Function spaces of all functions in the linear combination must match for assignment.")
function.x.array[:] += weight * func.x.array[:]
function.x.scatter_forward()
Loading
Loading