Skip to content

Core Solver

core_solver

Core semi-implicit solver for time-evolving models (ODE + IMEX multiphysics).

This solver advances a :class:op_engine.model_core.ModelCore instance over its configured time grid. In the updated semantics, ModelCore.time_grid is treated as output times: the times at which the user wants a stored solution state.

Between consecutive output times, the solver may take either: - one or more fixed steps bounded by RunConfig.fixed_max_step (adaptive=False), or - multiple internal adaptive substeps that land exactly on t_{i+1} (adaptive=True).

Supported methods (keyword method=): - "euler": Explicit Euler (order 1), adaptive via step-doubling. - "heun": Explicit Heun / RK2 (order 2), embedded Euler estimator. - "rk4": Classic explicit Runge--Kutta (order 4). - "dopri5": Dormand--Prince 5(4), embedded adaptive estimator. - "imex-euler": IMEX Euler: explicit Euler on F(t,y), implicit Euler on A. Adaptive via step-doubling (IMEX step-doubling). - "imex-heun-tr": IMEX Heun-Trapezoidal: Heun on F, trapezoidal/CN on A. Adaptive via embedded low/high (Euler vs Heun) mapped by the same implicit operator solve. - "imex-trbdf2": IMEX TR-BDF2 (order 2), adaptive via step-doubling. - "imex-ark3": ARS(4,4,3) additive Runge--Kutta with embedded order 2. - "implicit-euler": One-linearization Euler approximation (order 1). - "trapezoidal": One-linearization trapezoidal approximation (order 2). - "bdf2": One-linearization BDF2 approximation (order 2). - "ros2": L-stable Rosenbrock-W 2(1). - "sdirk2": L-stable Alexander SDIRK2 with full nonlinear stages.

IMEX structure

We assume a split system: y' = A(t,y) y + F(t,y) where F is provided by rhs_func(t, y), and A is represented by linear operators applied along a single tensor axis. Operators may be: - None (ODE-only / explicit-only behavior), or - provided as tuples (predictor?, L, R), or - provided as factories depending on dt, stage-scale, and context.

Operator application

Operators act along a configured axis (default "state"). All other axes are batched. The solve form is: L @ y_next = R @ x optionally with a preprocessing predictor: x_tilde = predictor @ x

Non-uniform dt
  • Explicit methods naturally support non-uniform dt.
  • Implicit/IMEX methods require operator factories whenever dt varies across steps (non-uniform output grid or adaptive stepping), because L/R depend on dt.
Performance hygiene
  • NumPy paths retain preallocated scratch arrays and in-place operations.
  • Immutable namespaces use functional stepping operations.
  • Dense implicit solves use Array-API linalg; sparse adapters cache factors.

AdaptiveAdvanceParams(plan, t0, t1, y0, adaptive_cfg, dt_ctrl) dataclass

Bundle of parameters for adaptive advancement to an output time.

Attributes:

Name Type Description
plan RunPlan

Resolved run plan.

t0 float

Start time.

t1 float

End/output time.

y0 NDArray[floating]

Initial state at t0.

adaptive_cfg AdaptiveConfig

Adaptive stepping configuration.

dt_ctrl DtControllerConfig

dt controller configuration.

AdaptiveConfig(rtol=1e-06, atol=1e-09, dt_init=None, max_reject=25, max_steps=1000000) dataclass

Configuration for adaptive stepping.

Attributes:

Name Type Description
rtol float

Relative tolerance.

atol float | Array

Absolute tolerance (scalar or array-like).

dt_init float | None

Optional initial dt guess; if None, use output dt.

max_reject int

Maximum number of rejected attempts per accepted step.

max_steps int

Maximum number of internal substeps per output interval.

__post_init__()

Validate static adaptive-step parameters without coercing arrays.

Raises:

Type Description
ValueError

If any static adaptive parameter is invalid.

Source code in src/op_engine/core_solver.py
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
def __post_init__(self) -> None:
    """Validate static adaptive-step parameters without coercing arrays.

    Raises:
        ValueError: If any static adaptive parameter is invalid.
    """
    if not np.isfinite(self.rtol) or self.rtol < 0.0:
        msg = "rtol must be finite and non-negative"
        raise ValueError(msg)

    if isinstance(self.atol, (float, int, np.floating, np.integer)):
        if not np.isfinite(self.atol) or self.atol < 0.0:
            msg = "scalar atol must be finite and non-negative"
            raise ValueError(msg)
    elif isinstance(self.atol, np.ndarray) and (
        not np.all(np.isfinite(self.atol)) or np.any(self.atol < 0.0)
    ):
        msg = "NumPy atol values must be finite and non-negative"
        raise ValueError(msg)

    if self.dt_init is not None and (
        not np.isfinite(self.dt_init) or self.dt_init <= 0.0
    ):
        msg = "dt_init must be finite and positive when provided"
        raise ValueError(msg)
    if (
        not isinstance(self.max_reject, Integral)
        or isinstance(self.max_reject, bool)
        or self.max_reject < 1
    ):
        msg = "max_reject must be a positive integer"
        raise ValueError(msg)
    if (
        not isinstance(self.max_steps, Integral)
        or isinstance(self.max_steps, bool)
        or self.max_steps < 1
    ):
        msg = "max_steps must be a positive integer"
        raise ValueError(msg)

AdaptiveStepSchedule(output_times, step_sizes) dataclass

Accepted step sizes for replaying one adaptive solve.

The schedule records controller decisions, not array values. Replaying it therefore uses the same numerical kernels and active Array-API namespace as a live solve while keeping loop lengths and step sizes static.

Attributes:

