@@ -48,6 +48,42 @@ def pytest_addoption(parser):
4848
4949
5050def pytest_configure (config ):
51+ # Monkey-patch JAX slice/slice_update to normalize mixed-dtype
52+ # start_indices. Keras >= 3.14.1 can pass Python ints (int64) alongside
53+ # JAX int32 arrays to jax.lax.dynamic_slice, which rejects them.
54+ if keras .config .backend () == "jax" :
55+ import jax
56+ import jax .numpy as jnp
57+ from keras .src .backend .jax import core as jax_core
58+
59+ _original_jax_slice = jax_core .slice
60+ _original_jax_slice_update = jax_core .slice_update
61+
62+ def _normalize_start_indices (start_indices ):
63+ arrays = [
64+ jnp .asarray (idx )
65+ if not isinstance (idx , jax .Array )
66+ else idx
67+ for idx in start_indices
68+ ]
69+ dtypes = {a .dtype for a in arrays }
70+ if len (dtypes ) > 1 :
71+ arrays = [jnp .astype (a , jnp .int32 ) for a in arrays ]
72+ return arrays
73+
74+ def _patched_slice (inputs , start_indices , shape ):
75+ start_indices = _normalize_start_indices (start_indices )
76+ return _original_jax_slice (inputs , start_indices , shape )
77+
78+ def _patched_slice_update (inputs , start_indices , updates ):
79+ start_indices = _normalize_start_indices (start_indices )
80+ return _original_jax_slice_update (
81+ inputs , start_indices , updates
82+ )
83+
84+ jax_core .slice = _patched_slice
85+ jax_core .slice_update = _patched_slice_update
86+
5187 # Monkey-patch training methods for OpenVINO backend
5288 if keras .config .backend () == "openvino" :
5389 keras .Model .fit = lambda * args , ** kwargs : pytest .skip (
0 commit comments