Skip to content

Commit e7a3523

Browse files
committed
Work around Keras 3.14.1 JAX backend dtype mismatch in slice/slice_update
1 parent 1b0efad commit e7a3523

1 file changed

Lines changed: 36 additions & 0 deletions

File tree

conftest.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,42 @@ def pytest_addoption(parser):
4848

4949

5050
def 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

Comments
 (0)