ProCreations's picture
Publish validated HiPPO Zoo reproduction nB0TrIRAs1
1cd8a52 verified
Raw
History Blame Contribute Delete
2.44 kB
"""Numerical integrators for multiscale polynomial memory systems."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Iterable, Optional, Sequence, Tuple, Union
import numpy as np
from scipy.integrate import solve_ivp
from .utils import vec, unvec
State = np.ndarray
Params = Union[Tuple[np.ndarray, np.ndarray], Tuple[np.ndarray, np.ndarray, np.ndarray]]
@dataclass
class IntegrationResult:
"""Container holding the final state and, optionally, the trajectory."""
state: State
trajectory: Optional[np.ndarray] = None
def _validate_inputs(xs: np.ndarray, dt: float) -> None:
if xs.ndim != 1:
raise ValueError("`xs` must be a one-dimensional array")
if not np.isscalar(dt) or dt <= 0:
raise ValueError("`dt` must be a positive scalar")
def integrate_system(
xs: np.ndarray,
dt: float,
params: Params,
*,
state: Optional[np.ndarray] = None,
return_trajectory: bool = False,
) -> IntegrationResult:
"""Integrate a discrete control signal with either single or multiscale parameters."""
_validate_inputs(xs, dt)
multiscale = len(params) == 3
if multiscale:
A, B, M = params
n, m = B.shape
if state is None:
state = np.zeros(n * m, dtype=float)
elif state.shape != (n * m,):
raise ValueError(f"Expected state of shape {(n * m,)}, got {state.shape}")
def get_fun(x_val: float):
return lambda _, y: vec(A @ unvec(y, n, m) @ M + x_val * B)
else:
A, b = params
d = b.shape[0]
if state is None:
state = np.zeros(d, dtype=float)
elif state.shape != (d,):
raise ValueError(f"Expected state of shape {(d,)}, got {state.shape}")
def get_fun(x_val: float):
return lambda _, y: A @ y + b * x_val
trajectory: Optional[list[np.ndarray]] = [] if return_trajectory else None
for x_val in xs:
fun = get_fun(float(x_val))
result = solve_ivp(fun, (0.0, float(dt)), state, method="RK45")
state = result.y[:, -1]
if return_trajectory:
trajectory.append(state.copy())
if return_trajectory and trajectory is not None:
history = np.stack(trajectory, axis=0)
else:
history = None
return IntegrationResult(state=state, trajectory=history)
__all__ = ["IntegrationResult", "integrate_system"]