Name Type Description
output_times tuple[float, ...]

Output grid used to create the schedule.

step_sizes tuple[tuple[float, ...], ...]

Accepted internal step sizes for each output interval.

__post_init__()

Normalize and validate the recorded mesh.

Raises:

Type Description
ValueError

If times or step sizes do not define a valid mesh.

Source code in src/op_engine/core_solver.py
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
def __post_init__(self) -> None:
    """Normalize and validate the recorded mesh.

    Raises:
        ValueError: If times or step sizes do not define a valid mesh.
    """
    output_times = tuple(float(value) for value in self.output_times)
    step_sizes = tuple(
        tuple(float(step_size) for step_size in interval)
        for interval in self.step_sizes
    )
    object.__setattr__(self, "output_times", output_times)
    object.__setattr__(self, "step_sizes", step_sizes)

    if not output_times:
        msg = "Adaptive schedule must contain at least one output time"
        raise ValueError(msg)
    if any(not math.isfinite(value) for value in output_times):
        msg = "Adaptive schedule output times must be finite"
        raise ValueError(msg)
    if any(end <= start for start, end in pairwise(output_times)):
        msg = "Adaptive schedule output times must be strictly increasing"
        raise ValueError(msg)
    if len(step_sizes) != len(output_times) - 1:
        msg = "Adaptive schedule must contain one step group per output interval"
        raise ValueError(msg)

    for interval_index, interval_steps in enumerate(step_sizes):
        if not interval_steps:
            msg = "Adaptive schedule step groups must not be empty"
            raise ValueError(msg)
        if any(
            not math.isfinite(step_size) or step_size <= 0.0
            for step_size in interval_steps
        ):
            msg = "Adaptive schedule step sizes must be finite and positive"
            raise ValueError(msg)

        interval = output_times[interval_index + 1] - output_times[interval_index]
        if not math.isclose(
            math.fsum(interval_steps),
            interval,
            rel_tol=1e-10,
            abs_tol=1e-12,
        ):
            msg = "Adaptive schedule steps must sum to each output interval"
            raise ValueError(msg)

CoreSolver(core, operators=None, *, operator_axis='state')

Semi-implicit solver operating on a ModelCore time/state grid.

Initialize CoreSolver.

Parameters:

Name Type Description Default
core ModelCore

ModelCore instance to solve.

required
operators CoreOperators | StageOperatorFactory | None

Default operator spec (tuple or factory) for implicit stages.

None
operator_axis str | int

Axis along which operators act (name or index).

'state'
Source code in src/op_engine/core_solver.py
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
def __init__(
    self,
    core: ModelCore,
    operators: CoreOperators | StageOperatorFactory | None = None,
    *,
    operator_axis: str | int = "state",
) -> None:
    """Initialize CoreSolver.

    Args:
        core: ModelCore instance to solve.
        operators: Default operator spec (tuple or factory) for implicit stages.
        operator_axis: Axis along which operators act (name or index).
    """
    self.core = core
    self.dtype = core.dtype
    self.state_shape = core.state_shape
    self.state_ndim = len(self.state_shape)

    # Operator axis resolution
    self._op_axis = operator_axis
    self._op_axis_idx: int | None = None
    self._op_axis_len: int | None = None

    # Default operator spec (tuple or factory or None)
    self._default_operator_spec: CoreOperators | StageOperatorFactory | None = (
        operators
    )

    # Preallocate buffers (full tensor shape)
    self._rhs_buffer: NDArray[np.floating] = np.zeros(
        self.state_shape,
        dtype=self.dtype,
    )
    self._next_state_buffer: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)

    # Shared stepping buffers
    self._f_n: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._f_pred: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._state_pred: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)

    # Adaptive buffers
    self._y_full: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._y_half: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._y_two_half: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._y_low: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._err: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)

    # Working state buffers (avoid allocating per substep)
    self._y_curr: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._y_try: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)

    # TR-BDF2 additional buffers
    self._y_stage1: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._f_stage1: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._f_extrap: NDArray[np.floating] = np.zeros_like(self._rhs_buffer)
    self._sparse_identity: csr_matrix | None = None

    self._last_adaptive_schedule: AdaptiveStepSchedule | None = None
    self._last_nonlinear_diagnostics: NonlinearIntegrationDiagnostics | None = None

    # Validate operator sizes if default spec is a static tuple
    if operators is not None and not callable(operators):
        _predictor, left_op, right_op = self._normalize_ops_tuple(operators)
        if left_op is not None and right_op is not None:
            self._resolve_operator_axis()
            self._validate_operator_sizes(left_op, right_op)

last_adaptive_schedule property

Return the schedule recorded or replayed by the latest adaptive run.

last_nonlinear_diagnostics property

Return nonlinear diagnostics from the latest SDIRK run or replay.

adaptive_explicit_step(rhs_func, *, method, t, dt, y, first_stage=None)

Attempt one explicit step and return its local-error estimate.

This functional boundary does not mutate :class:ModelCore history and does not extract host scalars. Provider integrations can compose it with backend-native adaptive controller loops.

Parameters:

Name Type Description Default
rhs_func RHSFunction

Function computing the explicit RHS F(t, y).

required
method str

Explicit solver method name.

required
t Scalar

Step start time as a Python float or backend-native scalar.

required
dt Scalar

Step size as a Python float or backend-native scalar.

required
y Array

State at the step start.

required
first_stage Array | None

Optional cached derivative at (t, y).

None

Returns:

Type Description
ExplicitStepResult

Candidate state, local error estimate, controller order, and

ExplicitStepResult

reusable derivative stages.

Raises:

