From e7dadccf69940e293d1c63f4135f25448e28244e Mon Sep 17 00:00:00 2001 From: Tom George Date: Tue, 25 Aug 2026 21:47:55 -0400 Subject: [PATCH] jit compilation issue --- src/simpl/simpl.py | 2 +- src/simpl/utils.py | 16 ++++++++++++++++ 2 files changed, 17 insertions(+), 1 deletion(-) diff --git a/src/simpl/simpl.py b/src/simpl/simpl.py index ecb7cb6..3466eaf 100755 --- a/src/simpl/simpl.py +++ b/src/simpl/simpl.py @@ -1893,7 +1893,7 @@ def _per_trial_initial_states(mode_l, trial_slices, is_1D_angular=False): modes = mode_np[trial_slice] if is_1D_angular: angles = modes[:, 0] - mean_angle, variance = utils._circular_mean_and_variance(angles=angles, weights=None) + mean_angle, variance = utils._circular_mean_and_variance_numpy(angles) mu = np.asarray(mean_angle)[None] sigma = np.asarray(variance)[None, None] else: diff --git a/src/simpl/utils.py b/src/simpl/utils.py index e758653..589256c 100644 --- a/src/simpl/utils.py +++ b/src/simpl/utils.py @@ -315,6 +315,22 @@ def _circular_mean_and_variance( return mean, variance +def _circular_mean_and_variance_numpy(angles: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + """Compute unweighted circular moments on the host without JAX compilation.""" + angles = np.asarray(angles) + sin_mean = np.mean(np.sin(angles), axis=-1) + cos_mean = np.mean(np.cos(angles), axis=-1) + mean = (np.arctan2(sin_mean, cos_mean) + np.pi) % (2 * np.pi) - np.pi + + resultant_squared = sin_mean**2 + cos_mean**2 + fallback = np.take(angles, 0, axis=-1) + mean = np.where(resultant_squared > 1e-12, mean, (fallback + np.pi) % (2 * np.pi) - np.pi) + + residuals = (angles - np.expand_dims(mean, axis=-1) + np.pi) % (2 * np.pi) - np.pi + variance = np.mean(residuals**2, axis=-1) + return mean, variance + + def _bin_indices_minuspi_pi(theta: jax.Array, n_bins: int) -> jax.Array: """Map theta in radians to integer bin indices [0, n_bins).