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


# ============================================================
# DOP853 (8th order) coefficients — exact values from scipy.
# Stage numbering: 1..12 (k1..k12). FSAL stage (k_f = f(y_new))
# is computed separately and does NOT enter the error estimate.
# ============================================================
# Stage combination coefficients (aIJ = coef of kJ in stage I)
a21 = 0.05260015195876773

a31 = 0.0197250569845379
a32 = 0.0591751709536137

a41 = 0.02958758547680685
a43 = 0.08876275643042054

a51 = 0.2413651341592667
a53 = -0.8845494793282861
a54 = 0.924834003261792

a61 = 0.037037037037037035
a64 = 0.17082860872947386
a65 = 0.12546768756682242

a71 = 0.037109375
a74 = 0.17025221101954405
a75 = 0.06021653898045596
a76 = -0.017578125

a81 = 0.03709200011850479
a84 = 0.17038392571223998
a85 = 0.10726203044637328
a86 = -0.015319437748624402
a87 = 0.008273789163814023

a91 = 0.6241109587160757
a94 = -3.3608926294469414
a95 = -0.868219346841726
a96 = 27.59209969944671
a97 = 20.154067550477894
a98 = -43.48988418106996

a10_1 = 0.47766253643826434
a10_4 = -2.4881146199716677
a10_5 = -0.590290826836843
a10_6 = 21.230051448181193
a10_7 = 15.279233632882423
a10_8 = -33.28821096898486
a10_9 = -0.020331201708508627

a11_1 = -0.9371424300859873
a11_4 = 5.186372428844064
a11_5 = 1.0914373489967295
a11_6 = -8.149787010746927
a11_7 = -18.52006565999696
a11_8 = 22.739487099350505
a11_9 = 2.4936055526796523
a11_10 = -3.0467644718982196

a12_1 = 2.273310147516538
a12_4 = -10.53449546673725
a12_5 = -2.0008720582248625
a12_6 = -17.9589318631188
a12_7 = 27.94888452941996
a12_8 = -2.8589982771350235
a12_9 = -8.87285693353063
a12_10 = 12.360567175794303
a12_11 = 0.6433927460157636

# Solution weights B (nonzero at stages 1, 6, 7, 8, 9, 10, 11, 12)
B_k1 = 0.054293734116568765
B_k6 = 4.450312892752409
B_k7 = 1.8915178993145003
B_k8 = -5.801203960010585
B_k9 = 0.3111643669578199
B_k10 = -0.1521609496625161
B_k11 = 0.20136540080403034
B_k12 = 0.04471061572777259

# Error coefficients E5 (nonzero at stages 1, 6, 7, 8, 9, 10, 11, 12)
E5_k1 = 0.013120044994195
E5_k6 = -1.225156446376204
E5_k7 = -0.49575894965725
E5_k8 = 1.664377182454986
E5_k9 = -0.350328848749974
E5_k10 = 0.334179118713017
E5_k11 = 0.081923206485116
E5_k12 = -0.022355307863886

# Error coefficients E3
E3_k1 = -0.189800754072408
E3_k6 = 4.450312892752409
E3_k7 = 1.8915178993145
E3_k8 = -5.801203960010585
E3_k9 = -0.422682321323792
E3_k10 = -0.152160949662516
E3_k11 = 0.20136540080403
E3_k12 = 0.022651792198361


@njit(cache=True, fastmath=True)
def _derivs(S, E, I, R, beta, sigma, gamma, omega):
    bsi = beta * S * I
    dS = -bsi + omega * R
    dE = bsi - sigma * E
    dI = sigma * E - gamma * I
    dR = gamma * I - omega * R
    return dS, dE, dI, dR