Type Description
ValueError

If method is not an explicit method.

Source code in src/op_engine/core_solver.py
2995
2996
2997
2998
2999
3000
3001
3002
3003
3004
3005
3006
3007
3008
3009
3010
3011
3012
3013
3014
3015
3016
3017
3018
3019
3020
3021
3022
3023
3024
3025
3026
3027
3028
3029
3030
3031
3032
3033
3034
3035
3036
3037
def adaptive_explicit_step(  # noqa: PLR0913
    self,
    rhs_func: RHSFunction,
    *,
    method: str,
    t: Scalar,
    dt: Scalar,
    y: Array,
    first_stage: Array | None = None,
) -> ExplicitStepResult:
    """Attempt one explicit step and return its local-error estimate.

    This functional boundary does not mutate :class:`ModelCore` history
    and does not extract host scalars. Provider integrations can compose
    it with backend-native adaptive controller loops.

    Args:
        rhs_func: Function computing the explicit RHS F(t, y).
        method: Explicit solver method name.
        t: Step start time as a Python float or backend-native scalar.
        dt: Step size as a Python float or backend-native scalar.
        y: State at the step start.
        first_stage: Optional cached derivative at ``(t, y)``.

    Returns:
        Candidate state, local error estimate, controller order, and
        reusable derivative stages.

    Raises:
        ValueError: If ``method`` is not an explicit method.
    """
    normalized_method = _normalize_method(method)
    if normalized_method not in _EXPLICIT_METHODS:
        msg = f"Method '{method}' is not an explicit solver method"
        raise ValueError(msg)
    return self._attempt_explicit_step(
        rhs_func,
        method=normalized_method,
        t=cast("float", t),
        dt=cast("float", dt),
        y=y,
        first_stage=first_stage,
    )

fixed_explicit_step(rhs_func, *, method, t, dt, y, first_stage=None)

Take one functional fixed step with an explicit method.

This boundary does not mutate :class:ModelCore history, so callers can compose it with backend-native loop primitives and apply the completed trajectory once. It is the supported public step kernel for external drivers such as jax.lax.scan; the method and state shape are static, while t, dt, y, and RHS parameters may be traced. RHS results must preserve the configured state shape and the namespace of y.

Parameters:

Name Type Description Default
rhs_func RHSFunction

Function computing the explicit RHS F(t, y).

required
method str

Explicit solver method name.

required
t Scalar

Step start time as a Python float or backend-native scalar.

required
dt Scalar

Step size as a Python float or backend-native scalar.

required
y Array

State at the step start.

required
first_stage Array | None

Optional derivative at (t, y). Dormand--Prince returns this cache for the next step; reuse it only when time, state, and RHS parameters are unchanged at that start.

None

Returns:

Type Description
Array

Next state and an FSAL derivative for Dormand--Prince, or

Array | None

None for other explicit methods. Seed a scan's FSAL carry

tuple[Array, Array | None]

with one step outside the scan so its structure stays fixed.

Raises:

Type Description
ValueError

If method is not an explicit method.

Source code in src/op_engine/core_solver.py
3039
3040
3041
3042
3043
3044
3045
3046
3047
3048
3049
3050
3051
3052
3053
3054
3055
3056
3057
3058
3059
3060
3061
3062
3063
3064
3065
3066
3067
3068
3069
3070
3071
3072
3073
3074
3075
3076
3077
3078
3079
3080
3081
3082
3083
3084
3085
3086
3087
3088
def fixed_explicit_step(  # noqa: PLR0913
    self,
    rhs_func: RHSFunction,
    *,
    method: str,
    t: Scalar,
    dt: Scalar,
    y: Array,
    first_stage: Array | None = None,
) -> tuple[Array, Array | None]:
    """Take one functional fixed step with an explicit method.

    This boundary does not mutate :class:`ModelCore` history, so callers
    can compose it with backend-native loop primitives and apply the
    completed trajectory once. It is the supported public step kernel
    for external drivers such as ``jax.lax.scan``; the method and state
    shape are static, while ``t``, ``dt``, ``y``, and RHS parameters may
    be traced. RHS results must preserve the configured state shape and
    the namespace of ``y``.

    Args:
        rhs_func: Function computing the explicit RHS F(t, y).
        method: Explicit solver method name.
        t: Step start time as a Python float or backend-native scalar.
        dt: Step size as a Python float or backend-native scalar.
        y: State at the step start.
        first_stage: Optional derivative at ``(t, y)``. Dormand--Prince
            returns this cache for the next step; reuse it only when
            time, state, and RHS parameters are unchanged at that start.

    Returns:
        Next state and an FSAL derivative for Dormand--Prince, or
        ``None`` for other explicit methods. Seed a scan's FSAL carry
        with one step outside the scan so its structure stays fixed.

    Raises:
        ValueError: If ``method`` is not an explicit method.
    """
    normalized_method = _normalize_method(method)
    if normalized_method not in _EXPLICIT_METHODS:
        msg = f"Method '{method}' is not an explicit solver method"
        raise ValueError(msg)
    return self._step_explicit_fixed(
        rhs_func,
        method=normalized_method,
        t=t,
        dt=dt,
        y=y,
        first_stage=first_stage,
    )

replay_adaptive_schedule(rhs_func, schedule, *, config)

Replay a recorded adaptive mesh through the configured method.

Replay bypasses error norms and accept/reject decisions while invoking the same high-order step kernels used by a live adaptive solve. With a JAX state, the static Python schedule can therefore be traced by jax.jit and differentiated with respect to array-valued model inputs.

Parameters:

Name Type Description Default
rhs_func RHSFunction

