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

@numba.njit(fastmath=True)
def seirs_deriv(t, S, E, I, R, beta, sigma, gamma, omega):
    dS = -beta * S * I + omega * R
    dE = beta * S * I - sigma * E
    dI = sigma * E - gamma * I
    dR = gamma * I - omega * R
    return dS, dE, dI, dR

@numba.njit(fastmath=True)
def dp5_step(t, S, E, I, R, h, beta, sigma, gamma, omega):
    k1_S, k1_E, k1_I, k1_R = seirs_deriv(t, S, E, I, R, beta, sigma, gamma, omega)
    
    k2_S, k2_E, k2_I, k2_R = seirs_deriv(t + 0.2*h, 
        S + h*0.2*k1_S, E + h*0.2*k1_E, I + h*0.2*k1_I, R + h*0.2*k1_R, beta, sigma, gamma, omega)
    
    k3_S, k3_E, k3_I, k3_R = seirs_deriv(t + 0.3*h, 
        S + h*(0.075*k1_S + 0.225*k2_S), E + h*(0.075*k1_E + 0.225*k2_E), I + h*(0.075*k1_I + 0.225*k2_I), R + h*(0.075*k1_R + 0.225*k2_R), beta, sigma, gamma, omega)
    
    k4_S, k4_E, k4_I, k4_R = seirs_deriv(t + 0.8*h, 
        S + h*(44./45.*k1_S - 56./15.*k2_S + 32./9.*k3_S), E + h*(44./45.*k1_E - 56./15.*k2_E + 32./9.*k3_E), I + h*(44./45.*k1_I - 56./15.*k2_I + 32./9.*k3_I), R + h*(44./45.*k1_R - 56./15.*k2_R + 32./9.*k3_R), beta, sigma, gamma, omega)
    
    k5_S, k5_E, k5_I, k5_R = seirs_deriv(t + 8./9.*h, 
        S + h*(19372./6561.*k1_S - 25360./2187.*k2_S + 64448./6561.*k3_S - 212./729.*k4_S), E + h*(19372./6561.*k1_E - 25360./2187.*k2_E + 64448./6561.*k3_E - 212./729.*k4_E), I + h*(19372./6561.*k1_I - 25360./2187.*k2_I + 64448./6561.*k3_I - 212./729.*k4_I), R + h*(19372./6561.*k1_R - 25360./2187.*k2_R + 64448./6561.*k3_R - 212./729.*k4_R), beta, sigma, gamma, omega)
    
    k6_S, k6_E, k6_I, k6_R = seirs_deriv(t + h, 
        S + h*(9017./3168.*k1_S - 355./33.*k2_S + 46732./5247.*k3_S + 49./176.*k4_S - 5103./18656.*k5_S), E + h*(9017./3168.*k1_E - 355./33.*k2_E + 46732./5247.*k3_E + 49./176.*k4_E - 5103./18656.*k5_E), I + h*(9017./3168.*k1_I - 355./33.*k2_I + 46732./5247.*k3_I + 49./176.*k4_I - 5103./18656.*k5_I), R + h*(9017./3168.*k1_R - 355./33.*k2_R + 46732./5247.*k3_R + 49./176.*k4_R - 5103./18656.*k5_R), beta, sigma, gamma, omega)
    
    y_next_S = S + h*(35./384.*k1_S + 500./1113.*k3_S + 125./192.*k4_S - 2187./6784.*k5_S + 11./84.*k6_S)
    y_next_E = E + h*(35./384.*k1_E + 500./1113.*k3_E + 125./192.*k4_E - 2187./6784.*k5_E + 11./84.*k6_E)
    y_next_I = I + h*(35./384.*k1_I + 500./1113.*k3_I + 125./192.*k4_I - 2187./6784.*k5_I + 11./84.*k6_I)
    y_next_R = R + h*(35./384.*k1_R + 500./1113.*k3_R + 125./192.*k4_R - 2187./6784.*k5_R + 11./84.*k6_R)
    
    k7_S, k7_E, k7_I, k7_R = seirs_deriv(t + h, y_next_S, y_next_E, y_next_I, y_next_R, beta, sigma, gamma, omega)
    
    err_S = y_next_S - (S + h*(5179./57600.*k1_S + 7571./16695.*k3_S + 393./640.*k4_S - 92097./339200.*k5_S + 187./2100.*k6_S + 1./40.*k7_S))
    err_E = y_next_E - (E + h*(5179./57600.*k1_E + 7571./16695.*k3_E + 393./640.*k4_E - 92097./339200.*k5_E + 187./2100.*k6_E + 1./40.*k7_E))
    err_I = y_next_I - (I + h*(5179./57600.*k1_I + 7571./16695.*k3_I + 393./640.*k4_I - 92097./339200.*k5_I + 187./2100.*k6_I + 1./40.*k7_I))
    err_R = y_next_R - (R + h*(5179./57600.*k1_R + 7571./16695.*k3_R + 393./640.*k4_R - 92097./339200.*k5_R + 187./2100.*k6_R + 1./40.*k7_R))
    
    return y_next_S, y_next_E, y_next_I, y_next_R, err_S, err_E, err_I, err_R, k7_S, k7_E, k7_I, k7_R

