def _stable_expit(x):
"""Numerically stable sigmoid 1/(1+exp(-x)), scalar or array."""
x = np.asarray(x, dtype=float)
return np.where(
x >= 0,
1.0 / (1.0 + np.exp(-x)),
np.exp(x) / (1.0 + np.exp(x)),
)
def _jr_sigmoid(v, v_max=5.0, r=0.56, v0=6.0):
"""Jansen-Rit population sigmoid: S(v) = v_max / (1 + exp(r*(v0 - v)))."""
return v_max * _stable_expit(r * (v - v0))
def rhs(t: float, y: np.ndarray) -> np.ndarray:
y0v, y1v, y2v, y3v, y4v = y[0], y[1], y[2], y[3], y[4]
y5v, y6v, y7v, y8v, y9v = y[5], y[6], y[7], y[8], y[9]
S_pyr = _jr_sigmoid(y1v - y2v - y3v)
S_exc = _jr_sigmoid(C1 * y0v)
S_slow = _jr_sigmoid(C3 * y0v)
S_fast = _jr_sigmoid(C5 * y0v - C6 * y4v)
S_self = _jr_sigmoid(C3 * y0v)
return np.array([
y5v, # ẏ₀
y6v, # ẏ₁
y7v, # ẏ₂
y8v, # ẏ₃
y9v, # ẏ₄
A * a * S_pyr - 2.0 * a * y5v - a2 * y0v, # ẏ₅
A * a * (p + C2 * S_exc) - 2.0 * a * y6v - a2 * y1v, # ẏ₆
B * b * C4 * S_slow - 2.0 * b * y7v - b2 * y2v, # ẏ₇
G * g * C7 * S_fast - 2.0 * g * y8v - g2 * y3v, # ẏ₈
B * b * S_self - 2.0 * b * y9v - b2 * y4v, # ẏ₉
])