{"record":{"id":"054073cd09243f55","repo":"keras-team/keras","slug":"dot-product-attention-only-supports-4d-inputs-r-054073","errorCode":null,"errorMessage":"`dot_product_attention` only supports 4D inputs. Received: query.shape={query.shape}, key.shape={key.shape}, value.shape={value.shape}.","messagePattern":"`dot_product_attention` only supports 4D inputs\\. Received: query\\.shape=(.+?), key\\.shape=(.+?), value\\.shape=(.+?)\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"keras/src/backend/jax/ops/nn.py","lineNumber":1688,"sourceCode":"        scale: Float. Optional scale that is applied to the attention\n            computation.\n        is_causal: Boolean. Specifying whether causal masking is applied.\n        flash_attention: Boolean. Whether to use flash attention optimization\n            for increased performance. Default to None, which means it will\n            be auto-determined based on the platform, input shapes and\n            compatibility.\n        attn_logits_soft_cap: Float. Optional float to softly cap attention\n            logits to avoid numerical stability issues. Applied as:\n            `logits = logits / (1.0 + abs(logits) / attn_logits_soft_cap)`.\n\n    Returns:\n        JAX Array of shape `[batch, time, heads, depth_v]`.\n    \"\"\"\n    query = convert_to_tensor(query)\n    key = convert_to_tensor(key)\n    value = convert_to_tensor(value)\n    if len(query.shape) != 4 or len(key.shape) != 4 or len(value.shape) != 4:\n        raise ValueError(\n            \"`dot_product_attention` only supports 4D inputs. \"\n            f\"Received: query.shape={query.shape}, key.shape={key.shape}, \"\n            f\"value.shape={value.shape}.\"\n        )\n    compute_dtype = backend.result_type(query.dtype, key.dtype, value.dtype)\n    query = cast(query, compute_dtype)\n    key = cast(key, compute_dtype)\n    value = cast(value, compute_dtype)\n    if bias is not None:\n        bias = convert_to_tensor(bias, dtype=compute_dtype)\n\n    # Check platform\n    platform = jax.devices()[0].platform\n    is_tpu = platform == \"tpu\"\n\n    # Determine flash attention compatibility\n    if flash_attention is None:\n        flash_attention = _can_use_flash_attention(query, key, value, bias)","sourceCodeStart":1670,"sourceCodeEnd":1706,"githubUrl":"https://github.com/keras-team/keras/blob/7a34a03db60bf60042242d6a556fc3be119046a5/keras/src/backend/jax/ops/nn.py#L1670-L1706","documentation":"Error \"`dot_product_attention` only supports 4D inputs. Received: query.shape={query.shape}, key.shape={key.shape}, value.shape={value.shape}.\" thrown in keras-team/keras.","triggerScenarios":"Thrown at keras/src/backend/jax/ops/nn.py:1688 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"}