Function computing the explicit RHS F(t, y).

required
schedule AdaptiveStepSchedule

Accepted step mesh recorded by an adaptive run.

required
config RunConfig

Matching adaptive run configuration.

required

Returns:

Type Description
NonlinearIntegrationDiagnostics | None

Array-valued nonlinear diagnostics for SDIRK2, otherwise None.

Raises:

Type Description
TypeError

If schedule has the wrong type or the state changes array ecosystems during replay.

ValueError

If config is not adaptive or the output grid differs.

Source code in src/op_engine/core_solver.py
4615
4616
4617
4618
4619
4620
4621
4622
4623
4624
4625
4626
4627
4628
4629
4630
4631
4632
4633
4634
4635
4636
4637
4638
4639
4640
4641
4642
4643
4644
4645
4646
4647
4648
4649
4650
4651
4652
4653
4654
4655
4656
4657
4658
4659
4660
4661
4662
4663
4664
4665
4666
4667
4668
4669
4670
4671
4672
4673
4674
4675
4676
4677
4678
4679
4680
4681
4682
4683
4684
4685
4686
4687
4688
4689
4690
4691
4692
4693
4694
4695
4696
4697
4698
4699
4700
4701
def replay_adaptive_schedule(
    self,
    rhs_func: RHSFunction,
    schedule: AdaptiveStepSchedule,
    *,
    config: RunConfig,
) -> NonlinearIntegrationDiagnostics | None:
    """Replay a recorded adaptive mesh through the configured method.

    Replay bypasses error norms and accept/reject decisions while invoking
    the same high-order step kernels used by a live adaptive solve. With a
    JAX state, the static Python schedule can therefore be traced by
    ``jax.jit`` and differentiated with respect to array-valued model inputs.

    Args:
        rhs_func: Function computing the explicit RHS F(t, y).
        schedule: Accepted step mesh recorded by an adaptive run.
        config: Matching adaptive run configuration.

    Returns:
        Array-valued nonlinear diagnostics for SDIRK2, otherwise ``None``.

    Raises:
        TypeError: If schedule has the wrong type or the state changes array
            ecosystems during replay.
        ValueError: If config is not adaptive or the output grid differs.
    """
    if not isinstance(schedule, AdaptiveStepSchedule):
        msg = "schedule must be an AdaptiveStepSchedule"
        raise TypeError(msg)
    if not config.adaptive:
        msg = "Schedule replay requires config.adaptive=True"
        raise ValueError(msg)

    self._last_adaptive_schedule = None
    self._last_nonlinear_diagnostics = None
    self._validate_schedule_time_grid(schedule)
    plan = self._resolve_run_plan(config)

    adapter = None
    if config.replay_loop != "unroll":
        adapter = get_loop_adapter(_namespace_of(self.core.get_current_state()))
        if plan.method == "sdirk2":
            adapter = None
        if config.replay_loop == "scan" and adapter is None:
            msg = (
                "Scan replay requires a loop adapter and a method other than SDIRK2"
            )
            raise ValueError(msg)
    if adapter is not None:
        self._replay_scan_schedule(
            rhs_func,
            plan=plan,
            schedule=schedule,
            adapter=adapter,
            checkpoint=config.replay_checkpoint,
        )
        self._last_adaptive_schedule = schedule
        return None

    if plan.method == "sdirk2":
        diagnostics = self._replay_sdirk_schedule(
            rhs_func,
            plan=plan,
            schedule=schedule,
        )
        self._last_adaptive_schedule = schedule
        return diagnostics
    if plan.method in _EXPLICIT_METHODS:
        self._replay_explicit_schedule(rhs_func, plan=plan, schedule=schedule)
    else:
        current_state = self.core.get_current_state()
        if plan.method == "imex-ark3" or not isinstance(current_state, np.ndarray):
            self._replay_array_implicit_schedule(
                rhs_func,
                plan=plan,
                schedule=schedule,
            )
        else:
            self._replay_numpy_implicit_schedule(
                rhs_func,
                plan=plan,
                schedule=schedule,
            )

    self._last_adaptive_schedule = schedule
    return None

run(rhs_func, *, config=None)

Advance the ModelCore state through its time grid.

Parameters:

Name Type Description Default
rhs_func RHSFunction

Function computing the explicit RHS F(t, y).

required
config RunConfig | None

Optional run configuration. If None, defaults are used.

None

Returns:

Type Description
NonlinearIntegrationDiagnostics | None

Array-valued nonlinear diagnostics for SDIRK2, otherwise None.

Raises:

Type Description
TypeError

If the state changes array ecosystems during a run.

ValueError

If invalid parameters are provided.

