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

@numba.njit(cache=True)
def triangular_sqrt_numba(T, S, n):
    """Compute S such that S^2 = T, where T is upper triangular."""
    for i in range(n):
        t_ii = T[i, i]
        sqrt_val = np.sqrt(t_ii)
        if sqrt_val.real < 0:
            sqrt_val = -sqrt_val
        S[i, i] = sqrt_val
    for j in range(1, n):
        for i in range(j-1, -1, -1):
            sum_val = 0.0 + 0.0j
            for k in range(i+1, j):
                sum_val += S[i, k] * S[k, j]
            denom = S[i, i] + S[j, j]
            if denom != 0:
                S[i, j] = (T[i, j] - sum_val) / denom
            else:
                S[i, j] = 0.0

def sqrtm_schur(A):
    """Compute principal matrix square root using Schur decomposition."""
    n = A.shape[0]
    if not np.issubdtype(A.dtype, np.complexfloating):
        A = A.astype(complex)
    # Use scipy.linalg.schur for reliability
    T, Q = scipy.linalg.schur(A, output='complex')
    S = np.zeros_like(T)
    triangular_sqrt_numba(T, S, n)
    X = Q @ S @ Q.conj().T
    return X

class Solver:
    def __init__(self):
        # Warm up numba JIT compilation for various sizes
        dummy = np.eye(2, dtype=complex)
        S = np.zeros((2, 2), dtype=complex)
        triangular_sqrt_numba(dummy, S, 2)
        for n in [5, 10, 20, 50, 100]:
            dummy_n = np.eye(n, dtype=complex)
            S_n = np.zeros((n, n), dtype=complex)
            triangular_sqrt_numba(dummy_n, S_n, n)
    
    def solve(self, problem: dict[str, Any]) -> dict[str, dict[str, list[list[complex]]]]:
        A = problem["matrix"]
        # Convert to numpy array
        if isinstance(A, list):
            rows = []
            for row in A:
                new_row = []
                for elem in row:
                    if isinstance(elem, str):
                        new_row.append(complex(elem))
                    else:
                        new_row.append(elem)
                rows.append(new_row)
            A = np.array(rows, dtype=complex)
        try:
            X = sqrtm_schur(A)
        except Exception as e:
            return {"sqrtm": {"X": []}}
        return {"sqrtm": {"X": X.tolist()}}

--- 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: 58.3023s
Total Solver Time:   22.7597s
Raw Speedup:         2.5616 x
Final Reward (Score): 2.5616
---------------------------
..

=============================== warnings summary ===============================
test_outputs.py: 1100 warnings
  /tests/evaluator.py:79: DeprecationWarning: The `disp` argument is deprecated and will be removed in SciPy 1.18.0.
    X, _ = scipy.linalg.sqrtm(

-- Docs: https://docs.pytest.org/en/stable/how-to/capture-warnings.html
==================================== 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, 1100 warnings in 1132.01s (0:18:52) =================
