keras-team/keras · error · NotImplementedError
Argument synchronized=True is not supported with JAX.
Error message
Argument synchronized=True is not supported with JAX.
What it means
Error "Argument synchronized=True is not supported with JAX." thrown in keras-team/keras.
Source
Thrown at keras/src/backend/jax/ops/nn.py:1044
"Arguments `target` and `output` must have the same shape. "
"Received: "
f"target.shape={target.shape}, output.shape={output.shape}"
)
if from_logits:
log_logits = jax.nn.log_sigmoid(output)
log_neg_logits = jax.nn.log_sigmoid(-output)
return -1.0 * target * log_logits - (1.0 - target) * log_neg_logits
output = jnp.clip(output, backend.epsilon(), 1.0 - backend.epsilon())
bce = target * jnp.log(output)
bce += (1.0 - target) * jnp.log(1.0 - output)
return -bce
def moments(x, axes, keepdims=False, synchronized=False):
if synchronized:
raise NotImplementedError(
"Argument synchronized=True is not supported with JAX."
)
# The dynamic range of float16 is too limited for statistics. As a
# workaround, we simply perform the operations on float32 and convert back
# to float16
need_cast = False
ori_dtype = backend.standardize_dtype(x.dtype)
if ori_dtype in ("float16", "bfloat16"):
need_cast = True
x = cast(x, "float32")
mean = jnp.mean(x, axes, keepdims=True)
variance = jnp.var(x, axis=axes, keepdims=True)
if not keepdims:
mean = jnp.squeeze(mean, axes)
variance = jnp.squeeze(variance, axes)
if need_cast:View on GitHub (pinned to 7a34a03db6)
When it happens
Trigger: Thrown at keras/src/backend/jax/ops/nn.py:1044 when the library encounters an invalid state.
Common situations: See trigger scenarios.
AI-assisted analysis of keras-team/keras@7a34a03db6 (2026-08-25).
Data as JSON: /api/errors/4a4560e2d835a6f8.
Report an issue: GitHub.