{"record":{"id":"4a4560e2d835a6f8","repo":"keras-team/keras","slug":"argument-synchronized-true-is-not-supported-with-j-4a4560","errorCode":null,"errorMessage":"Argument synchronized=True is not supported with JAX.","messagePattern":"Argument synchronized=True is not supported with JAX\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"keras/src/backend/jax/ops/nn.py","lineNumber":1044,"sourceCode":"            \"Arguments `target` and `output` must have the same shape. \"\n            \"Received: \"\n            f\"target.shape={target.shape}, output.shape={output.shape}\"\n        )\n\n    if from_logits:\n        log_logits = jax.nn.log_sigmoid(output)\n        log_neg_logits = jax.nn.log_sigmoid(-output)\n        return -1.0 * target * log_logits - (1.0 - target) * log_neg_logits\n\n    output = jnp.clip(output, backend.epsilon(), 1.0 - backend.epsilon())\n    bce = target * jnp.log(output)\n    bce += (1.0 - target) * jnp.log(1.0 - output)\n    return -bce\n\n\ndef moments(x, axes, keepdims=False, synchronized=False):\n    if synchronized:\n        raise NotImplementedError(\n            \"Argument synchronized=True is not supported with JAX.\"\n        )\n    # The dynamic range of float16 is too limited for statistics. As a\n    # workaround, we simply perform the operations on float32 and convert back\n    # to float16\n    need_cast = False\n    ori_dtype = backend.standardize_dtype(x.dtype)\n    if ori_dtype in (\"float16\", \"bfloat16\"):\n        need_cast = True\n        x = cast(x, \"float32\")\n\n    mean = jnp.mean(x, axes, keepdims=True)\n    variance = jnp.var(x, axis=axes, keepdims=True)\n\n    if not keepdims:\n        mean = jnp.squeeze(mean, axes)\n        variance = jnp.squeeze(variance, axes)\n    if need_cast:","sourceCodeStart":1026,"sourceCodeEnd":1062,"githubUrl":"https://github.com/keras-team/keras/blob/7a34a03db60bf60042242d6a556fc3be119046a5/keras/src/backend/jax/ops/nn.py#L1026-L1062","documentation":"Error \"Argument synchronized=True is not supported with JAX.\" thrown in keras-team/keras.","triggerScenarios":"Thrown at keras/src/backend/jax/ops/nn.py:1044 when the library encounters an invalid state.","commonSituations":"See trigger scenarios.","solutions":[],"exampleFix":null,"handlingStrategy":null,"validationCode":null,"typeGuard":null,"tryCatchPattern":null,"preventionTips":[],"tags":[],"backgroundTag":null,"analyzedSha":"7a34a03db60bf60042242d6a556fc3be119046a5","analyzedAt":"2026-08-25T21:25:25.994Z","schemaVersion":2},"datasetVersion":"2026-08-26T02:17:13.382Z"}