--- Final Solver (used for performance test) ---
import numpy as np
from numba import njit
from typing import Any


@njit(cache=True)
def rk45_scalar(t0, t1, S0, E0, I0, R0, beta, sigma, gamma, omega, rtol, atol):
    S, E, I, R = S0, E0, I0, R0
    t = t0
    h = (t1 - t0) * 1e-3
    if h <= 0.0:
        h = 1e-6

    # Initial k1
    bSI = beta * S * I
    k1_S = -bSI + omega * R
    k1_E = bSI - sigma * E
    k1_I = sigma * E - gamma * I
    k1_R = gamma * I - omega * R

    while t < t1:
        if t + h > t1:
            h = t1 - t

        # Stage 2
        S2 = S + h * 0.2 * k1_S
        E2 = E + h * 0.2 * k1_E
        I2 = I + h * 0.2 * k1_I
        R2 = R + h * 0.2 * k1_R
        bSI = beta * S2 * I2
        k2_S = -bSI + omega * R2
        k2_E = bSI - sigma * E2
        k2_I = sigma * E2 - gamma * I2
        k2_R = gamma * I2 - omega * R2

        # Stage 3
        S3 = S + h * (0.075*k1_S + 0.225*k2_S)
        E3 = E + h * (0.075*k1_E + 0.225*k2_E)
        I3 = I + h * (0.075*k1_I + 0.225*k2_I)
        R3 = R + h * (0.075*k1_R + 0.225*k2_R)
        bSI = beta * S3 * I3
        k3_S = -bSI + omega * R3
        k3_E = bSI - sigma * E3
        k3_I = sigma * E3 - gamma * I3
        k3_R = gamma * I3 - omega * R3

        # Stage 4
        S4 = S + h * (44.0/45.0*k1_S + (-56.0/15.0)*k2_S + 32.0/9.0*k3_S)
        E4 = E + h * (44.0/45.0*k1_E + (-56.0/15.0)*k2_E + 32.0/9.0*k3_E)
        I4 = I + h * (44.0/45.0*k1_I + (-56.0/15.0)*k2_I + 32.0/9.0*k3_I)
        R4 = R + h * (44.0/45.0*k1_R + (-56.0/15.0)*k2_R + 32.0/9.0*k3_R)
        bSI = beta * S4 * I4
        k4_S = -bSI + omega * R4
        k4_E = bSI - sigma * E4
        k4_I = sigma * E4 - gamma * I4
        k4_R = gamma * I4 - omega * R4

        # Stage 5
        S5 = S + h * (19372.0/6561.0*k1_S + (-25360.0/2187.0)*k2_S + 64448.0/6561.0*k3_S + (-212.0/729.0)*k4_S)
        E5 = E + h * (19372.0/6561.0*k1_E + (-25360.0/2187.0)*k2_E + 64448.0/6561.0*k3_E + (-212.0/729.0)*k4_E)
        I5 = I + h * (19372.0/6561.0*k1_I + (-25360.0/2187.0)*k2_I + 64448.0/6561.0*k3_I + (-212.0/729.0)*k4_I)
        R5 = R + h * (19372.0/6561.0*k1_R + (-25360.0/2187.0)*k2_R + 64448.0/6561.0*k3_R + (-212.0/729.0)*k4_R)
        bSI = beta * S5 * I5
        k5_S = -bSI + omega * R5
        k5_E = bSI - sigma * E5
        k5_I = sigma * E5 - gamma * I5
        k5_R = gamma * I5 - omega * R5

        # Stage 6
        S6 = S + h * (9017.0/3168.0*k1_S + (-355.0/33.0)*k2_S + 46732.0/5247.0*k3_S + 49.0/176.0*k4_S + (-5103.0/18656.0)*k5_S)
        E6 = E + h * (9017.0/3168.0*k1_E + (-355.0/33.0)*k2_E + 46732.0/5247.0*k3_E + 49.0/176.0*k4_E + (-5103.0/18656.0)*k5_E)
        I6 = I + h * (9017.0/3168.0*k1_I + (-355.0/33.0)*k2_I + 46732.0/5247.0*k3_I + 49.0/176.0*k4_I + (-5103.0/18656.0)*k5_I)
        R6 = R + h * (9017.0/3168.0*k1_R + (-355.0/33.0)*k2_R + 46732.0/5247.0*k3_R + 49.0/176.0*k4_R + (-5103.0/18656.0)*k5_R)
        bSI = beta * S6 * I6
        k6_S = -bSI + omega * R6
        k6_E = bSI - sigma * E6
        k6_I = sigma * E6 - gamma * I6
        k6_R = gamma * I6 - omega * R6

        # Stage 7
        S7 = S + h * (35.0/384.0*k1_S + 500.0/1113.0*k3_S + 125.0/192.0*k4_S + (-2187.0/6784.0)*k5_S + 11.0/84.0*k6_S)
        E7 = E + h * (35.0/384.0*k1_E + 500.0/1113.0*k3_E + 125.0/192.0*k4_E + (-2187.0/6784.0)*k5_E + 11.0/84.0*k6_E)
        I7 = I + h * (35.0/384.0*k1_I + 500.0/1113.0*k3_I + 125.0/192.0*k4_I + (-2187.0/6784.0)*k5_I + 11.0/84.0*k6_I)
        R7 = R + h * (35.0/384.0*k1_R + 500.0/1113.0*k3_R + 125.0/192.0*k4_R + (-2187.0/6784.0)*k5_R + 11.0/84.0*k6_R)
        bSI = beta * S7 * I7
        k7_S = -bSI + omega * R7
        k7_E = bSI - sigma * E7
        k7_I = sigma * E7 - gamma * I7
        k7_R = gamma * I7 - omega * R7

        # 5th order solution
        S_new = S + h * (35.0/384.0*k1_S + 500.0/1113.0*k3_S + 125.0/192.0*k4_S + (-2187.0/6784.0)*k5_S + 11.0/84.0*k6_S)
        E_new = E + h * (35.0/384.0*k1_E + 500.0/1113.0*k3_E + 125.0/192.0*k4_E + (-2187.0/6784.0)*k5_E + 11.0/84.0*k6_E)
        I_new = I + h * (35.0/384.0*k1_I + 500.0/1113.0*k3_I + 125.0/192.0*k4_I + (-2187.0/6784.0)*k5_I + 11.0/84.0*k6_I)
        R_new = R + h * (35.0/384.0*k1_R + 500.0/1113.0*k3_R + 125.0/192.0*k4_R + (-2187.0/6784.0)*k5_R + 11.0/84.0*k6_R)

        # Error estimate
        err_S = h * ((35.0/384.0 - 5179.0/57600.0)*k1_S + (500.0/1113.0 - 7571.0/16695.0)*k3_S + (125.0/192.0 - 393.0/640.0)*k4_S + (-2187.0/6784.0 + 92097.0/339200.0)*k5_S + (11.0/84.0 - 187.0/2100.0)*k6_S + (-1.0/40.0)*k7_S)
        err_E = h * ((35.0/384.0 - 5179.0/57600.0)*k1_E + (500.0/1113.0 - 7571.0/16695.0)*k3_E + (125.0/192.0 - 393.0/640.0)*k4_E + (-2187.0/6784.0 + 92097.0/339200.0)*k5_E + (11.0/84.0 - 187.0/2100.0)*k6_E + (-1.0/40.0)*k7_E)
        err_I = h * ((35.0/384.0 - 5179.0/57600.0)*k1_I + (500.0/1113.0 - 7571.0/16695.0)*k3_I + (125.0/192.0 - 393.0/640.0)*k4_I + (-2187.0/6784.0 + 92097.0/339200.0)*k5_I + (11.0/84.0 - 187.0/2100.0)*k6_I + (-1.0/40.0)*k7_I)
        err_R = h * ((35.0/384.0 - 5179.0/57600.0)*k1_R + (500.0/1113.0 - 7571.0/16695.0)*k3_R + (125.0/192.0 - 393.0/640.0)*k4_R + (-2187.0/6784.0 + 92097.0/339200.0)*k5_R + (11.0/84.0 - 187.0/2100.0)*k6_R + (-1.0/40.0)*k7_R)

        scale_S = atol + rtol * max(abs(S_new), abs(S))
        scale_E = atol + rtol * max(abs(E_new), abs(E))
        scale_I = atol + rtol * max(abs(I_new), abs(I))
        scale_R = atol + rtol * max(abs(R_new), abs(R))

        err_norm = (err_S/scale_S)**2 + (err_E/scale_E)**2 + (err_I/scale_I)**2 + (err_R/scale_R)**2
        err_norm = np.sqrt(err_norm / 4.0)

        if err_norm <= 1.0:
            t += h
            S, E, I, R = S_new, E_new, I_new, R_new
            k1_S, k1_E, k1_I, k1_R = k7_S, k7_E, k7_I, k7_R

        if err_norm == 0.0:
            factor = 5.0
        else:
            factor = 0.9 * err_norm**(-0.2)
        if factor > 5.0:
            factor = 5.0
        elif factor < 0.2:
            factor = 0.2
        h *= factor

    return S, E, I, R


class Solver:
    def __init__(self):
        # Warmup the JIT to avoid compilation cost during timing
        rk45_scalar(0.0, 1.0, 0.99, 0.01, 0.0, 0.0, 0.35, 0.2, 0.1, 0.002, 1e-6, 1e-9)

    def solve(self, problem, **kwargs) -> Any:
        t0 = problem["t0"]
        t1 = problem["t1"]
        y0 = problem["y0"]
        params = problem["params"]
        beta = params["beta"]
        sigma = params["sigma"]
        gamma = params["gamma"]
        omega = params["omega"]

        S, E, I, R = rk45_scalar(float(t0), float(t1),
                                  float(y0[0]), float(y0[1]), float(y0[2]), float(y0[3]),
                                  float(beta), float(sigma), float(gamma), float(omega),
                                  1e-6, 1e-9)
        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.4896s
Total Solver Time:   0.0144s
Raw Speedup:         590.7078 x
Final Reward (Score): 590.7078
---------------------------
..

==================================== 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 114.55s (0:01:54) =========================
