--- Final Solver (used for performance test) ---
import numpy as np
from numba import njit
from scipy.integrate._ivp import dop853_coefficients as _dc

_A = np.ascontiguousarray(_dc.A[:12, :12])
_B = np.ascontiguousarray(_dc.B)
_C = np.ascontiguousarray(_dc.C[:12])
_E3 = np.ascontiguousarray(_dc.E3)
_E5 = np.ascontiguousarray(_dc.E5)


@njit(cache=True, fastmath=True, inline='always')
def _rhs3(S, E, I, beta, sigma, gamma, omega):
    bSI = beta * S * I
    return (-bSI + omega * (1.0 - S - E - I), bSI - sigma * E, sigma * E - gamma * I)


@njit(cache=True, fastmath=True)
def _rk45_3d(t0, t1, S, E, I, beta, sigma, gamma, omega, rtol, atol):
    a21 = 0.2
    a31 = 3.0/40.0; a32 = 9.0/40.0
    a41 = 44.0/45.0; a42 = -56.0/15.0; a43 = 32.0/9.0
    a51 = 19372.0/6561.0; a52 = -25360.0/2187.0; a53 = 64448.0/6561.0; a54 = -212.0/729.0
    a61 = 9017.0/3168.0; a62 = -355.0/33.0; a63 = 46732.0/5247.0; a64 = 49.0/176.0; a65 = -5103.0/18656.0
    b1 = 35.0/384.0; b3 = 500.0/1113.0; b4 = 125.0/192.0; b5 = -2187.0/6784.0; b6 = 11.0/84.0
    e1 = 71.0/57600.0; e3 = -71.0/16695.0; e4 = 71.0/1920.0; e5 = -17253.0/339200.0; e6 = 22.0/525.0; e7 = -1.0/40.0
    t = t0
    direction = 1.0 if t1 >= t0 else -1.0
    k1S, k1E, k1I = _rhs3(S, E, I, beta, sigma, gamma, omega)
    sc0 = atol + abs(S)*rtol; sc1 = atol + abs(E)*rtol; sc2 = atol + abs(I)*rtol
    d0 = np.sqrt(((S/sc0)**2 + (E/sc1)**2 + (I/sc2)**2)/3.0)
    d1 = np.sqrt(((k1S/sc0)**2 + (k1E/sc1)**2 + (k1I/sc2)**2)/3.0)
    if d0 < 1e-5 or d1 < 1e-5:
        h0 = 1e-6
    else:
        h0 = 0.01 * d0 / d1
    yS = S + h0*direction*k1S; yE = E + h0*direction*k1E; yI = I + h0*direction*k1I
    f2S, f2E, f2I = _rhs3(yS, yE, yI, beta, sigma, gamma, omega)
    d2 = np.sqrt((((f2S-k1S)/sc0)**2 + ((f2E-k1E)/sc1)**2 + ((f2I-k1I)/sc2)**2)/3.0)/h0
    if d1 <= 1e-15 and d2 <= 1e-15:
        h1 = max(1e-6, h0*1e-3)
    else:
        h1 = (0.01/max(d1, d2))**0.2
    h = min(100.0*h0, h1)
    span = abs(t1-t0)
    if h > span:
        h = span
    SAFETY = 0.9; MIN_FACTOR = 0.2; MAX_FACTOR = 10.0
    rejected = False
    while (t - t1)*direction < 0.0:
        rem = abs(t1-t)
        if h > rem:
            h = rem
        hd = h*direction
        yS = S+hd*a21*k1S; yE = E+hd*a21*k1E; yI = I+hd*a21*k1I
        k2S,k2E,k2I = _rhs3(yS,yE,yI,beta,sigma,gamma,omega)
        yS = S+hd*(a31*k1S+a32*k2S); yE = E+hd*(a31*k1E+a32*k2E); yI = I+hd*(a31*k1I+a32*k2I)
        k3S,k3E,k3I = _rhs3(yS,yE,yI,beta,sigma,gamma,omega)
        yS = S+hd*(a41*k1S+a42*k2S+a43*k3S); yE = E+hd*(a41*k1E+a42*k2E+a43*k3E); yI = I+hd*(a41*k1I+a42*k2I+a43*k3I)
        k4S,k4E,k4I = _rhs3(yS,yE,yI,beta,sigma,gamma,omega)
        yS = S+hd*(a51*k1S+a52*k2S+a53*k3S+a54*k4S); yE = E+hd*(a51*k1E+a52*k2E+a53*k3E+a54*k4E); yI = I+hd*(a51*k1I+a52*k2I+a53*k3I+a54*k4I)
        k5S,k5E,k5I = _rhs3(yS,yE,yI,beta,sigma,gamma,omega)
        yS = S+hd*(a61*k1S+a62*k2S+a63*k3S+a64*k4S+a65*k5S); yE = E+hd*(a61*k1E+a62*k2E+a63*k3E+a64*k4E+a65*k5E); yI = I+hd*(a61*k1I+a62*k2I+a63*k3I+a64*k4I+a65*k5I)
        k6S,k6E,k6I = _rhs3(yS,yE,yI,beta,sigma,gamma,omega)
        newS = S+hd*(b1*k1S+b3*k3S+b4*k4S+b5*k5S+b6*k6S)
        newE = E+hd*(b1*k1E+b3*k3E+b4*k4E+b5*k5E+b6*k6E)
        newI = I+hd*(b1*k1I+b3*k3I+b4*k4I+b5*k5I+b6*k6I)
        k7S,k7E,k7I = _rhs3(newS,newE,newI,beta,sigma,gamma,omega)
        errS = hd*(e1*k1S+e3*k3S+e4*k4S+e5*k5S+e6*k6S+e7*k7S)
        errE = hd*(e1*k1E+e3*k3E+e4*k4E+e5*k5E+e6*k6E+e7*k7E)
        errI = hd*(e1*k1I+e3*k3I+e4*k4I+e5*k5I+e6*k6I+e7*k7I)
        scS = atol+max(abs(S),abs(newS))*rtol; scE = atol+max(abs(E),abs(newE))*rtol; scI = atol+max(abs(I),abs(newI))*rtol
        err = np.sqrt(((errS/scS)**2+(errE/scE)**2+(errI/scI)**2)/3.0)
        if err <= 1.0:
            t += hd
            S = newS; E = newE; I = newI
            k1S = k7S; k1E = k7E; k1I = k7I
            if err == 0.0:
                factor = MAX_FACTOR
            else:
                factor = SAFETY*err**(-0.2)
                if factor > MAX_FACTOR:
                    factor = MAX_FACTOR
            if rejected and factor > 1.0:
                factor = 1.0
            h *= factor
            rejected = False
        else:
            factor = SAFETY*err**(-0.2)
            if factor < MIN_FACTOR:
                factor = MIN_FACTOR
            h *= factor
            rejected = True
    return S, E, I, 1.0 - S - E - I