@numba.njit(fastmath=True)
def solve_rk45_numba(t0, t1, y0_0, y0_1, y0_2, y0_3, beta, sigma, gamma, omega, rtol, atol):
    t = t0
    h = 1e-4
    if t1 - t0 < h:
        h = t1 - t0
    h_min = 1e-12
    h_max = t1 - t0
    
    S, E, I, R = y0_0, y0_1, y0_2, y0_3
    
    while t < t1:
        if t + h > t1: 
            h = t1 - t
            
        nS, nE, nI, nR, eS, eE, eI, eR, k7_S, k7_E, k7_I, k7_R = dp5_step(t, S, E, I, R, h, beta, sigma, gamma, omega)
        
        err_norm = 0.0
        
        tol_S = atol + rtol * max(abs(S), abs(nS))
        if tol_S > 0: err_norm = max(err_norm, abs(eS) / tol_S)
        
        tol_E = atol + rtol * max(abs(E), abs(nE))
        if tol_E > 0: err_norm = max(err_norm, abs(eE) / tol_E)
        
        tol_I = atol + rtol * max(abs(I), abs(nI))
        if tol_I > 0: err_norm = max(err_norm, abs(eI) / tol_I)
        
        tol_R = atol + rtol * max(abs(R), abs(nR))
        if tol_R > 0: err_norm = max(err_norm, abs(eR) / tol_R)
            
        if err_norm <= 1.0:
            t += h
            S, E, I, R = nS, nE, nI, nR
            
            max_deriv = max(abs(k7_S), abs(k7_E), abs(k7_I), abs(k7_R))
            if max_deriv < 1e-11:
                break
                
            if err_norm < 1e-5: err_norm = 1e-5
            h_new = h * 0.9 * (err_norm ** -0.2)
            if h_new > 5.0 * h: h_new = 5.0 * h
        else:
            h_new = h * 0.9 * (err_norm ** -0.25)
            if h_new < 0.2 * h: h_new = 0.2 * h
            
        h = h_new
        if h < h_min: h = h_min
        if h > h_max: h = h_max
        
    total = S + E + I + R
    return S/total, E/total, I/total, R/total

class Solver:
    def __init__(self):
        # Trigger JIT compilation so it doesn't count towards evaluation time.
        solve_rk45_numba(0.0, 1.0, 0.89, 0.01, 0.005, 0.095, 0.35, 0.2, 0.1, 0.002, 1e-12, 1e-12)

    def solve(self, problem: dict, **kwargs) -> Any:
        t0 = float(problem['t0'])
        t1 = float(problem['t1'])
        y0 = problem['y0']
        p = problem['params']
        beta = float(p['beta'])
        sigma = float(p['sigma'])
        gamma = float(p['gamma'])
        omega = float(p['omega'])
        
        # Use exact 1e-12 tolerance. Numba JIT makes the loops insanely fast. 
        # Equilibrium detection early stops perfectly safely before the integrator wastes time.
        S, E, I, R = solve_rk45_numba(
            t0, t1, 
            float(y0[0]), float(y0[1]), float(y0[2]), float(y0[3]), 
            beta, sigma, gamma, omega, 
            1e-12, 1e-12
        )
        return [float(S), float(E), float(I), float(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.3381s
Total Solver Time:   0.0202s
Raw Speedup:         412.9805 x
Final Reward (Score): 412.9805
---------------------------
..

==================================== 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 111.23s (0:01:51) =========================
