Source code for yonderdrake.time.stepper

"""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, )