"""Residual-first eliminated-recurrence time-memory stepper."""
from __future__ import annotations
import warnings
from math import gamma, isfinite
from typing import Any
import numpy as np
from yonderdrake._firedrake import require_real_float64_petsc
from yonderdrake.time._stepper_lifecycle import StepperLifecycle
from yonderdrake.time.checkpointing import (
load_checkpoint_file,
save_checkpoint_file,
stepper_metadata,
validate_stepper_metadata,
)
from yonderdrake.time.coefficients import exp_neg, phi1
from yonderdrake.time.formulations import AuxiliaryODE, Recurrence
from yonderdrake.time.representations import (
BirkSong,
Diethelm,
Diethelm2022,
FullHistory,
YuanAgrawal,
_SingleExponential,
validate_checkpoint_representation,
)
class _RecurrenceStepper(StepperLifecycle):
"""Advance one physical Firedrake field with constant-memory modes."""
def __init__(
self,
F: Any,
representation: Any,
t: Any,
dt: Any,
u: Any,
*,
formulation: Any = None,
u0: Any = None,
bcs: Any = None,
solver_parameters: Any = None,
appctx: Any = None,
allow_exponential: bool = False,
) -> None:
try:
import firedrake as fd
except ImportError as error:
raise RuntimeError(
"time-memory stepping requires an active Firedrake environment"
) from error
require_real_float64_petsc()
from yonderdrake.time._ufl_marker import (
CaputoDerivativeMarker,
ExponentialMemoryMarker,
RiemannLiouvilleDerivativeMarker,
evaluate_form_at_end_time,
find_time_memory_markers,
replace_time_memory_markers,
)
allowed_representations = (
BirkSong,
Diethelm,
Diethelm2022,
YuanAgrawal,
)
if representation is not None and not isinstance(
representation,
allowed_representations,
):
raise TypeError(
"representation must be BirkSong, Diethelm, "
"Diethelm2022, or YuanAgrawal"
)
if formulation is None:
formulation = Recurrence()
if not isinstance(formulation, Recurrence):
raise NotImplementedError(
"AuxiliaryODE uses the monolithic mixed-space implementation"
)
end_time_form = evaluate_form_at_end_time(F, t, dt)
markers = find_time_memory_markers(end_time_form)
if not markers:
raise ValueError(
"the recurrence stepper requires at least one time-memory marker"
)
exponential_markers = tuple(
marker
for marker in markers
if isinstance(marker, ExponentialMemoryMarker)
)
fractional_markers = tuple(
marker
for marker in markers
if isinstance(
marker,
(CaputoDerivativeMarker, RiemannLiouvilleDerivativeMarker),
)
)
if exponential_markers and not allow_exponential:
raise ValueError(
"ExponentialMemory requires TimeMemoryStepper"
)
if fractional_markers and representation is None:
raise TypeError(
"a representation is required for Caputo and "
"Riemann-Liouville markers"
)
if any(marker.field is not u for marker in markers):
raise ValueError(
"every time-memory marker must wrap the stepper solution u"
)
element = u.function_space().ufl_element()
if element.family() != "Lagrange":
raise NotImplementedError(
"Recurrence stepping supports continuous Lagrange spaces only"
)
self._fd = fd
self.F = F
self.representation: Any = representation
self.formulation = formulation
self.t = t
self.dt = dt
self.u = u
self.bcs = bcs
self._solver_parameters = dict(solver_parameters or {})
self._appctx = dict(appctx or {})
self._markers = markers
self._parameter_operands = tuple(
(
marker.decay_rate
if isinstance(marker, ExponentialMemoryMarker)
else marker.alpha
)
for marker in markers
)
self._operator_kinds = tuple(
(
"exponential_memory"
if isinstance(marker, ExponentialMemoryMarker)
else (
"riemann_liouville"
if isinstance(marker, RiemannLiouvilleDerivativeMarker)
else "caputo"
)
)
for marker in markers
)
single_exponential = _SingleExponential()
self._term_representations: tuple[Any, ...] = tuple(
(
single_exponential
if kind == "exponential_memory"
else representation
)
for kind in self._operator_kinds
)
self.parameters = self._read_parameters()
self.spectra = tuple(
term_representation.spectrum(parameter)
for term_representation, parameter in zip(
self._term_representations,
self.parameters,
strict=True,
)
)
self.decay_rates = tuple(
parameter
for kind, parameter in zip(
self._operator_kinds,
self.parameters,
strict=True,
)
if kind == "exponential_memory"
)
self._operator_kind = self._operator_kinds[0]
self.spectrum = (
self.spectra[0] if len(self.spectra) == 1 else self.spectra
)
self._space = u.function_space()
self._lower_limit = self._read_time()
if u0 is not None:
self.u.assign(u0)
self._previous = fd.Function(self._space, name="time_memory_previous")
self._previous.assign(self.u)
self._initial = fd.Function(self._space, name="time_memory_initial")
self._initial.assign(self.u)
# Reused rollback snapshot.
self._committed_u = fd.Function(
self._space,
name="time_memory_committed",
)
self._increment = fd.Function(
self._space,
name="time_memory_increment",
)
self._mode_groups = tuple(
tuple(
fd.Function(
self._space,
name=f"time_memory_term_{term:02d}_mode_{index:04d}",
)
for index in range(spectrum.rates.size)
)
for term, spectrum in enumerate(self.spectra)
)
self._history_terms = tuple(
fd.Function(
self._space,
name=f"time_memory_history_term_{term:02d}",
)
for term in range(len(markers))
)
for history_term in self._history_terms:
history_term.assign(0.0)
self._implicit_weights = tuple(
fd.Constant(0.0) for _ in markers
)
self._initial_trace_terms = tuple(
fd.Function(
self._space,
name=f"fractional_initial_trace_term_{term:02d}",
)
for term in range(len(markers))
)
for initial_trace_term in self._initial_trace_terms:
initial_trace_term.assign(0.0)
replacements = {
marker: (
history_term
+ implicit_weight * (self.u - self._previous)
+ initial_trace_term
)
for marker, history_term, implicit_weight, initial_trace_term in zip(
markers,
self._history_terms,
self._implicit_weights,
self._initial_trace_terms,
strict=True,
)
}
self._transformed_residual = replace_time_memory_markers(
end_time_form,
replacements.__getitem__,
)
self._solver: Any = None
self._rebuild_solver = True
self._last_step_size: float | None = None
self._coefficient_step_size: float | None = None
self._recurrence_coefficients: (
tuple[tuple[np.ndarray, np.ndarray, np.ndarray], ...] | None
) = None
self._reset_solver_counters()
@staticmethod
def _read_parameter_operand(
parameter_operand: Any,
operator_kind: str,
) -> float:
try:
value = float(parameter_operand)
except (TypeError, ValueError) as error:
name = (
"decay_rate"
if operator_kind == "exponential_memory"
else "alpha"
)
raise TypeError(f"{name} must be a real scalar") from error
if operator_kind == "exponential_memory":
if not isfinite(value) or value <= 0.0:
raise ValueError("decay_rate must be finite and positive")
return value
if not isfinite(value) or not 0.0 < value < 1.0:
raise ValueError("alpha must satisfy 0 < alpha < 1")
return value
def _read_parameters(self) -> tuple[float, ...]:
return tuple(
self._read_parameter_operand(operand, kind)
for operand, kind in zip(
self._parameter_operands,
self._operator_kinds,
strict=True,
)
)
def _update_initial_trace(self, step_size: float) -> None:
if "riemann_liouville" not in self._operator_kinds:
return
elapsed = self._read_time() + step_size - self._lower_limit
if not isfinite(elapsed) or elapsed <= 0.0:
raise ValueError(
"Riemann-Liouville evaluation time must exceed its lower limit"
)
for kind, parameter, initial_trace_term in zip(
self._operator_kinds,
self.parameters,
self._initial_trace_terms,
strict=True,
):
if kind == "riemann_liouville":
scale = elapsed ** (-parameter) / gamma(1.0 - parameter)
initial_trace_term.assign(scale * self._initial)
def _build_solver(self) -> None:
problem = self._fd.NonlinearVariationalProblem(
self._transformed_residual,
self.u,
bcs=self.bcs,
)
self._solver = self._fd.NonlinearVariationalSolver(
problem,
solver_parameters=self._solver_parameters,
appctx=self._appctx,
)
self._rebuild_solver = False
def _prepare_step(
self,
step_size: float,
) -> tuple[tuple[np.ndarray, np.ndarray], ...]:
self._update_initial_trace(step_size)
if (
self._coefficient_step_size != step_size
or self._recurrence_coefficients is None
):
recurrence_coefficients = []
for spectrum, implicit_weight_field in zip(
self.spectra,
self._implicit_weights,
strict=True,
):
arguments = spectrum.rates * step_size
decay = np.asarray(exp_neg(arguments), dtype=np.float64)
interpolation = np.asarray(phi1(arguments), dtype=np.float64)
history_weights = spectrum.weights * decay
implicit_weight_field.assign(
float(np.dot(spectrum.weights, interpolation))
)
recurrence_coefficients.append(
(decay, interpolation, history_weights)
)
self._coefficient_step_size = step_size
self._recurrence_coefficients = tuple(recurrence_coefficients)
coefficients_by_term = self._recurrence_coefficients
assert coefficients_by_term is not None
for modes, history_term, coefficients in zip(
self._mode_groups,
self._history_terms,
coefficients_by_term,
strict=True,
):
_, _, history_weights = coefficients
with history_term.dat.vec as history_vector:
history_vector.set(0.0)
for weight, mode in zip(
history_weights,
modes,
strict=True,
):
with mode.dat.vec_ro as mode_vector:
history_vector.axpy(float(weight), mode_vector)
self._last_step_size = step_size
return tuple(
(decay, interpolation)
for decay, interpolation, _ in coefficients_by_term
)
def advance(self) -> None:
"""Solve for the next state and atomically commit all recurrence modes."""
current_parameters = self._read_parameters()
if current_parameters != self.parameters:
if "exponential_memory" not in self._operator_kinds:
raise RuntimeError(
"changing alpha after construction is unsupported; "
"rebuild the stepper"
)
raise RuntimeError(
"changing a time-memory parameter after construction is "
"unsupported; rebuild the stepper"
)
step_size = self._read_step_size()
recurrence_coefficients = self._prepare_step(step_size)
if self._rebuild_solver:
self._build_solver()
self._committed_u.assign(self.u)
try:
self._solver.solve()
except Exception:
self.u.assign(self._committed_u)
self._failure_count += 1
raise
with (
self.u.dat.vec_ro as solution_vector,
self._previous.dat.vec_ro as previous_vector,
self._increment.dat.vec as increment_vector,
):
increment_vector.waxpy(-1.0, previous_vector, solution_vector)
for modes, coefficients in zip(
self._mode_groups,
recurrence_coefficients,
strict=True,
):
decay, interpolation = coefficients
for mode, decay_value, interpolation_value in zip(
modes,
decay,
interpolation,
strict=True,
):
with mode.dat.vec as mode_vector:
mode_vector.scale(float(decay_value))
mode_vector.axpy(
float(interpolation_value),
increment_vector,
)
self._previous.assign(self.u)
self._solve_count += 1
self._nonlinear_iterations += self._solver.snes.getIterationNumber()
self._linear_iterations += self._solver.snes.getLinearSolveIterations()
@property
def history(self) -> tuple[Any, ...]:
"""Return defensive copies of the current diffusive mode fields."""
return tuple(
mode.copy(deepcopy=True)
for modes in self._mode_groups
for mode in modes
)
@property
def term_histories(self) -> tuple[tuple[Any, ...], ...]:
"""Return internal modes grouped by time-memory term."""
return tuple(
tuple(mode.copy(deepcopy=True) for mode in modes)
for modes in self._mode_groups
)
@property
def transformed_residual(self) -> Any:
"""The ordinary UFL residual supplied to Firedrake."""
return self._transformed_residual
def solver_stats(self) -> dict[str, Any]:
num_fractional_terms = sum(
kind != "exponential_memory" for kind in self._operator_kinds
)
return {
"solves": self._solve_count,
"failures": self._failure_count,
"nonlinear_iterations": self._nonlinear_iterations,
"linear_iterations": self._linear_iterations,
"num_modes": sum(len(modes) for modes in self._mode_groups),
"num_time_memory_terms": len(self._markers),
"num_fractional_terms": num_fractional_terms,
"num_exponential_memory_terms": (
len(self._markers) - num_fractional_terms
),
"modes_per_term": tuple(
len(modes) for modes in self._mode_groups
),
"last_step_size": self._last_step_size,
}
def reset(self, u0: Any, t0: Any = None) -> None:
"""Reset the physical state and zero every diffusive mode."""
self.u.assign(u0)
self._previous.assign(self.u)
self._initial.assign(self.u)
for modes in self._mode_groups:
for mode in modes:
mode.assign(0.0)
for history_term in self._history_terms:
history_term.assign(0.0)
for initial_trace_term in self._initial_trace_terms:
initial_trace_term.assign(0.0)
if t0 is not None:
self.t.assign(t0)
self._lower_limit = self._read_time()
self._last_step_size = None
self._reset_solver_counters()
def _checkpoint_metadata(self) -> dict[str, Any]:
payload = stepper_metadata(
kind="recurrence",
operator_kinds=self._operator_kinds,
parameters=self.parameters,
representations=[
term_representation.describe(parameter)
for term_representation, parameter in zip(
self._term_representations,
self.parameters,
strict=True,
)
],
formulation={
"kind": "recurrence",
"interpolant": self.formulation.interpolant,
},
)
payload.update({
"lower_limit": self._lower_limit,
"time": float(self.t),
"dt": float(self.dt),
"stats": self.solver_stats(),
})
return payload
def checkpoint_state(self) -> dict[str, Any]:
"""Return a serializable local-state checkpoint payload."""
payload = self._checkpoint_metadata()
payload.update(
{
"u": self.u.dat.data_ro.tolist(),
"initial": self._initial.dat.data_ro.tolist(),
"previous": self._previous.dat.data_ro.tolist(),
}
)
payload["mode_groups"] = [
[mode.dat.data_ro.tolist() for mode in modes]
for modes in self._mode_groups
]
return payload
def _checkpoint_file_fields(self) -> dict[str, Any]:
fields = {
"u": self.u,
"initial": self._initial,
"previous": self._previous,
}
fields.update(
{
f"mode_{term:04d}_{index:04d}": mode
for term, modes in enumerate(self._mode_groups)
for index, mode in enumerate(modes)
}
)
return fields
def save_checkpoint(self, checkpoint: Any, *, name: str = "state") -> None:
"""Save collective state to a Firedrake CheckpointFile."""
save_checkpoint_file(
checkpoint,
name=name,
metadata=self._checkpoint_metadata(),
fields=self._checkpoint_file_fields(),
)
def load_checkpoint(self, checkpoint: Any, *, name: str = "state") -> None:
"""Load collective state from a Firedrake CheckpointFile."""
fields = self._checkpoint_file_fields()
state, loaded = load_checkpoint_file(
checkpoint,
name=name,
mesh=self._space.mesh(),
expected_fields=tuple(fields),
)
state["u"] = loaded["u"].dat.data_ro.tolist()
state["initial"] = loaded["initial"].dat.data_ro.tolist()
state["previous"] = loaded["previous"].dat.data_ro.tolist()
mode_groups = [
[
loaded[f"mode_{term:04d}_{index:04d}"].dat.data_ro.tolist()
for index in range(len(modes))
]
for term, modes in enumerate(self._mode_groups)
]
state["mode_groups"] = mode_groups
self.restore_checkpoint(state)
def restore_checkpoint(self, state: dict[str, Any]) -> None:
"""Restore a payload produced by :meth:`checkpoint_state`."""
representations = validate_stepper_metadata(
state,
kind="recurrence",
operator_kinds=self._operator_kinds,
parameters=self.parameters,
)
formulation = dict(state.get("formulation") or {})
if (
formulation.get("kind") != "recurrence"
or formulation.get("interpolant") != self.formulation.interpolant
):
raise ValueError("checkpoint formulation does not match")
for metadata, term_representation, parameter in zip(
representations,
self._term_representations,
self.parameters,
strict=True,
):
validate_checkpoint_representation(
metadata,
term_representation,
parameter,
)
mode_groups = state.get("mode_groups")
if (
not isinstance(mode_groups, list)
or len(mode_groups) != len(self._mode_groups)
or any(not isinstance(modes, list) for modes in mode_groups)
or any(
len(values) != len(modes)
for values, modes in zip(
mode_groups,
self._mode_groups,
strict=True,
)
)
):
raise ValueError("checkpoint mode count does not match the stepper")
physical = np.asarray(state.get("u"), dtype=np.float64)
initial = np.asarray(state.get("initial"), dtype=np.float64)
previous = np.asarray(state.get("previous"), dtype=np.float64)
local_shape = self.u.dat.data_ro.shape
if (
physical.shape != local_shape
or initial.shape != local_shape
or previous.shape != local_shape
):
raise ValueError("checkpoint physical field has the wrong local shape")
mode_arrays = [
[
np.asarray(values, dtype=np.float64)
for values in group_values
]
for group_values in mode_groups
]
if any(
values.shape != local_shape
for group_values in mode_arrays
for values in group_values
):
raise ValueError("checkpoint mode field has the wrong local shape")
if not all(
np.all(np.isfinite(values))
for values in (physical, initial, previous)
) or any(
not np.all(np.isfinite(values))
for group_values in mode_arrays
for values in group_values
):
raise ValueError("checkpoint fields must contain finite values")
try:
checkpoint_time = float(state["time"])
checkpoint_dt = float(state["dt"])
checkpoint_lower_limit = float(state["lower_limit"])
except (KeyError, TypeError, ValueError) as error:
raise ValueError("checkpoint time metadata is invalid") from error
if not isfinite(checkpoint_time) or not isfinite(checkpoint_lower_limit):
raise ValueError("checkpoint times must be finite")
if not isfinite(checkpoint_dt) or checkpoint_dt <= 0.0:
raise ValueError("checkpoint dt must be finite and positive")
try:
stats = dict(state.get("stats") or {})
solve_count = int(stats.get("solves", 0))
failure_count = int(stats.get("failures", 0))
nonlinear_iterations = int(stats.get("nonlinear_iterations", 0))
linear_iterations = int(stats.get("linear_iterations", 0))
last_step_size = stats.get("last_step_size")
if last_step_size is not None:
last_step_size = float(last_step_size)
except (TypeError, ValueError, OverflowError) as error:
raise ValueError("checkpoint solver statistics are invalid") from error
if min(
solve_count,
failure_count,
nonlinear_iterations,
linear_iterations,
) < 0:
raise ValueError("checkpoint solver statistics must be nonnegative")
if last_step_size is not None and (
not isfinite(last_step_size) or last_step_size <= 0.0
):
raise ValueError(
"checkpoint last step size must be finite and positive"
)
self.u.dat.data[:] = physical
self._initial.dat.data[:] = initial
self._previous.dat.data[:] = previous
for modes, values_group in zip(
self._mode_groups,
mode_arrays,
strict=True,
):
for mode, values in zip(modes, values_group, strict=True):
mode.dat.data[:] = values
self.t.assign(checkpoint_time)
self.dt.assign(checkpoint_dt)
self._lower_limit = checkpoint_lower_limit
for initial_trace_term in self._initial_trace_terms:
initial_trace_term.assign(0.0)
self._solve_count = solve_count
self._failure_count = failure_count
self._nonlinear_iterations = nonlinear_iterations
self._linear_iterations = linear_iterations
self._last_step_size = last_step_size
self._rebuild_solver = True
[docs]
class ExponentialMemoryCompatibilityWarning(UserWarning):
"""Warn that bounded-kernel evolution requires compatible initial data."""
def _construct_time_stepper(
F: Any,
representation: Any | None,
t: Any,
dt: Any,
u: Any,
*,
formulation: Any = None,
u0: Any = None,
bcs: Any = None,
solver_parameters: Any = None,
appctx: Any = None,
allow_exponential: bool,
warn_initial_compatibility: bool,
) -> Any:
from yonderdrake.time._ufl_marker import (
ExponentialMemoryMarker,
find_time_memory_markers,
)
markers = find_time_memory_markers(F)
has_exponential = any(
isinstance(marker, ExponentialMemoryMarker) for marker in markers
)
if has_exponential and not allow_exponential:
raise ValueError("ExponentialMemory requires TimeMemoryStepper")
if has_exponential and warn_initial_compatibility:
warnings.warn(
"ExponentialMemory is zero at the initial time. If it is used as "
"the leading evolution operator, the remaining residual and "
"initial data must satisfy the corresponding compatibility "
"condition.",
ExponentialMemoryCompatibilityWarning,
stacklevel=3,
)
if isinstance(representation, FullHistory):
if has_exponential:
raise NotImplementedError(
"FullHistory cannot be combined with ExponentialMemory"
)
from yonderdrake.time.full_history import FullHistoryStepper
return FullHistoryStepper(
F,
representation,
t,
dt,
u,
formulation=formulation,
u0=u0,
bcs=bcs,
solver_parameters=solver_parameters,
appctx=appctx,
)
if formulation is None or isinstance(formulation, Recurrence):
return _RecurrenceStepper(
F,
representation,
t,
dt,
u,
formulation=formulation,
u0=u0,
bcs=bcs,
solver_parameters=solver_parameters,
appctx=appctx,
allow_exponential=allow_exponential,
)
if isinstance(formulation, AuxiliaryODE):
from yonderdrake.time.auxiliary_ode import AuxiliaryODEStepper
return AuxiliaryODEStepper(
F,
representation,
t,
dt,
u,
formulation=formulation,
u0=u0,
bcs=bcs,
solver_parameters=solver_parameters,
appctx=appctx,
allow_exponential=allow_exponential,
)
raise TypeError("formulation must be Recurrence or AuxiliaryODE")
[docs]
def TimeMemoryStepper(
F: Any,
t: Any,
dt: Any,
u: Any,
*,
representation: Any = None,
formulation: Any = None,
u0: Any = None,
bcs: Any = None,
solver_parameters: Any = None,
appctx: Any = None,
warn_initial_compatibility: bool = True,
) -> Any:
"""Advance exponential memory and optional fractional time markers."""
return _construct_time_stepper(
F,
representation,
t,
dt,
u,
formulation=formulation,
u0=u0,
bcs=bcs,
solver_parameters=solver_parameters,
appctx=appctx,
allow_exponential=True,
warn_initial_compatibility=warn_initial_compatibility,
)
[docs]
def FractionalTimeStepper(
F: Any,
representation: Any,
t: Any,
dt: Any,
u: Any,
*,
formulation: Any = None,
u0: Any = None,
bcs: Any = None,
solver_parameters: Any = None,
appctx: Any = None,
) -> Any:
"""Construct a native formulation for a fractional derivative marker."""
return _construct_time_stepper(
F,
representation,
t,
dt,
u,
formulation=formulation,
u0=u0,
bcs=bcs,
solver_parameters=solver_parameters,
appctx=appctx,
allow_exponential=False,
warn_initial_compatibility=False,
)