{"record":{"id":"1bdf8623b038c170","repo":"jax-ml/jax","slug":"aval-out-dtype","errorCode":null,"errorMessage":"{aval_out.dtype}","messagePattern":"\\{aval_out\\.dtype\\}","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/lowering.py","lineNumber":3473,"sourceCode":"        _dtype_to_ir_type(y_dtype))\n    y = vector.broadcast(y_ty, y)\n  return x, y\n\n\n@register_lowering_rule(\n    lax.add_p, kernel_types=[*tpu_core.CoreType], ensure_mlir_values=False\n)\n@register_lowering_rule(ad_util.add_any_p, ensure_mlir_values=False)\ndef _add_lowering_rule(ctx: LoweringRuleContext, x, y):\n  x, y = _bcast(x, y, ctx.avals_in[0], ctx.avals_in[1], ctx.avals_out[0],\n      ctx.lowering_context.dynamic_shape_replacement_fn,\n  )\n  (aval_out,) = ctx.avals_out\n  if jnp.issubdtype(aval_out.dtype, jnp.integer):\n    return arith.addi(x, y)\n  if jnp.issubdtype(aval_out.dtype, jnp.floating):\n    return arith.addf(x, y)\n  raise NotImplementedError(aval_out.dtype)\n\n\nclass FoldingError(Exception):\n  pass\n\n\ndef _fold(x, fuel):\n  if fuel <= 0:\n    raise FoldingError()\n  op_name = getattr(x.owner, \"name\", None)\n  binop_folds = {\n      \"arith.maxsi\": max,\n      \"arith.minsi\": min,\n  }\n  if op_name == \"arith.constant\":\n    if isinstance(x.type, ir.IntegerType):\n      return ir.IntegerAttr(x.owner.attributes[\"value\"]).value\n    elif isinstance(x.type, ir.FloatType):","sourceCodeStart":3455,"sourceCodeEnd":3491,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/lowering.py#L3455-L3491","documentation":"The Mosaic lowering for lax.add_p dispatches on the output dtype: integers map to arith.addi, floats to arith.addf. Any other dtype category (complex, bool, extended dtypes) reaches the bare raise NotImplementedError(aval_out.dtype).","triggerScenarios":"Adding arrays with complex dtype, bool, or a non-standard dtype inside a Pallas TPU kernel.","commonSituations":"Using complex64 weights inside a TPU kernel (e.g. FFT pipelines); adding boolean masks with + instead of |; custom dtypes enabled via experimental config.","solutions":["Convert operands to float32/float(bfloat16) before adding: lax.convert_element_type(x, jnp.float32)","Replace boolean addition with logical_or","Do the complex arithmetic outside the kernel or via lower_fun on supported parts","Check Mosaic dtype support table for your TPU generation"],"exampleFix":"// before\nz = x + y  # complex64 operands in kernel\n// after\nz = (x.real + y.real) + 1j*(x.imag + y.imag)  # or convert to float32 pairs outside kernel","handlingStrategy":"type-guard","validationCode":"import jax.numpy as jnp\ndef add_supported(dtype):\n    return jnp.issubdtype(dtype, jnp.integer) or jnp.issubdtype(dtype, jnp.floating)","typeGuard":"def is_kernel_safe_dtype(dt): return jnp.issubdtype(dt, jnp.integer) or jnp.issubdtype(dt, jnp.floating)","tryCatchPattern":null,"preventionTips":["Standardize kernels on float32/bfloat16/int32 dtypes","Avoid complex and bool operands in arithmetic inside Pallas kernels"],"tags":["jax","pallas","tpu","dtype","add"],"backgroundTag":"unsupported-dtype-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}