{"record":{"id":"d2969e75a1d2b762","repo":"jax-ml/jax","slug":"expected-int32-input-but-got-array-dtype","errorCode":null,"errorMessage":"Expected int32 input, but got {array.dtype}.","messagePattern":"Expected int32 input, but got (.+?)\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_mask_info.py","lineNumber":482,"sourceCode":"    data_next_per_head_list.append(data_next_per_head)\n    mask_next_per_head = jnp.concatenate(\n        mask_next_sequence_slices, axis=q_sequence_axis\n    )\n    mask_next_per_head_list.append(mask_next_per_head)\n\n  # Concatenate (or broadcast) the head shards.\n  data_next = jnp.concatenate(data_next_per_head_list, axis=head_axis)\n  mask_next = jnp.concatenate(mask_next_per_head_list, axis=head_axis)\n\n  if is_dkv:\n    partial_mask_blocks = partial_mask_blocks.swapaxes(-1, -2)\n\n  def _downcast(array: jax.Array, max_value: int) -> jax.Array:\n    if array.size == 0:\n      return array\n\n    if array.dtype != np.int32:\n      raise ValueError(f'Expected int32 input, but got {array.dtype}.')\n\n    if max_value <= np.iinfo(np.int8).max:\n      return array.astype(np.int8)\n    elif max_value <= np.iinfo(np.int16).max:\n      return array.astype(np.int16)\n    else:\n      return array.astype(np.int32)\n\n  if downcast_smem_data:\n    block_mask = block_mask.astype(np.int8)  # values are in the range [0, 1, 2]\n    data_next = _downcast(\n        data_next, q_blocks_per_shard if is_dkv else kv_blocks_count\n    )\n    mask_next = _downcast(\n        mask_next, heads_per_shard * q_blocks_per_shard * kv_blocks_count\n    )\n\n  return (","sourceCodeStart":464,"sourceCodeEnd":500,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_mask_info.py#L464-L500","documentation":"The splash attention mask preprocessing downcasts mask index arrays from int32 to int8/int16 to save TPU memory. _downcast only accepts int32 arrays; any other integer width (int64, uint32, etc.) is rejected before the astype.","triggerScenarios":"Passing a dynamic mask index array created with jnp.arange(..., dtype=jnp.int64) or numpy int64/uint arrays to make_splash_attention_mask's dynamic-mask path. On 64-bit-enabled JAX (jax_enable_x64=True) literals and arange default to int64 and trigger this.","commonSituations":"Enabling jax_enable_x64 in a training script then reusing the same mask-building code; constructing mask indices with numpy defaults (int64 on Linux) instead of jnp.int32.","solutions":["Cast mask index arrays to jnp.int32 before passing them: arr.astype(jnp.int32)","Audit mask construction under jax_enable_x64=True; wrap dynamic mask creation in a helper that forces int32"],"exampleFix":"# before\ndynamic_mask = (positions.astype(jnp.int64))  # x64 enabled\n# after\ndynamic_mask = (positions.astype(jnp.int32))","handlingStrategy":"type-guard","validationCode":"assert dynamic_mask_indices.dtype == jnp.int32","typeGuard":"def is_int32(a) -> bool:\n    return a.dtype == jnp.int32","tryCatchPattern":null,"preventionTips":["Force .astype(jnp.int32) on all mask-building arrays","Be extra careful when jax_enable_x64=True"],"tags":["jax","dtype","splash-attention","tpu","int64"],"backgroundTag":"dtype-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}