Source code in src/op_engine/core_solver.py
4703
4704
4705
4706
4707
4708
4709
4710
4711
4712
4713
4714
4715
4716
4717
4718
4719
4720
4721
4722
4723
4724
4725
4726
4727
4728
4729
4730
4731
4732
4733
4734
4735
4736
4737
4738
4739
4740
4741
4742
4743
4744
4745
4746
4747
4748
4749
4750
4751
4752
4753
4754
4755
4756
4757
4758
4759
4760
4761
4762
4763
4764
4765
4766
4767
4768
4769
4770
4771
4772
4773
4774
4775
4776
4777
4778
4779
4780
4781
4782
4783
4784
4785
4786
4787
4788
4789
4790
4791
4792
4793
4794
4795
def run(
    self,
    rhs_func: RHSFunction,
    *,
    config: RunConfig | None = None,
) -> NonlinearIntegrationDiagnostics | None:
    """Advance the ModelCore state through its time grid.

    Args:
        rhs_func: Function computing the explicit RHS F(t, y).
        config: Optional run configuration. If None, defaults are used.

    Returns:
        Array-valued nonlinear diagnostics for SDIRK2, otherwise ``None``.

    Raises:
        TypeError: If the state changes array ecosystems during a run.
        ValueError: If invalid parameters are provided.
    """
    self._last_adaptive_schedule = None
    self._last_nonlinear_diagnostics = None
    cfg = config or RunConfig()
    plan = self._resolve_run_plan(cfg)

    if plan.method == "sdirk2":
        return self._run_sdirk(rhs_func, plan=plan, config=cfg)

    if plan.method in _EXPLICIT_METHODS:
        self._run_explicit(rhs_func, plan=plan, config=cfg)
        return None

    current_state = self.core.get_current_state()
    if plan.method == "imex-ark3" or not isinstance(current_state, np.ndarray):
        self._run_array_implicit(rhs_func, plan=plan, config=cfg)
        return None

    time_grid = np.asarray(self.core.time_grid, dtype=float)
    n_steps = int(self.core.n_timesteps)
    history = MultistepHistory[NDArray[np.floating]](
        capacity=BDF2.stored_state_history if plan.method == "bdf2" else 0,
    )
    schedule_steps: list[tuple[float, ...]] = []

    for idx in range(n_steps - 1):
        t0 = float(time_grid[idx])
        t1 = float(time_grid[idx + 1])
        if t1 <= t0:
            raise ValueError(_TIME_GRID_INCREASING_ERROR_MSG)
        dt_out = t1 - t0

        state = self.core.get_current_state()
        if not isinstance(state, np.ndarray):
            raise TypeError(
                _NUMPY_ARRAY_ERROR_MSG.format(type_name=type(state).__name__)
            )
        np.copyto(self._y_curr, state)

        if not cfg.adaptive:
            y_next = self._advance_nonadaptive_to_time(
                rhs_func,
                plan=plan,
                t0=t0,
                dt_out=dt_out,
                y0=self._y_curr,
                history=history.states,
            )
            history = history.push(self._y_curr)
            self.core.advance_timestep(y_next)
            continue

        accepted_steps: list[float] = []
        y_end = self._advance_adaptive_to_time(
            rhs_func,
            AdaptiveAdvanceParams(
                plan=plan,
                t0=t0,
                t1=t1,
                y0=self._y_curr,
                adaptive_cfg=cfg.adaptive_cfg,
                dt_ctrl=cfg.dt_controller,
            ),
            accepted_steps,
        )
        schedule_steps.append(tuple(accepted_steps))
        history = history.push(self._y_curr)
        self.core.advance_timestep(y_end)

    if cfg.adaptive:
        self._last_adaptive_schedule = AdaptiveStepSchedule(
            output_times=tuple(float(value) for value in time_grid),
            step_sizes=tuple(schedule_steps),
        )
    return None

DtControllerConfig(dt_min=0.0, dt_max=float('inf'), safety=0.9, fac_min=0.2, fac_max=5.0) dataclass

Configuration for adaptive timestep control.

Attributes:

Name Type Description
dt_min float

Minimum allowed dt.

dt_max float

Maximum allowed dt.

safety float

Safety factor applied to dt updates.

fac_min float

Minimum multiplicative change factor.

fac_max float

Maximum multiplicative change factor.

__post_init__()

Validate static timestep-controller parameters.

Raises:

Type Description
ValueError

If any controller parameter is outside its valid range.

Source code in src/op_engine/core_solver.py
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
def __post_init__(self) -> None:
    """Validate static timestep-controller parameters.

    Raises:
        ValueError: If any controller parameter is outside its valid range.
    """
    if not np.isfinite(self.dt_min) or self.dt_min < 0.0:
        msg = "dt_min must be finite and non-negative"
        raise ValueError(msg)
    if np.isnan(self.dt_max) or self.dt_max <= 0.0:
        msg = "dt_max must be positive and not NaN"
        raise ValueError(msg)
    if self.dt_max < self.dt_min:
        msg = "dt_max must be greater than or equal to dt_min"
        raise ValueError(msg)
    if not np.isfinite(self.safety) or self.safety <= 0.0:
        msg = "safety must be finite and positive"
        raise ValueError(msg)
    if not np.isfinite(self.fac_min) or self.fac_min <= 0.0:
        msg = "fac_min must be finite and positive"
        raise ValueError(msg)
    if not np.isfinite(self.fac_max) or self.fac_max <= 0.0:
        msg = "fac_max must be finite and positive"
        raise ValueError(msg)
    if self.fac_max < self.fac_min:
        msg = "fac_max must be greater than or equal to fac_min"
        raise ValueError(msg)

ExplicitStepResult(state, error, controller_order, first_stage, last_stage) dataclass

Result and reusable stages from one explicit adaptive attempt.

Attributes:

Name Type Description
state Array

Accepted-order candidate state.

error Array

Local error estimate.

controller_order int

Order supplied to the existing step-size controller.

first_stage Array

Derivative at the attempted step's initial state.

last_stage Array | None

Derivative reusable by an FSAL method after acceptance.

ImexEulerOnceParams(t, y, dt, op_spec, out) dataclass

Bundle of parameters for one IMEX Euler step (non-doubling).

Attributes:

Name Type Description
t float

Current time.

y NDArray[floating]

Current state.

dt float

Step size.

op_spec CoreOperators | StageOperatorFactory | None

Operator spec for implicit stage.

out NDArray[floating]

Output state array (written in-place).

ImplicitStageParams(spec, dt, scale, t_stage, y_stage, stage, x, out) dataclass

