"""
kr_reservoir.py
===============
Core K-R Reservoir Architecture implementation.

Three components:
    1. Multi-timescale dual reservoir (fast alpha=1.0, slow alpha<=3/M)
    2. Adaptive delay embedding (d=10 delays spanning [1, M])
    3. State normalisation before ridge regression

References:
    Pasupuleti, R. (2026). Improved K-R Reservoir Architecture for
    Long-Memory Time Series Prediction. Neurocomputing.

Usage:
    from kr_reservoir import KRReservoir
    model = KRReservoir(N=200, M=50, seed=0)
    model.fit(u_train, y_train)
    nrmse = model.evaluate(u_test, y_test)
"""

import numpy as np
from typing import List, Optional, Tuple


# ── Adaptive-K scaling formula (Eq. 6 in paper) ─────────────────
def adaptive_k(N: int, K0: float = 1.5, N0: int = 200, beta: float = 0.40) -> float:
    """
    Adaptive input weight scaling to prevent tanh saturation.
    K = K0 * (N0 / N)^beta
    Keeps P(|Ku| > 2) < 18% for u ~ N(0,1), N in [50, 500].
    """
    return K0 * (N0 / N) ** beta


def make_reservoir(N: int, rho: float, seed: int) -> np.ndarray:
    """
    Create a random reservoir weight matrix W with spectral radius rho.
    W ~ N(0,1), then scaled so max(|eigenvalues(W)|) = rho.
    """
    np.random.seed(seed)
    W = np.random.randn(N, N)
    eigvals = np.linalg.eigvals(W)
    W *= rho / (np.max(np.abs(eigvals)) + 1e-12)
    return W


def make_input_weights(N: int, scale: float, seed: int) -> np.ndarray:
    """Create random input weight vector W_in ~ N(0, scale)."""
    np.random.seed(seed + 1000)
    return np.random.randn(N) * scale


def adaptive_delays(M: int, d: int = 10) -> List[int]:
    """
    Adaptive delay set: d uniformly spaced delays covering [1, M].
    Eq. 7: tau = {ceil(M/d), 2*ceil(M/d), ..., M}

    Parameters
    ----------
    M : int   Task memory order
    d : int   Number of delays (default 10)

    Returns
    -------
    List[int]  Delay indices (1-indexed)
    """
    step = max(1, int(np.ceil(M / d)))
    return list(range(step, M + 1, step))[:d]


