{"record":{"id":"758fa61dc6133a46","repo":"jax-ml/jax","slug":"unexpected-mask-shape-mask-shape","errorCode":null,"errorMessage":"Unexpected mask shape: {mask.shape}","messagePattern":"Unexpected mask shape: (.+?)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py","lineNumber":2560,"sourceCode":"    )\n\n\ndef _make_splash_attention(\n    mask: np.ndarray | jax.Array | mask_lib.MultiHeadMask,\n    *,\n    block_sizes: BlockSizes | None = None,\n    is_mqa: bool,\n    save_residuals: bool = False,\n    mask_value: float = DEFAULT_MASK_VALUE,\n    attn_logits_soft_cap: float | None = None,\n    downcast_smem_data: bool = True,\n    head_shards: int,\n    q_seq_shards: int,\n    residual_checkpoint_name: str | None = None,\n    interpret: bool = False,\n):\n  if len(mask.shape) != 3:\n    raise ValueError(f'Unexpected mask shape: {mask.shape}')\n\n  if isinstance(mask, np.ndarray):\n    mask = mask_lib.MultiHeadMask(\n        [mask_lib.NumpyMask(head_mask) for head_mask in mask]\n    )\n\n  if block_sizes is None:\n    block_sizes = BlockSizes.get_default()\n\n  process_mask_fn = (\n      mask_info_lib.process_dynamic_mask\n      if isinstance(mask, jax.Array)\n      else mask_info_lib.process_mask\n  )\n\n  process_mask_dvk_fn = (\n      mask_info_lib.process_dynamic_mask_dkv\n      if isinstance(mask, jax.Array)","sourceCodeStart":2542,"sourceCodeEnd":2578,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py#L2542-L2578","documentation":"_make_splash_attention expects a rank-3 mask array shaped (num_heads, q_seq_len, kv_seq_len) (or an equivalent Mask object). A mask of any other rank is rejected before kernel construction.","triggerScenarios":"Passing a 2D (single-head) numpy mask, a 4D batch mask, or a scalar mask to make_splash_attention.","commonSituations":"Migrating code from single-head attention that built (q, kv) masks; passing batched masks shaped (batch, heads, q, kv); forgetting to broadcast a causal mask across heads.","solutions":["Reshape/broadcast the mask to (num_heads, q_seq_len, kv_seq_len)","Wrap a per-head 2D mask: mask_lib.MultiHeadMask([NumpyMask(m) for m in masks])","For a shared causal mask, stack it num_heads times along axis 0"],"exampleFix":"// before\nmask = make_causal_mask((q_len, kv_len))  # rank 2\n// after\nmask = np.tile(make_causal_mask((q_len, kv_len)), (num_heads, 1, 1))","handlingStrategy":"type-guard","validationCode":"assert isinstance(mask, Mask) or (isinstance(mask, np.ndarray) and mask.ndim == 3), mask.shape","typeGuard":"def is_splash_mask(m):\n    return hasattr(m, 'shape') or (isinstance(m, np.ndarray) and m.ndim == 3)","tryCatchPattern":null,"preventionTips":["Broadcast 2D causal masks to (heads, q, kv) with np.tile","Wrap arrays in MultiHeadMask/NumpyMask"],"tags":["jax","splash-attention","mask","shape-validation"],"backgroundTag":"mask-rank-invalid","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}