@njit(cache=True, fastmath=True)
def _dop853_3d(t0, t1, y0, beta, sigma, gamma, omega, rtol, atol, A, B, E3, E5):
    n = 3
    K = np.empty((13, n))
    y = y0.copy()
    ytmp = np.empty(n)
    t = t0
    direction = 1.0 if t1 >= t0 else -1.0
    r0 = _rhs3(y[0], y[1], y[2], beta, sigma, gamma, omega)
    K[0,0]=r0[0]; K[0,1]=r0[1]; K[0,2]=r0[2]
    sc0 = atol+abs(y[0])*rtol; sc1 = atol+abs(y[1])*rtol; sc2 = atol+abs(y[2])*rtol
    d0 = np.sqrt(((y[0]/sc0)**2+(y[1]/sc1)**2+(y[2]/sc2)**2)/n)
    d1 = np.sqrt(((K[0,0]/sc0)**2+(K[0,1]/sc1)**2+(K[0,2]/sc2)**2)/n)
    if d0 < 1e-5 or d1 < 1e-5:
        h0 = 1e-6
    else:
        h0 = 0.01*d0/d1
    ytmp[0]=y[0]+h0*direction*K[0,0]; ytmp[1]=y[1]+h0*direction*K[0,1]; ytmp[2]=y[2]+h0*direction*K[0,2]
    f1 = _rhs3(ytmp[0], ytmp[1], ytmp[2], beta, sigma, gamma, omega)
    d2 = np.sqrt((((f1[0]-K[0,0])/sc0)**2+((f1[1]-K[0,1])/sc1)**2+((f1[2]-K[0,2])/sc2)**2)/n)/h0
    if d1 <= 1e-15 and d2 <= 1e-15:
        h1 = max(1e-6, h0*1e-3)
    else:
        h1 = (0.01/max(d1, d2))**(1.0/8.0)
    h_abs = min(100.0*h0, h1)
    span = abs(t1-t0)
    if h_abs > span:
        h_abs = span
    SAFETY = 0.9; MIN_FACTOR = 0.2; MAX_FACTOR = 10.0
    rejected = False
    while (t - t1)*direction < 0.0:
        rem = abs(t1-t)
        if h_abs > rem:
            h_abs = rem
        h = h_abs*direction
        for s in range(1, 12):
            for d in range(n):
                acc = 0.0
                for j in range(s):
                    acc += A[s, j]*K[j, d]
                ytmp[d] = y[d] + h*acc
            ks = _rhs3(ytmp[0], ytmp[1], ytmp[2], beta, sigma, gamma, omega)
            K[s,0]=ks[0]; K[s,1]=ks[1]; K[s,2]=ks[2]
        for d in range(n):
            acc = 0.0
            for j in range(12):
                acc += B[j]*K[j, d]
            ytmp[d] = y[d] + h*acc
        fn = _rhs3(ytmp[0], ytmp[1], ytmp[2], beta, sigma, gamma, omega)
        K[12,0]=fn[0]; K[12,1]=fn[1]; K[12,2]=fn[2]
        err5_2 = 0.0; err3_2 = 0.0
        for d in range(n):
            sc = atol + max(abs(y[d]), abs(ytmp[d]))*rtol
            e5 = 0.0; e3 = 0.0
            for j in range(13):
                e5 += E5[j]*K[j, d]
                e3 += E3[j]*K[j, d]
            e5 /= sc; e3 /= sc
            err5_2 += e5*e5; err3_2 += e3*e3
        if err5_2 == 0.0 and err3_2 == 0.0:
            err_norm = 0.0
        else:
            denom = err5_2 + 0.01*err3_2
            err_norm = abs(h)*err5_2/np.sqrt(denom*n)
        if err_norm < 1.0:
            t += h
            for d in range(n):
                y[d] = ytmp[d]
                K[0, d] = K[12, d]
            if err_norm == 0.0:
                factor = MAX_FACTOR
            else:
                factor = SAFETY*err_norm**(-1.0/8.0)
                if factor > MAX_FACTOR:
                    factor = MAX_FACTOR
            if rejected and factor > 1.0:
                factor = 1.0
            h_abs *= factor
            rejected = False
        else:
            factor = SAFETY*err_norm**(-1.0/8.0)
            if factor < MIN_FACTOR:
                factor = MIN_FACTOR
            h_abs *= factor
            rejected = True
    return y[0], y[1], y[2], 1.0 - y[0] - y[1] - y[2]


