File size: 2,306 Bytes
9147d47 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 | from __future__ import annotations
from dataclasses import dataclass
import numpy as np
@dataclass
class RecurrentFeatures:
input_weight: np.ndarray
recurrent_weight: np.ndarray
bias: np.ndarray
readout: np.ndarray | None = None
@property
def hidden_size(self) -> int:
return len(self.bias)
def states(self, series: np.ndarray) -> np.ndarray:
hidden = np.zeros(self.hidden_size, dtype=np.float64)
outputs = np.empty((len(series), self.hidden_size), dtype=np.float64)
for index, value in enumerate(series):
hidden = np.tanh(
self.input_weight * value
+ self.recurrent_weight @ hidden
+ self.bias
)
outputs[index] = hidden
return outputs
def design(self, series: np.ndarray) -> np.ndarray:
states = self.states(series)
return np.column_stack([states, series, np.ones(len(series))])
def predict(self, series: np.ndarray) -> np.ndarray:
if self.readout is None:
raise RuntimeError("Fit or load a readout before prediction.")
return self.design(series) @ self.readout
def unpack_genome(genome: np.ndarray, hidden_size: int) -> RecurrentFeatures:
cursor = 0
input_weight = genome[cursor : cursor + hidden_size]
cursor += hidden_size
raw_recurrent = genome[cursor : cursor + hidden_size**2].reshape(
hidden_size,
hidden_size,
)
cursor += hidden_size**2
bias = genome[cursor : cursor + hidden_size]
radius = max(abs(np.linalg.eigvals(raw_recurrent)).max(), 1e-8)
recurrent = raw_recurrent * (0.95 / max(float(radius), 0.95))
return RecurrentFeatures(input_weight, recurrent, bias)
def fit_readout(
model: RecurrentFeatures,
series: np.ndarray,
targets: np.ndarray,
indices: np.ndarray,
ridge: float = 1e-4,
) -> np.ndarray:
design = model.design(series)[indices]
gram = design.T @ design + ridge * np.eye(design.shape[1])
model.readout = np.linalg.solve(gram, design.T @ targets[indices])
return model.readout
def parameter_count(hidden_size: int, horizons: int) -> int:
recurrent = hidden_size + hidden_size**2 + hidden_size
readout = (hidden_size + 2) * horizons
return recurrent + readout
|