class KRReservoir:
    """
    Improved K-R Reservoir Architecture.

    Parameters
    ----------
    N       : int    Total reservoir size (split equally fast/slow)
    M       : int    Task memory order (controls alpha_s and delay range)
    rho_f   : float  Spectral radius of fast reservoir
    rho_s   : float  Spectral radius of slow reservoir
    alpha_s : float  Leak rate of slow reservoir (None -> auto = 3/M)
    lam     : float  Ridge regression regularisation parameter
    d       : int    Number of delay features (default 10)
    seed    : int    Random seed
    """

    def __init__(
        self,
        N: int = 200,
        M: int = 50,
        rho_f: float = 0.9,
        rho_s: float = 0.9,
        alpha_s: Optional[float] = None,
        lam: float = 1e-4,
        d: int = 10,
        seed: int = 0,
        washout: int = 200,
    ):
        self.N       = N
        self.M       = M
        self.rho_f   = rho_f
        self.rho_s   = rho_s
        self.alpha_s = alpha_s if alpha_s is not None else min(3.0 / M, 1.0)
        self.lam     = lam
        self.d       = d
        self.seed    = seed
        self.washout = washout

        self.Nf = N // 2
        self.Ns = N - self.Nf

        # Build fast reservoir
        KF = adaptive_k(N)
        self.Wf   = make_reservoir(self.Nf, rho_f, seed)
        self.Wif  = make_input_weights(self.Nf, KF, seed)

        # Build slow reservoir
        KS = adaptive_k(N, K0=1.0)
        self.Ws   = make_reservoir(self.Ns, rho_s, seed + 500)
        self.Wis  = make_input_weights(self.Ns, KS, seed + 500)

        # Delay indices
        self.delays = adaptive_delays(M, d)
        self.max_delay = max(self.delays) + 1

        # Will be set during fit
        self.W_out = None
        self.mu_train  = None
        self.sig_train = None

    # ── State collection ─────────────────────────────────────────
    def _collect_states(self, u: np.ndarray) -> np.ndarray:
        """
        Run K-R reservoir on input u, return feature matrix Phi.
        Shape: (T, Nf + Ns + d)
        """
        T = len(u)
        xf = np.zeros(self.Nf)
        xs = np.zeros(self.Ns)
        db = np.zeros(self.max_delay)  # delay buffer (circular)
        Phi = np.zeros((T, self.Nf + self.Ns + self.d))

        for t in range(T):
            # Fast reservoir update (no leak)
            xf = np.tanh(self.Wf @ xf + self.Wif * u[t])
            # Slow reservoir update (leaky)
            xs = ((1 - self.alpha_s) * xs
                  + self.alpha_s * np.tanh(self.Ws @ xs + self.Wis * u[t]))
            # Delay buffer: shift right, insert new value at index 0
            db = np.roll(db, 1)
            db[0] = u[t]
            # Feature vector: [x_f | x_s | u(t-tau_1) ... u(t-tau_d)]
            Phi[t] = np.concatenate([xf, xs, db[self.delays]])

        return Phi

    # ── Normalisation ─────────────────────────────────────────────
    def _normalise(self, Phi: np.ndarray, fit: bool = False) -> np.ndarray:
        """
        Normalise reservoir states (first N columns only).
        Delay features are pre-normalised inputs — excluded.
        """
        Phi = Phi.copy()
        if fit:
            self.mu_train  = Phi[:, :self.N].mean(axis=0)
            self.sig_train = Phi[:, :self.N].std(axis=0) + 1e-8
        Phi[:, :self.N] = (Phi[:, :self.N] - self.mu_train) / self.sig_train
        return Phi

    # ── Fit ────────────────────────────────────────────────────────
    def fit(self, u: np.ndarray, y: np.ndarray) -> "KRReservoir":
        """
        Train the K-R readout on training data.

        Parameters
        ----------
        u : ndarray (T,)  Normalised input time series
        y : ndarray (T,)  Normalised target time series

        Returns
        -------
        self
        """
        T = len(u)
        Phi_raw = self._collect_states(u)

        # Apply normalisation using training statistics
        Phi = self._normalise(Phi_raw[self.washout:], fit=True)
        y_tr = y[self.washout:]

        # Ridge regression: W_out = (Phi^T Phi + lambda I)^{-1} Phi^T y
        A = Phi.T @ Phi + self.lam * np.eye(Phi.shape[1])
        self.W_out = np.linalg.solve(A, Phi.T @ y_tr)
        return self

    # ── Predict ────────────────────────────────────────────────────
    def predict(self, u: np.ndarray) -> np.ndarray:
        """
        Predict on new input u (no washout applied).

        Parameters
        ----------
        u : ndarray (T,)  Normalised input

        Returns
        -------
        y_hat : ndarray (T,)  Predictions
        """
        if self.W_out is None:
            raise RuntimeError("Call fit() before predict()")
        Phi_raw = self._collect_states(u)
        Phi = self._normalise(Phi_raw, fit=False)
        return Phi @ self.W_out

    # ── Evaluate ───────────────────────────────────────────────────
    def evaluate(self, u: np.ndarray, y: np.ndarray,
                 test_size: int = 500) -> Tuple[float, np.ndarray, np.ndarray]:
        """
        Run prediction on last test_size steps, return NRMSE + traces.

        Parameters
        ----------
        u         : ndarray (T,)  Full input series (train + test)
        y         : ndarray (T,)  Full target series
        test_size : int           Number of test steps

        Returns
        -------
        nrmse : float      NRMSE on test segment
        y_te  : ndarray    True test targets
        y_hat : ndarray    Predicted values
        """
        T = len(u)
        Phi_raw = self._collect_states(u)
        # Normalise using training statistics
        Phi = self._normalise(Phi_raw, fit=False)
        Phi_te = Phi[T - test_size:]
        y_te   = y[T - test_size:]
        y_hat  = Phi_te @ self.W_out
        nrmse  = float(np.sqrt(np.mean((y_te - y_hat) ** 2))
                       / (np.std(y_te) + 1e-12))
        return nrmse, y_te, y_hat

    def __repr__(self):
        return (f"KRReservoir(N={self.N}, M={self.M}, "
                f"rho_f={self.rho_f}, rho_s={self.rho_s}, "
                f"alpha_s={self.alpha_s:.4f}, d={self.d}, "
                f"delays={self.delays[:3]}...)")