Bundle of parameters for one implicit operator application.

Attributes:

Name Type Description
spec CoreOperators | StageOperatorFactory | None

Operator spec (tuple or factory) or None for identity.

dt float

Full-step dt for operator factory context.

scale float

Stage scaling factor for dt-dependent operators.

t_stage float

Stage time.

y_stage NDArray[floating]

Stage state proxy for operator factories.

stage str

Stage label (e.g., "be", "tr", "bdf2").

x NDArray[floating]

Input array to map.

out NDArray[floating]

Output array (written in-place).

NonlinearIntegrationConvergenceError(diagnostics)

Bases: RuntimeError

Raised after validating failed array-valued integration diagnostics.

Store diagnostics for the invalid integration or frozen mesh.

Source code in src/op_engine/core_solver.py
545
546
547
548
549
550
551
552
553
554
555
def __init__(self, diagnostics: NonlinearIntegrationDiagnostics) -> None:
    """Store diagnostics for the invalid integration or frozen mesh."""
    self.diagnostics = diagnostics
    failed = np.logical_and(
        np.asarray(diagnostics.step_accepted),
        np.logical_not(np.asarray(diagnostics.step_converged)),
    )
    super().__init__(
        f"{int(np.count_nonzero(failed))} accepted nonlinear step(s) "
        "did not converge"
    )

NonlinearIntegrationDiagnostics

Bases: NamedTuple

Array-valued nonlinear diagnostics for one integration or replay.

Every field is in the state array's namespace, so the record can cross a compiled JAX boundary. stages_per_step maps flattened stage fields back to attempted steps. Rejected adaptive attempts remain present with step_accepted=False.

require_converged()

Return valid diagnostics or invalidate a failed compiled replay.

Returns:

Type Description
NonlinearIntegrationDiagnostics

This unchanged diagnostic record.

Raises:

Type Description
NonlinearIntegrationConvergenceError

If an accepted or replayed step contains a failed nonlinear stage.

Source code in src/op_engine/core_solver.py
527
528
529
530
531
532
533
534
535
536
537
538
539
def require_converged(self) -> NonlinearIntegrationDiagnostics:
    """Return valid diagnostics or invalidate a failed compiled replay.

    Returns:
        This unchanged diagnostic record.

    Raises:
        NonlinearIntegrationConvergenceError: If an accepted or replayed
            step contains a failed nonlinear stage.
    """
    if not bool(self.converged.item()):
        raise NonlinearIntegrationConvergenceError(self)
    return self

NonlinearMethodConfig(rhs_jacobian, solver=DenseNewtonSolver()) dataclass

Configuration for fully nonlinear integration methods.

The Jacobian acts on the entire flattened RHS state. It is deliberately separate from :attr:RunConfig.jacobian, whose operators act only along CoreSolver.operator_axis for linearly implicit methods.

Attributes:

Name Type Description
rhs_jacobian FullRhsJacobianFunction

Dense full-system Jacobian with shape (state.size, state.size).

solver NonlinearSolver

Backend-neutral nonlinear solver implementation.

__post_init__()

Validate the method boundary without selecting an array backend.

Raises:

Type Description
TypeError

If a callback or nonlinear solver is invalid.

Source code in src/op_engine/core_solver.py
707
708
709
710
711
712
713
714
715
716
717
718
def __post_init__(self) -> None:
    """Validate the method boundary without selecting an array backend.

    Raises:
        TypeError: If a callback or nonlinear solver is invalid.
    """
    if not callable(self.rhs_jacobian):
        msg = "rhs_jacobian must be callable"
        raise TypeError(msg)
    if not isinstance(self.solver, NonlinearSolver):
        msg = "solver must implement the NonlinearSolver protocol"
        raise TypeError(msg)

OperatorLike

Bases: Protocol

Minimal operator interface required by CoreSolver.

Implementations are expected to behave like 2D linear operators suitable for implicit_solve(L, R, rhs2d). Only shape is required for validation.

shape property

Operator shape.

OperatorSpecs(default=None, tr=None, bdf2=None) dataclass

Operator specifications for implicit/IMEX methods.

Attributes:

Name Type Description
default CoreOperators | StageOperatorFactory | None

Default operator spec (tuple or factory) used by IMEX Euler/Heun-TR and as a fallback for TR/BDF2 stages.

tr CoreOperators | StageOperatorFactory | None

Operator spec for trapezoidal stage of TR-BDF2 (optional).

bdf2 CoreOperators | StageOperatorFactory | None

Operator spec for BDF2 stage of TR-BDF2 (optional).

PredictorLike

Bases: Protocol

Minimal predictor interface required by CoreSolver.

The predictor is an optional preprocessing operator applied as

rhs2d = predictor @ rhs2d

__matmul__(other)

Apply the predictor to a 2D array.

Source code in src/op_engine/core_solver.py
256
257
258
def __matmul__(self, other: Array) -> Array:
    """Apply the predictor to a 2D array."""
    ...

RunConfig(method='heun', adaptive=False, strict=True, dt_controller=DtControllerConfig(), adaptive_cfg=AdaptiveConfig(), operators=OperatorSpecs(), jacobian=None, nonlinear=None, gamma=None, fixed_max_step=None, replay_loop='unroll', replay_checkpoint=False) dataclass

Configuration for CoreSolver.run.

Attributes:

Name Type Description
method str

Method name.

adaptive bool

Whether to use adaptive substepping between output times.

strict bool

If True, invalid configurations raise; otherwise warnings and method downshifts may occur.

dt_controller DtControllerConfig

