Source code for memory_esn.multi

"""
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_]})" )