104 lines
3.0 KiB
Python
104 lines
3.0 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Callable
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SolveIVPConfig:
|
|
t_start: float = 0.0
|
|
t_stop: float = 20.0
|
|
method: str = "BDF"
|
|
rtol: float = 1e-6
|
|
atol: float = 1e-8
|
|
max_step: float = 1e-3
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ODESolution:
|
|
t: list[float]
|
|
y: list[list[float]]
|
|
success: bool
|
|
message: str
|
|
|
|
|
|
def _vector_add(a: list[float], b: list[float], scale: float = 1.0) -> list[float]:
|
|
return [x + scale * y for x, y in zip(a, b)]
|
|
|
|
|
|
def _runge_kutta_4(
|
|
rhs: Callable[[float, list[float]], list[float]],
|
|
initial_state: list[float],
|
|
config: SolveIVPConfig,
|
|
t_eval: list[float] | None,
|
|
) -> ODESolution:
|
|
if t_eval is None:
|
|
point_count = max(
|
|
2,
|
|
int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1,
|
|
)
|
|
step = (config.t_stop - config.t_start) / (point_count - 1)
|
|
t_eval = [config.t_start + index * step for index in range(point_count)]
|
|
|
|
state = list(initial_state)
|
|
states = [[value] for value in state]
|
|
times = [float(t_eval[0])]
|
|
current_time = float(t_eval[0])
|
|
|
|
for target_time in t_eval[1:]:
|
|
while current_time < target_time - 1e-15:
|
|
dt = min(config.max_step, target_time - current_time)
|
|
k1 = rhs(current_time, state)
|
|
k2 = rhs(current_time + 0.5 * dt, _vector_add(state, k1, 0.5 * dt))
|
|
k3 = rhs(current_time + 0.5 * dt, _vector_add(state, k2, 0.5 * dt))
|
|
k4 = rhs(current_time + dt, _vector_add(state, k3, dt))
|
|
state = [
|
|
value + (dt / 6.0) * (a + 2.0 * b + 2.0 * c + d)
|
|
for value, a, b, c, d in zip(state, k1, k2, k3, k4)
|
|
]
|
|
current_time += dt
|
|
|
|
times.append(float(target_time))
|
|
for index, value in enumerate(state):
|
|
states[index].append(value)
|
|
|
|
return ODESolution(
|
|
t=times,
|
|
y=states,
|
|
success=True,
|
|
message="Integrated with built-in RK4 fallback because SciPy is unavailable.",
|
|
)
|
|
|
|
|
|
def integrate_ode(
|
|
rhs: Callable[[float, list[float]], list[float]],
|
|
initial_state: list[float],
|
|
config: SolveIVPConfig,
|
|
t_eval: list[float] | None = None,
|
|
):
|
|
"""Thin wrapper around scipy.integrate.solve_ivp with a pure-Python fallback."""
|
|
|
|
if abs(config.t_stop - config.t_start) <= 1e-15:
|
|
return ODESolution(
|
|
t=[float(config.t_start)],
|
|
y=[[value] for value in initial_state],
|
|
success=True,
|
|
message="Skipped integration because t_start equals t_stop.",
|
|
)
|
|
|
|
try:
|
|
from scipy.integrate import solve_ivp
|
|
except ImportError:
|
|
return _runge_kutta_4(rhs, initial_state, config, t_eval)
|
|
|
|
return solve_ivp(
|
|
fun=rhs,
|
|
t_span=(config.t_start, config.t_stop),
|
|
y0=initial_state,
|
|
method=config.method,
|
|
rtol=config.rtol,
|
|
atol=config.atol,
|
|
max_step=config.max_step,
|
|
t_eval=t_eval,
|
|
)
|