Parameters for dt controller when adaptive=True.

adaptive_cfg AdaptiveConfig

Parameters controlling error tolerances and limits.

operators OperatorSpecs

Operator specifications for implicit/IMEX methods.

jacobian JacobianFunction | None

Optional Jacobian function for linearly implicit methods.

nonlinear NonlinearMethodConfig | None

Full-system Jacobian and backend-neutral nonlinear solver.

gamma float | None

Optional TR-BDF2 gamma (if None, uses default).

fixed_max_step float | None

Maximum explicit fixed-step size between output times. None retains one step per output interval.

replay_loop Literal['auto', 'scan', 'unroll']

Iteration strategy for frozen adaptive replay. unroll retains eager loops; auto uses an available namespace adapter; scan requires one. SDIRK2 diagnostics retain eager replay.

replay_checkpoint bool

Rematerialize the scan body during reverse mode. Used only when replay selects an adapter.

__post_init__()

Normalize the method and validate context-free configuration.

Raises:

Type Description
TypeError

If a nested configuration has the wrong type.

ValueError

If the method or TR-BDF2 gamma is invalid.

Source code in src/op_engine/core_solver.py
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
def __post_init__(self) -> None:
    """Normalize the method and validate context-free configuration.

    Raises:
        TypeError: If a nested configuration has the wrong type.
        ValueError: If the method or TR-BDF2 gamma is invalid.
    """
    method = _normalize_method(self.method)
    object.__setattr__(self, "method", method)

    self._validate_replay_options()

    if self.fixed_max_step is not None:
        fixed_max_step = float(self.fixed_max_step)
        if not np.isfinite(fixed_max_step) or fixed_max_step <= 0.0:
            raise ValueError(_FIXED_MAX_STEP_ERROR_MSG)
        object.__setattr__(self, "fixed_max_step", fixed_max_step)
        if self.adaptive or method not in _EXPLICIT_METHODS:
            raise ValueError(_FIXED_STEP_MODE_ERROR_MSG)

    if not isinstance(self.dt_controller, DtControllerConfig):
        msg = "dt_controller must be a DtControllerConfig"
        raise TypeError(msg)
    if not isinstance(self.adaptive_cfg, AdaptiveConfig):
        msg = "adaptive_cfg must be an AdaptiveConfig"
        raise TypeError(msg)
    if not isinstance(self.operators, OperatorSpecs):
        msg = "operators must be an OperatorSpecs"
        raise TypeError(msg)
    if self.nonlinear is not None and not isinstance(
        self.nonlinear, NonlinearMethodConfig
    ):
        msg = "nonlinear must be a NonlinearMethodConfig"
        raise TypeError(msg)

    if method == "imex-trbdf2" and self.gamma is not None:
        gamma = float(self.gamma)
        if not np.isfinite(gamma) or not (0.0 < gamma < 1.0):
            raise ValueError(_GAMMA_RANGE_ERROR_MSG)

RunPlan(method, gamma, op_default, op_tr, op_bdf2, jacobian, nonlinear=None) dataclass

Resolved execution plan derived from RunConfig.

This is the internal, validated form used by the stepping loops.

Attributes:

Name Type Description
method MethodName

Final method after any strict=False downshifts.

gamma float | None

TR-BDF2 gamma, or None for non-TR-BDF2 methods.

op_default CoreOperators | StageOperatorFactory | None

Operator spec for IMEX Euler/Heun-TR.

op_tr CoreOperators | StageOperatorFactory | None

TR-stage operator spec for TR-BDF2.

op_bdf2 CoreOperators | StageOperatorFactory | None

BDF2-stage operator spec for TR-BDF2.

jacobian JacobianFunction | None

Optional Jacobian function for linearly implicit methods.

nonlinear NonlinearMethodConfig | None

Configuration for fully nonlinear methods.

StepIO(t, dt, y, out, err_out=None, history=()) dataclass

Bundle of per-step state for stepping kernels.

Attributes:

Name Type Description
t float

Current time.

dt float

Step size.

y NDArray[floating]

Current state array (input).

out NDArray[floating]

Output state array (written in-place).

err_out NDArray[floating] | None

Error estimate array (written in-place) for adaptive methods.

history tuple[NDArray[floating], ...]

Older accepted states for multistep methods, newest first.

Trbdf2OnceParams(t, y, dt, operators_tr, operators_bdf2, gamma, out) dataclass

Bundle of parameters for one TR-BDF2 step (non-doubling).

Attributes:

Name Type Description
t float

Current time.

y NDArray[floating]

Current state.

dt float

Step size.

operators_tr CoreOperators | StageOperatorFactory | None

TR stage operator spec.

operators_bdf2 CoreOperators | StageOperatorFactory | None

BDF2 stage operator spec.

gamma float

TR-BDF2 gamma.

out NDArray[floating]

Output state array (written in-place).

fixed_step_sizes(t0, t1, max_step)

Partition one output interval into deterministic fixed steps.

The first steps use max_step and the final step absorbs the remainder, so the partition lands on t1 without changing the stored output grid.

Parameters:

Name Type Description Default
t0 float

Output interval start.

required
t1 float

Output interval end.

required
max_step float | None

Maximum internal step, or None for one full interval step.

required

Returns:

Type Description
tuple[float, ...]

Positive internal step sizes that sum to t1 - t0.

Raises:

Type Description
ValueError

If the interval or maximum step is invalid, or if the partition would be unreasonably large.

