From d7c21c450020c7c50fc1f980737957589a280cd4 Mon Sep 17 00:00:00 2001 From: Praharsh Suryadevara Date: Sat, 19 Sep 2026 10:42:16 -0500 Subject: [PATCH 1/2] remove ensure_compile_time_eval because hlo shows it's already optimized --- diffrax/_solution.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/diffrax/_solution.py b/diffrax/_solution.py index 392c447e..67b9b8cb 100644 --- a/diffrax/_solution.py +++ b/diffrax/_solution.py @@ -74,8 +74,7 @@ def update_result(old_result: RESULTS, new_result: RESULTS) -> RESULTS: error_n | error_n error_n error_o """ out_result = RESULTS.where(is_okay(old_result), new_result, old_result) - with jax.ensure_compile_time_eval(): - pred = is_okay(new_result) & is_event(old_result) + pred = is_okay(new_result) & is_event(old_result) return RESULTS.where(pred, old_result, out_result) From e1e651a8c5d18dba57680e13a7e28362eeb3d69a Mon Sep 17 00:00:00 2001 From: Praharsh Suryadevara Date: Sat, 19 Sep 2026 11:14:34 -0500 Subject: [PATCH 2/2] missed compile time removal --- diffrax/_solution.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/diffrax/_solution.py b/diffrax/_solution.py index 67b9b8cb..8af6f444 100644 --- a/diffrax/_solution.py +++ b/diffrax/_solution.py @@ -1,7 +1,6 @@ import warnings from typing import Any -import jax import optimistix as optx from jaxtyping import Array, Bool, PyTree, Real, Shaped @@ -50,8 +49,7 @@ def discrete_terminating_event_occurred(self): def is_okay(result: RESULTS) -> Bool[Array, ""]: - with jax.ensure_compile_time_eval(): # for the `|` between two `Bool[Array, ""]`. - return is_successful(result) | is_event(result) + return is_successful(result) | is_event(result) def is_successful(result: RESULTS) -> Bool[Array, ""]: