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


@njit(cache=True)
def _kernel_full(y):
    """All-in-one numba kernel: sort + find rho + project.

    Used for medium n where the overhead of the sort inside numba is acceptable.
    """
    n = y.shape[0]
    # Sort ascending
    sorted_y = np.sort(y)
    # Reverse in-place to get descending
    for i in range(n // 2):
        j = n - 1 - i
        tmp = sorted_y[i]
        sorted_y[i] = sorted_y[j]
        sorted_y[j] = tmp

    # Combined cumulative sum + rho-finding with early termination.
    # cumsum holds cumsum_y[k] = sum_{i<=k} sorted_y[i] - 1.
    cumsum = -1.0
    rho = -1
    saved_cumsum = -1.0
    for k in range(n):
        cumsum += sorted_y[k]
        # Condition: sorted_y[k] * (k+1) > cumsum_y[k] <=> sorted_y[k] > cumsum_y[k] / (k+1)
        if sorted_y[k] * (k + 1) > cumsum:
            rho = k
            saved_cumsum = cumsum
        else:
            break

    theta = saved_cumsum / (rho + 1)

    result = np.empty(n, dtype=np.float64)
    for i in range(n):
        val = y[i] - theta
        if val > 0.0:
            result[i] = val
        else:
            result[i] = 0.0

    return result


@njit(cache=True)
def _kernel_post(sorted_y, y_orig, n):
    """Post-processing kernel for already sorted (descending) y.

    Used for large n where numpy's sort is faster than numba's.
    """
    cumsum = -1.0
    rho = -1
    saved_cumsum = -1.0
    for k in range(n):
        cumsum += sorted_y[k]
        if sorted_y[k] * (k + 1) > cumsum:
            rho = k
            saved_cumsum = cumsum
        else:
            break

    theta = saved_cumsum / (rho + 1)

    result = np.empty(n, dtype=np.float64)
    for i in range(n):
        val = y_orig[i] - theta
        if val > 0.0:
            result[i] = val
        else:
            result[i] = 0.0

    return result


def _solve_small(y_list):
    """Pure Python path for very small n to avoid numpy/numba overhead."""
    n = len(y_list)
    sorted_y = sorted(y_list, reverse=True)
    cumsum = -1.0
    rho = -1
    saved_cumsum = -1.0
    for k in range(n):
        cumsum += sorted_y[k]
        if sorted_y[k] * (k + 1) > cumsum:
            rho = k
            saved_cumsum = cumsum
        else:
            break
    theta = saved_cumsum / (rho + 1)
    return np.array([max(yi - theta, 0.0) for yi in y_list], dtype=np.float64)


class Solver:
    def __init__(self):
        # Warm up numba JIT compilation (this is in __init__, doesn't count toward runtime).
        dummy = np.array([1.0, 2.0, 3.0], dtype=np.float64)
        _ = _kernel_full(dummy)
        _ = _kernel_post(np.sort(dummy)[::-1], dummy, 3)

    def solve(self, problem, **kwargs):
        y_list = problem["y"]
        n = len(y_list)

        # Pure Python for very small n - avoids numpy/numba call overhead.
        if n < 30:
            x = _solve_small(y_list)
        elif n < 1000:
            # Full numba kernel for medium n: sort inside numba is faster than
            # Python overhead of numpy + kernel_post.
            y = np.asarray(y_list, dtype=np.float64)
            x = _kernel_full(y)
        else:
            # For large n, numpy's sort is faster than numba's, so use it directly.
            y = np.asarray(y_list, dtype=np.float64)
            sorted_y = np.sort(y)[::-1]
            x = _kernel_post(sorted_y, y, n)

        return {"solution": x}--- 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: 6.2166s
Total Solver Time:   5.4781s
Raw Speedup:         1.1348 x
Final Reward (Score): 1.1348
---------------------------
..

==================================== 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 295.11s (0:04:55) =========================
