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


@njit(cache=True, fastmath=True)
def _solve_rk4(y0, t1, beta, sigma, gamma, omega, n_steps):
    h = t1 / n_steps
    S, E, I, R = y0[0], y0[1], y0[2], y0[3]
    for _ in range(n_steps):
        bSI = beta * S * I
        k1S = -bSI + omega * R
        k1E = bSI - sigma * E
        k1I = sigma * E - gamma * I
        k1R = gamma * I - omega * R

        h2 = 0.5 * h
        S2 = S + h2 * k1S
        E2 = E + h2 * k1E
        I2 = I + h2 * k1I
        R2 = R + h2 * k1R
        bSI = beta * S2 * I2
        k2S = -bSI + omega * R2
        k2E = bSI - sigma * E2
        k2I = sigma * E2 - gamma * I2
        k2R = gamma * I2 - omega * R2

        S3 = S + h2 * k2S
        E3 = E + h2 * k2E
        I3 = I + h2 * k2I
        R3 = R + h2 * k2R
        bSI = beta * S3 * I3
        k3S = -bSI + omega * R3
        k3E = bSI - sigma * E3
        k3I = sigma * E3 - gamma * I3
        k3R = gamma * I3 - omega * R3

        S4 = S + h * k3S
        E4 = E + h * k3E
        I4 = I + h * k3I
        R4 = R + h * k3R
        bSI = beta * S4 * I4
        k4S = -bSI + omega * R4
        k4E = bSI - sigma * E4
        k4I = sigma * E4 - gamma * I4
        k4R = gamma * I4 - omega * R4

        h6 = h / 6.0
        S += h6 * (k1S + 2.0 * k2S + 2.0 * k3S + k4S)
        E += h6 * (k1E + 2.0 * k2E + 2.0 * k3E + k4E)
        I += h6 * (k1I + 2.0 * k2I + 2.0 * k3I + k4I)
        R += h6 * (k1R + 2.0 * k2R + 2.0 * k3R + k4R)
    return S, E, I, R


class Solver:
    def __init__(self):
        # Warm up Numba JIT compilation
        y0 = np.array([0.89, 0.01, 0.005, 0.095])
        _solve_rk4(y0, 1.0, 0.35, 0.2, 0.1, 0.002, 10)

    def solve(self, problem, **kwargs):
        y0 = np.asarray(problem["y0"])
        p = problem["params"]
        T = problem["t1"] - problem["t0"]
        # Adaptive step count based on max rate
        mr = p["beta"]
        s = p["sigma"]
        g = p["gamma"]
        o = p["omega"]
        if s > mr: mr = s
        if g > mr: mr = g
        if o > mr: mr = o
        # h * max_rate ~ 1 gives accuracy ~1e-5
        n_steps = max(50, int(T * (mr * 1.0 if mr > 0.1 else 0.05)))

        S, E, I, R = _solve_rk4(y0, T, p["beta"], p["sigma"], p["gamma"], p["omega"], n_steps)

        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.4009s
Total Solver Time:   0.0081s
Raw Speedup:         1037.8561 x
Final Reward (Score): 1037.8561
---------------------------
..

==================================== 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 110.82s (0:01:50) =========================