@njit(cache=True, fastmath=True)
def _integrate(t0, t1, y0, beta, sigma, gamma, omega, rtol, atol):
    if t1 == t0:
        return y0[0], y0[1], y0[2], y0[3]

    direction = 1.0 if t1 > t0 else -1.0
    span = abs(t1 - t0)

    S = y0[0]
    E = y0[1]
    I = y0[2]
    R = y0[3]
    t = t0

    k1S, k1E, k1I, k1R = _derivs(S, E, I, R, beta, sigma, gamma, omega)

    scS = atol + rtol * abs(S)
    scE = atol + rtol * abs(E)
    scI = atol + rtol * abs(I)
    scR = atol + rtol * abs(R)
    d0 = (scS * scS + scE * scE + scI * scI + scR * scR) ** 0.5
    d1 = (k1S * k1S + k1E * k1E + k1I * k1I + k1R * k1R) ** 0.5
    if d0 < 1e-5 or d1 < 1e-5:
        h_abs = 1e-6
    else:
        h_abs = 0.01 * d0 / d1
    if h_abs > span:
        h_abs = span
    if h_abs < 1e-12:
        h_abs = 1e-12

    min_step = 1e-13
    error_exponent = -1.0 / 8.0

    while (t - t1) * direction < 0.0:
        if (t + direction * h_abs - t1) * direction >= 0.0:
            h_abs = abs(t1 - t)
        h = direction * h_abs

        # Stage 2
        S2 = S + h * (a21 * k1S)
        E2 = E + h * (a21 * k1E)
        I2 = I + h * (a21 * k1I)
        R2 = R + h * (a21 * k1R)
        k2S, k2E, k2I, k2R = _derivs(S2, E2, I2, R2, beta, sigma, gamma, omega)

        # Stage 3
        S3 = S + h * (a31 * k1S + a32 * k2S)
        E3 = E + h * (a31 * k1E + a32 * k2E)
        I3 = I + h * (a31 * k1I + a32 * k2I)
        R3 = R + h * (a31 * k1R + a32 * k2R)
        k3S, k3E, k3I, k3R = _derivs(S3, E3, I3, R3, beta, sigma, gamma, omega)

        # Stage 4
        S4 = S + h * (a41 * k1S + a43 * k3S)
        E4 = E + h * (a41 * k1E + a43 * k3E)
        I4 = I + h * (a41 * k1I + a43 * k3I)
        R4 = R + h * (a41 * k1R + a43 * k3R)
        k4S, k4E, k4I, k4R = _derivs(S4, E4, I4, R4, beta, sigma, gamma, omega)

        # Stage 5
        S5 = S + h * (a51 * k1S + a53 * k3S + a54 * k4S)
        E5 = E + h * (a51 * k1E + a53 * k3E + a54 * k4E)
        I5 = I + h * (a51 * k1I + a53 * k3I + a54 * k4I)
        R5 = R + h * (a51 * k1R + a53 * k3R + a54 * k4R)
        k5S, k5E, k5I, k5R = _derivs(S5, E5, I5, R5, beta, sigma, gamma, omega)

        # Stage 6
        S6 = S + h * (a61 * k1S + a64 * k4S + a65 * k5S)
        E6 = E + h * (a61 * k1E + a64 * k4E + a65 * k5E)
        I6 = I + h * (a61 * k1I + a64 * k4I + a65 * k5I)
        R6 = R + h * (a61 * k1R + a64 * k4R + a65 * k5R)
        k6S, k6E, k6I, k6R = _derivs(S6, E6, I6, R6, beta, sigma, gamma, omega)

        # Stage 7
        S7 = S + h * (a71 * k1S + a74 * k4S + a75 * k5S + a76 * k6S)
        E7 = E + h * (a71 * k1E + a74 * k4E + a75 * k5E + a76 * k6E)
        I7 = I + h * (a71 * k1I + a74 * k4I + a75 * k5I + a76 * k6I)
        R7 = R + h * (a71 * k1R + a74 * k4R + a75 * k5R + a76 * k6R)
        k7S, k7E, k7I, k7R = _derivs(S7, E7, I7, R7, beta, sigma, gamma, omega)

        # Stage 8
        S8 = S + h * (a81 * k1S + a84 * k4S + a85 * k5S + a86 * k6S + a87 * k7S)
        E8 = E + h * (a81 * k1E + a84 * k4E + a85 * k5E + a86 * k6E + a87 * k7E)
        I8 = I + h * (a81 * k1I + a84 * k4I + a85 * k5I + a86 * k6I + a87 * k7I)
        R8 = R + h * (a81 * k1R + a84 * k4R + a85 * k5R + a86 * k6R + a87 * k7R)
        k8S, k8E, k8I, k8R = _derivs(S8, E8, I8, R8, beta, sigma, gamma, omega)

        # Stage 9
        S9 = S + h * (a91 * k1S + a94 * k4S + a95 * k5S + a96 * k6S + a97 * k7S + a98 * k8S)
        E9 = E + h * (a91 * k1E + a94 * k4E + a95 * k5E + a96 * k6E + a97 * k7E + a98 * k8E)
        I9 = I + h * (a91 * k1I + a94 * k4I + a95 * k5I + a96 * k6I + a97 * k7I + a98 * k8I)
        R9 = R + h * (a91 * k1R + a94 * k4R + a95 * k5R + a96 * k6R + a97 * k7R + a98 * k8R)
        k9S, k9E, k9I, k9R = _derivs(S9, E9, I9, R9, beta, sigma, gamma, omega)

        # Stage 10
        S10 = S + h * (a10_1 * k1S + a10_4 * k4S + a10_5 * k5S + a10_6 * k6S + a10_7 * k7S + a10_8 * k8S + a10_9 * k9S)
        E10 = E + h * (a10_1 * k1E + a10_4 * k4E + a10_5 * k5E + a10_6 * k6E + a10_7 * k7E + a10_8 * k8E + a10_9 * k9E)
        I10 = I + h * (a10_1 * k1I + a10_4 * k4I + a10_5 * k5I + a10_6 * k6I + a10_7 * k7I + a10_8 * k8I + a10_9 * k9I)
        R10 = R + h * (a10_1 * k1R + a10_4 * k4R + a10_5 * k5R + a10_6 * k6R + a10_7 * k7R + a10_8 * k8R + a10_9 * k9R)
        k10S, k10E, k10I, k10R = _derivs(S10, E10, I10, R10, beta, sigma, gamma, omega)

        # Stage 11
        S11 = S + h * (a11_1 * k1S + a11_4 * k4S + a11_5 * k5S + a11_6 * k6S + a11_7 * k7S + a11_8 * k8S + a11_9 * k9S + a11_10 * k10S)
        E11 = E + h * (a11_1 * k1E + a11_4 * k4E + a11_5 * k5E + a11_6 * k6E + a11_7 * k7E + a11_8 * k8E + a11_9 * k9E + a11_10 * k10E)
        I11 = I + h * (a11_1 * k1I + a11_4 * k4I + a11_5 * k5I + a11_6 * k6I + a11_7 * k7I + a11_8 * k8I + a11_9 * k9I + a11_10 * k10I)
        R11 = R + h * (a11_1 * k1R + a11_4 * k4R + a11_5 * k5R + a11_6 * k6R + a11_7 * k7R + a11_8 * k8R + a11_9 * k9R + a11_10 * k10R)
        k11S, k11E, k11I, k11R = _derivs(S11, E11, I11, R11, beta, sigma, gamma, omega)

        # Stage 12
        S12 = S + h * (a12_1 * k1S + a12_4 * k4S + a12_5 * k5S + a12_6 * k6S + a12_7 * k7S + a12_8 * k8S + a12_9 * k9S + a12_10 * k10S + a12_11 * k11S)
        E12 = E + h * (a12_1 * k1E + a12_4 * k4E + a12_5 * k5E + a12_6 * k6E + a12_7 * k7E + a12_8 * k8E + a12_9 * k9E + a12_10 * k10E + a12_11 * k11E)
        I12 = I + h * (a12_1 * k1I + a12_4 * k4I + a12_5 * k5I + a12_6 * k6I + a12_7 * k7I + a12_8 * k8I + a12_9 * k9I + a12_10 * k10I + a12_11 * k11I)
        R12 = R + h * (a12_1 * k1R + a12_4 * k4R + a12_5 * k5R + a12_6 * k6R + a12_7 * k7R + a12_8 * k8R + a12_9 * k9R + a12_10 * k10R + a12_11 * k11R)
        k12S, k12E, k12I, k12R = _derivs(S12, E12, I12, R12, beta, sigma, gamma, omega)

        # New solution
        Sn = S + h * (B_k1 * k1S + B_k6 * k6S + B_k7 * k7S + B_k8 * k8S + B_k9 * k9S + B_k10 * k10S + B_k11 * k11S + B_k12 * k12S)
        En = E + h * (B_k1 * k1E + B_k6 * k6E + B_k7 * k7E + B_k8 * k8E + B_k9 * k9E + B_k10 * k10E + B_k11 * k11E + B_k12 * k12E)
        In = I + h * (B_k1 * k1I + B_k6 * k6I + B_k7 * k7I + B_k8 * k8I + B_k9 * k9I + B_k10 * k10I + B_k11 * k11I + B_k12 * k12I)
        Rn = R + h * (B_k1 * k1R + B_k6 * k6R + B_k7 * k7R + B_k8 * k8R + B_k9 * k9R + B_k10 * k10R + B_k11 * k11R + B_k12 * k12R)

        # Scales
        scS = atol + rtol * max(abs(S), abs(Sn))
        scE = atol + rtol * max(abs(E), abs(En))
        scI = atol + rtol * max(abs(I), abs(In))
        scR = atol + rtol * max(abs(R), abs(Rn))

        # Error estimate (E5 and E3, using stages 1, 6, 7, 8, 9, 10, 11, 12)
        e5S = (E5_k1 * k1S + E5_k6 * k6S + E5_k7 * k7S + E5_k8 * k8S + E5_k9 * k9S + E5_k10 * k10S + E5_k11 * k11S + E5_k12 * k12S) / scS
        e5E = (E5_k1 * k1E + E5_k6 * k6E + E5_k7 * k7E + E5_k8 * k8E + E5_k9 * k9E + E5_k10 * k10E + E5_k11 * k11E + E5_k12 * k12E) / scE
        e5I = (E5_k1 * k1I + E5_k6 * k6I + E5_k7 * k7I + E5_k8 * k8I + E5_k9 * k9I + E5_k10 * k10I + E5_k11 * k11I + E5_k12 * k12I) / scI
        e5R = (E5_k1 * k1R + E5_k6 * k6R + E5_k7 * k7R + E5_k8 * k8R + E5_k9 * k9R + E5_k10 * k10R + E5_k11 * k11R + E5_k12 * k12R) / scR

        e3S = (E3_k1 * k1S + E3_k6 * k6S + E3_k7 * k7S + E3_k8 * k8S + E3_k9 * k9S + E3_k10 * k10S + E3_k11 * k11S + E3_k12 * k12S) / scS
        e3E = (E3_k1 * k1E + E3_k6 * k6E + E3_k7 * k7E + E3_k8 * k8E + E3_k9 * k9E + E3_k10 * k10E + E3_k11 * k11E + E3_k12 * k12E) / scE
        e3I = (E3_k1 * k1I + E3_k6 * k6I + E3_k7 * k7I + E3_k8 * k8I + E3_k9 * k9I + E3_k10 * k10I + E3_k11 * k11I + E3_k12 * k12I) / scI
        e3R = (E3_k1 * k1R + E3_k6 * k6R + E3_k7 * k7R + E3_k8 * k8R + E3_k9 * k9R + E3_k10 * k10R + E3_k11 * k11R + E3_k12 * k12R) / scR

        e5_norm_2 = e5S * e5S + e5E * e5E + e5I * e5I + e5R * e5R
        e3_norm_2 = e3S * e3S + e3E * e3E + e3I * e3I + e3R * e3R

        if e5_norm_2 == 0.0 and e3_norm_2 == 0.0:
            err_norm = 0.0
        else:
            denom = e5_norm_2 + 0.01 * e3_norm_2
            err_norm = abs(h) * e5_norm_2 / ((denom * 4.0) ** 0.5)

        if err_norm < 1.0:
            # FSAL: next k1 = f(y_new)
            kfS, kfE, kfI, kfR = _derivs(Sn, En, In, Rn, beta, sigma, gamma, omega)
            t = t + h
            S = Sn; E = En; I = In; R = Rn
            k1S = kfS; k1E = kfE; k1I = kfI; k1R = kfR

            if err_norm == 0.0:
                factor = 10.0
            else:
                factor = 0.9 * err_norm ** error_exponent
                if factor > 10.0:
                    factor = 10.0
                if factor < 0.2:
                    factor = 0.2
            new_h = h_abs * factor
            if new_h > span:
                new_h = span
            h_abs = new_h
        else:
            factor = 0.9 * err_norm ** error_exponent
            if factor < 0.2:
                factor = 0.2
            new_h = h_abs * factor
            if new_h < min_step:
                new_h = min_step
                kfS, kfE, kfI, kfR = _derivs(Sn, En, In, Rn, beta, sigma, gamma, omega)
                t = t + h
                S = Sn; E = En; I = In; R = Rn
                k1S = kfS; k1E = kfE; k1I = kfI; k1R = kfR
            h_abs = new_h

    return S, E, I, R


class Solver:
    def __init__(self):
        y0 = np.array([0.89, 0.01, 0.005, 0.095], dtype=float)
        _integrate(0.0, 1.0, y0, 0.35, 0.2, 0.1, 0.002, 1e-7, 1e-9)
        self._rtol = 1e-7
        self._atol = 1e-9

    def solve(self, problem, **kwargs):
        t0 = float(problem["t0"])
        t1 = float(problem["t1"])
        y0 = np.asarray(problem["y0"], dtype=float)
        p = problem["params"]
        beta = float(p["beta"])
        sigma = float(p["sigma"])
        gamma = float(p["gamma"])
        omega = float(p["omega"])
        S, E, I, R = _integrate(t0, t1, y0, beta, sigma, gamma, omega,
                                self._rtol, self._atol)
        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.3623s
Total Solver Time:   0.0126s
Raw Speedup:         664.2217 x
Final Reward (Score): 664.2217
---------------------------
..

==================================== 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 108.95s (0:01:48) =========================