class Solver:
    def __init__(self):
        self.method = 'rk45'
        self.rtol = 1e-6
        self.atol = 1e-8
        _rk45_3d(0.0, 5.0, 0.89, 0.01, 0.005, 0.35, 0.2, 0.1, 0.002, 1e-6, 1e-8)
        _dop853_3d(0.0, 5.0, np.array([0.89, 0.01, 0.005]), 0.35, 0.2, 0.1, 0.002, 1e-6, 1e-8, _A, _B, _E3, _E5)

    def solve(self, problem, **kwargs):
        y0 = problem['y0']; p = problem['params']
        t0 = float(problem['t0']); t1 = float(problem['t1'])
        b = float(p['beta']); sg = float(p['sigma']); g = float(p['gamma']); om = float(p['omega'])
        if self.method == 'rk45':
            S, E, I, R = _rk45_3d(t0, t1, float(y0[0]), float(y0[1]), float(y0[2]), b, sg, g, om, self.rtol, self.atol)
        else:
            S, E, I, R = _dop853_3d(t0, t1, np.array([float(y0[0]), float(y0[1]), float(y0[2])]), b, sg, g, om, self.rtol, self.atol, _A, _B, _E3, _E5)
        return [S, E, I, R]
--- End Solver ---
Running performance test...
============================= test session starts ==============================
platform linux -- Python 3.12.13, pytest-9.0.2, pluggy-1.6.0
rootdir: /tests
plugins: anyio-4.13.0, jaxtyping-0.3.9
collected 3 items

../tests/test_outputs.py .
--- Performance Summary ---
Validity: True
Total Baseline Time: 8.1638s
Total Solver Time:   0.0145s
Raw Speedup:         562.8215 x
Final Reward (Score): 562.8215
---------------------------
..

==================================== PASSES ====================================
=========================== short test summary info ============================
PASSED ../tests/test_outputs.py::test_solver_exists
PASSED ../tests/test_outputs.py::test_solver_validity
PASSED ../tests/test_outputs.py::test_solver_speedup
======================== 3 passed in 104.10s (0:01:44) =========================
