"""
MultiESN -- an Echo State Network with several parallel reservoirs.
Each reservoir receives its own input array; the reservoirs' state trajectories
are concatenated and fed to a single shared ridge readout. Internally this is a
composition of :class:`~memory_esn.base.BaseESN` instances (one per reservoir).
"""
from __future__ import annotations
from typing import List, Optional, Tuple, Union
import numpy as np
from sklearn.linear_model import RidgeCV
from .base import BaseESN, PersistenceMixin
Number = Union[int, float]
def _resolve_random_state(random_state, n_reservoirs: int) -> List[Optional[int]]:
"""Resolve ``random_state`` to one seed per reservoir.
* ``None`` -> ``[None] * n_reservoirs`` (non-reproducible).
* a single ``int`` -> distinct per-reservoir seeds derived from it, so the
reservoirs are reproducible *and* different from one another (a shared seed
would make identical reservoirs).
* a list/tuple/array -> used verbatim as the per-reservoir seeds; its length
must equal ``n_reservoirs`` (validated), mirroring the hyperparameter rule.
"""
if random_state is None:
return [None] * n_reservoirs
if isinstance(random_state, (list, tuple, np.ndarray)):
seeds = list(random_state)
if len(seeds) != n_reservoirs:
raise ValueError(
f"random_state must have length {n_reservoirs}, got {len(seeds)}"
)
return seeds
master = np.random.RandomState(random_state)
return [int(master.randint(0, 2 ** 31)) for _ in range(n_reservoirs)]
[docs]
class MultiESN(PersistenceMixin):
"""Multi-reservoir Echo State Network with separate inputs per reservoir.
Every per-reservoir hyperparameter accepts either a single value (broadcast
to all reservoirs) or a list of length ``n_reservoirs``.
Parameters
----------
n_reservoirs : int, default=3
Number of parallel reservoirs.
n_reservoir, spectral_radius, input_scaling, input_init, reservoir_init, bias_init, leaky, activation, bias_scaling, noise, sparsity :
Per-reservoir hyperparameters -- scalar (broadcast) or list of length
``n_reservoirs``. See :class:`~memory_esn.base.BaseESN` (``*_init`` select
the weight distribution: 'gaussian', 'uniform', 'bernoulli', 'laplace').
random_state : int, list of int, or None, default=None
Seed(s). A single int gives each reservoir a distinct *derived* seed
(reproducible but not identical). A list is used verbatim as the
per-reservoir seeds and must have length ``n_reservoirs``. None disables
reproducibility.
alphas : tuple of float, default=(0.01, 0.1, 1.0, 10.0)
Candidate ridge penalties for the shared readout.
ridge_cv_params : dict, optional
Extra keyword arguments for ``RidgeCV``.
concatenate_inputs : bool, default=True
If True, append all raw inputs to the combined states before the readout.
verbose : bool, default=False
Print a short fit summary.
Examples
--------
>>> esn = MultiESN(n_reservoirs=3, n_reservoir=100, random_state=0)
>>> esn.fit([X1, X2, X3], y, washout=100)
>>> y_pred = esn.predict([X1_test, X2_test, X3_test])
"""
def __init__(
self,
n_reservoirs: int = 3,
n_reservoir: Union[int, List[int]] = 100,
spectral_radius: Union[float, List[float]] = 0.9,
input_scaling: Union[float, List[float]] = 0.5,
input_init: Union[str, List[str]] = "uniform",
reservoir_init: Union[str, List[str]] = "uniform",
bias_init: Union[str, List[str]] = "uniform",
leaky: Union[float, List[float]] = 1.0,
activation: Union[str, List[str]] = "tanh",
bias_scaling: Union[float, List[float]] = 0.0,
noise: Union[float, List[float]] = 0.0,
sparsity: Union[float, List[float]] = 0.9,
random_state: Union[int, List[int], None] = None,
alphas: Tuple[float, ...] = (0.01, 0.1, 1.0, 10.0),
ridge_cv_params: Optional[dict] = None,
concatenate_inputs: bool = True,
verbose: bool = False,
):
self.n_reservoirs = n_reservoirs
self.alphas = alphas
self.ridge_cv_params = ridge_cv_params or {}
self.random_state = random_state
self.concatenate_inputs = concatenate_inputs
self.verbose = verbose
def to_list(param, name):
if isinstance(param, (list, tuple)):
if len(param) != n_reservoirs:
raise ValueError(
f"{name} must have length {n_reservoirs}, got {len(param)}"
)
return list(param)
return [param] * n_reservoirs
n_reservoir_list = to_list(n_reservoir, "n_reservoir")
spectral_radius_list = to_list(spectral_radius, "spectral_radius")
input_scaling_list = to_list(input_scaling, "input_scaling")
input_init_list = to_list(input_init, "input_init")
reservoir_init_list = to_list(reservoir_init, "reservoir_init")
bias_init_list = to_list(bias_init, "bias_init")
leaky_list = to_list(leaky, "leaky")
activation_list = to_list(activation, "activation")
bias_scaling_list = to_list(bias_scaling, "bias_scaling")
noise_list = to_list(noise, "noise")
sparsity_list = to_list(sparsity, "sparsity")
random_state_list = _resolve_random_state(random_state, n_reservoirs)
# One BaseESN per reservoir (composition). Concatenation of inputs is
# handled at this level, so the sub-reservoirs do not concatenate.
self.reservoirs_: List[BaseESN] = []
for i in range(n_reservoirs):
self.reservoirs_.append(
BaseESN(
n_reservoir=n_reservoir_list[i],
spectral_radius=spectral_radius_list[i],
input_scaling=input_scaling_list[i],
input_init=input_init_list[i],
reservoir_init=reservoir_init_list[i],
bias_init=bias_init_list[i],
leaky=leaky_list[i],
activation=activation_list[i],
bias_scaling=bias_scaling_list[i],
noise=noise_list[i],
sparsity=sparsity_list[i],
random_state=random_state_list[i],
alphas=alphas,
ridge_cv_params=ridge_cv_params,
concatenate_input=False,
)
)
self.readout_ = None
self.n_outputs_ = None
self._is_fitted = False
# ------------------------------------------------------------- helpers
def _validate_input_list(self, X_list, name="X_list") -> List[np.ndarray]:
if not isinstance(X_list, (list, tuple)):
raise TypeError(
f"{name} must be a list/tuple of arrays, one per reservoir. "
f"Got {type(X_list).__name__}"
)
if len(X_list) != self.n_reservoirs:
raise ValueError(
f"{name} must have {self.n_reservoirs} elements (one per reservoir), "
f"got {len(X_list)}"
)
processed = []
for X in X_list:
# 1-D univariate input -> (N, 1), not (1, N).
X = np.asarray(X)
if X.ndim == 1:
X = X.reshape(-1, 1)
processed.append(X)
n_timesteps = processed[0].shape[0]
for i, X in enumerate(processed[1:], 1):
if X.shape[0] != n_timesteps:
raise ValueError(
"All inputs must share the number of timesteps. "
f"Input 0 has {n_timesteps}, input {i} has {X.shape[0]}"
)
return processed
def _collect_states(self, X_list, x0_list, continuation) -> np.ndarray:
"""Compute and horizontally stack states (+ optional inputs) for the readout."""
all_states = []
for reservoir, X, x0 in zip(self.reservoirs_, X_list, x0_list):
if reservoir.W_in_ is None:
reservoir._initialize_weights(X.shape[1])
all_states.append(reservoir._compute_states(X, x0=x0, continuation=continuation))
combined = np.hstack(all_states)
if self.concatenate_inputs:
combined = np.hstack([combined, np.hstack(X_list)])
return combined
# ----------------------------------------------------------------- API
[docs]
def fit(
self,
X_list: List[np.ndarray],
y: np.ndarray,
washout: int = 0,
x0_list: Optional[List[np.ndarray]] = None,
) -> "MultiESN":
"""Fit the shared readout over all reservoirs.
Parameters
----------
X_list : list of ndarray
One input sequence per reservoir (all with the same length).
y : ndarray, shape (n_timesteps, n_outputs)
washout : int, default=0
x0_list : list of ndarray, optional
Initial state per reservoir.
"""
X_list = self._validate_input_list(X_list)
# 1-D target -> (N, 1), not (1, N).
y = np.asarray(y)
if y.ndim == 1:
y = y.reshape(-1, 1)
n_timesteps = X_list[0].shape[0]
self.n_outputs_ = y.shape[1]
if y.shape[0] != n_timesteps:
raise ValueError("X_list and y must have the same number of timesteps.")
if x0_list is None:
x0_list = [None] * self.n_reservoirs
extended = self._collect_states(X_list, x0_list, continuation=False)
if washout > 0:
extended = extended[washout:]
y = y[washout:]
self.readout_ = RidgeCV(alphas=self.alphas, **self.ridge_cv_params)
self.readout_.fit(extended, y)
self._is_fitted = True
if self.verbose:
print(
f"MultiESN fitted: {self.n_reservoirs} reservoirs, "
f"sizes={[r.n_reservoir for r in self.reservoirs_]}, "
f"extended_dim={extended.shape[1]}, best_alpha={self.readout_.alpha_}"
)
return self
[docs]
def predict(
self,
X_list: List[np.ndarray],
x0_list: Optional[List[np.ndarray]] = None,
continuation: bool = False,
) -> np.ndarray:
"""Predict outputs for a list of per-reservoir input sequences."""
if not self._is_fitted:
raise RuntimeError("Model must be fitted before prediction. Call fit() first.")
X_list = self._validate_input_list(X_list)
if x0_list is None:
x0_list = [None] * self.n_reservoirs
extended = self._collect_states(X_list, x0_list, continuation=continuation)
pred = self.readout_.predict(extended)
if pred.ndim == 1:
pred = pred.reshape(-1, self.n_outputs_)
return pred
# ------------------------------------------------------------- inspection
[docs]
def get_reservoir_states(
self,
X_list: List[np.ndarray],
x0_list: Optional[List[np.ndarray]] = None,
reservoir_idx: Optional[int] = None,
continuation: bool = False,
) -> Union[np.ndarray, List[np.ndarray]]:
"""Return states for one reservoir (``reservoir_idx``) or all of them."""
X_list = self._validate_input_list(X_list)
if x0_list is None:
x0_list = [None] * self.n_reservoirs
if reservoir_idx is not None:
return self.reservoirs_[reservoir_idx]._compute_states(
X_list[reservoir_idx],
x0=x0_list[reservoir_idx],
continuation=continuation,
)
return [
reservoir._compute_states(X, x0=x0, continuation=continuation)
for reservoir, X, x0 in zip(self.reservoirs_, X_list, x0_list)
]
[docs]
def get_extended_states(
self,
X_list: List[np.ndarray],
x0_list: Optional[List[np.ndarray]] = None,
continuation: bool = False,
) -> np.ndarray:
"""Return the concatenated states (+ inputs) exactly as fed to the readout."""
X_list = self._validate_input_list(X_list)
if x0_list is None:
x0_list = [None] * self.n_reservoirs
return self._collect_states(X_list, x0_list, continuation=continuation)
[docs]
def reset_states(self) -> None:
"""Forget the cached state of every reservoir."""
for reservoir in self.reservoirs_:
reservoir.reset_state()
def __repr__(self) -> str:
return (
f"MultiESN(n_reservoirs={self.n_reservoirs}, "
f"sizes={[r.n_reservoir for r in self.reservoirs_]})"
)