Source code in src/op_engine/core_solver.py
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
def fixed_step_sizes(
    t0: float,
    t1: float,
    max_step: float | None,
) -> tuple[float, ...]:
    """Partition one output interval into deterministic fixed steps.

    The first steps use ``max_step`` and the final step absorbs the remainder,
    so the partition lands on ``t1`` without changing the stored output grid.

    Args:
        t0: Output interval start.
        t1: Output interval end.
        max_step: Maximum internal step, or ``None`` for one full interval step.

    Returns:
        Positive internal step sizes that sum to ``t1 - t0``.

    Raises:
        ValueError: If the interval or maximum step is invalid, or if the
            partition would be unreasonably large.
    """
    start = float(t0)
    end = float(t1)
    if not math.isfinite(start) or not math.isfinite(end) or end <= start:
        raise ValueError(_TIME_GRID_INCREASING_ERROR_MSG)

    interval = end - start
    if max_step is None:
        return (interval,)
    step_limit = float(max_step)
    if not math.isfinite(step_limit) or step_limit <= 0.0:
        raise ValueError(_FIXED_MAX_STEP_ERROR_MSG)
    if step_limit >= interval:
        return (interval,)

    ratio = interval / step_limit
    if not math.isfinite(ratio) or ratio > _MAX_FIXED_STEPS_PER_INTERVAL:
        msg = (
            "fixed_max_step requires more than "
            f"{_MAX_FIXED_STEPS_PER_INTERVAL} steps in one output interval"
        )
        raise ValueError(msg)

    tolerance = 16.0 * np.finfo(float).eps * max(1.0, abs(ratio))
    step_count = max(1, math.ceil(ratio - tolerance))
    if step_count == 1:
        return (interval,)

    final_step = interval - step_limit * (step_count - 1)
    if final_step <= 0.0:
        # A ratio within rounding tolerance of an integer may make the direct
        # subtraction non-positive. Keep the same count and absorb roundoff in
        # the final step instead of adding a spurious tiny step.
        final_step = step_limit + final_step
        step_count -= 1
    return (step_limit,) * (step_count - 1) + (final_step,)

propose_step_size(step_size, error_norm, order, *, config)

Return the next adaptive step size without host scalar extraction.

Returns:

Type Description
Array

Scalar array in error_norm's namespace, clamped to the configured

Array

controller bounds.

Source code in src/op_engine/core_solver.py
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
def propose_step_size(
    step_size: Scalar,
    error_norm: Array,
    order: int,
    *,
    config: DtControllerConfig,
) -> Array:
    """Return the next adaptive step size without host scalar extraction.

    Returns:
        Scalar array in ``error_norm``'s namespace, clamped to the configured
        controller bounds.
    """
    xp = _namespace_of(error_norm)
    error_value = xp.asarray(error_norm)
    positive = xp.greater(error_value, 0.0)
    safe_error = xp.where(positive, error_value, xp.ones_like(error_value))
    exponent = -1.0 / float(order + 1)
    scaled = xp.multiply(config.safety, xp.pow(safe_error, exponent))
    factor = xp.where(positive, scaled, config.fac_max)
    factor = xp.maximum(factor, config.fac_min)
    factor = xp.minimum(factor, config.fac_max)
    proposed = xp.multiply(xp.asarray(step_size, dtype=error_value.dtype), factor)
    proposed = xp.maximum(proposed, config.dt_min)
    return cast("Array", xp.minimum(proposed, config.dt_max))

scaled_error_norm(error, reference, previous, *, rtol, atol, reduction='rms')

Return a backend-native scaled local-error norm.

Unlike the host controller wrapper used by :class:CoreSolver, this function never extracts a Python scalar. Provider integrations can therefore compose it with backend loop primitives such as jax.lax.while_loop while retaining the core solver's error semantics.

Returns:

Type Description
Array

Scalar array in error's namespace. Non-finite norms are mapped to

Array

positive infinity so compiled controllers reject the attempted step.

Raises:

Type Description
ValueError

If reduction is not "rms" or "max".

Source code in src/op_engine/core_solver.py
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
def scaled_error_norm(  # noqa: PLR0913
    error: Array,
    reference: Array,
    previous: Array,
    *,
    rtol: float,
    atol: float | Array,
    reduction: Literal["rms", "max"] = "rms",
) -> Array:
    """Return a backend-native scaled local-error norm.

    Unlike the host controller wrapper used by :class:`CoreSolver`, this
    function never extracts a Python scalar. Provider integrations can
    therefore compose it with backend loop primitives such as
    ``jax.lax.while_loop`` while retaining the core solver's error semantics.

    Returns:
        Scalar array in ``error``'s namespace. Non-finite norms are mapped to
        positive infinity so compiled controllers reject the attempted step.

    Raises:
        ValueError: If ``reduction`` is not ``"rms"`` or ``"max"``.
    """
    xp = _namespace_of(error)
    scale = xp.maximum(xp.abs(reference), xp.abs(previous))
    scale = xp.multiply(scale, rtol)
    if isinstance(atol, (float, int, np.floating)):
        scale = xp.add(scale, float(atol))
    else:
        atol_array = xp.asarray(atol, dtype=error.dtype)
        scale = xp.add(scale, atol_array)

    ratio = xp.divide(error, scale)
    if reduction == "rms":
        squared = xp.multiply(ratio, ratio)
        norm = cast("Array", xp.sqrt(xp.mean(squared)))
    elif reduction == "max":
        norm = cast("Array", xp.max(xp.abs(ratio)))
    else:
        msg = f"Unknown scaled-error reduction: {reduction!r}"
        raise ValueError(msg)
    infinity = xp.asarray(float("inf"), dtype=norm.dtype)
    return cast("Array", xp.where(xp.isfinite(norm), norm, infinity))