{"record":{"id":"80e7d6974fdc7bee","repo":"jax-ml/jax","slug":"unsupported-dtype-for-sign-x-dtype","errorCode":null,"errorMessage":"Unsupported dtype for sign: {x.dtype}","messagePattern":"Unsupported dtype for sign: (.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/lowering.py","lineNumber":3824,"sourceCode":"      sign_val = lax.bitcast_convert_type(sign_val_i, jnp.float32)\n      # By checking abs(x32) > 0.0 we handle NaN and +/-0.0.\n      res = jnp.where(jnp.abs(x32) > 0.0, sign_val, x32)\n\n      if dtype == jnp.bfloat16:\n        assert not tpu_has_native_bf16\n        # Drop the rightmost 16 bits, which are all zero.\n        res_i = lax.bitcast_convert_type(res, jnp.uint32)\n        res_u16 = lax.convert_element_type(\n            lax.shift_right_logical(res_i, jnp.uint32(16)), jnp.uint16\n        )\n        return lax.bitcast_convert_type(res_u16, jnp.bfloat16)\n      return res.astype(dtype)\n\n    if jnp.issubdtype(x.dtype, jnp.signedinteger):\n      return (x > 0).astype(x.dtype) - (x < 0).astype(x.dtype)\n    if jnp.issubdtype(x.dtype, jnp.unsignedinteger):\n      return (x != 0).astype(x.dtype)\n    raise ValueError(f\"Unsupported dtype for sign: {x.dtype}\")\n\n  return lower_fun(_lower_fun)(ctx, x)\n\n\n@register_lowering_rule(lax.nextafter_p)\ndef _nextafter_lowering_rule(ctx: LoweringRuleContext, x, y):\n  return lower_fun(\n      pallas_utils.nextafter_lowering_helper,\n  )(ctx, x, y)\n\n\n@register_lowering_rule(\n    lax.rsqrt_p,\n    kernel_types=(tpu_core.CoreType.TC, tpu_core.CoreType.SC_VECTOR_SUBCORE),\n)\ndef _rsqrt_lowering_rule(ctx: LoweringRuleContext, x, accuracy=None):\n  if accuracy is not None:\n    raise NotImplementedError(\"Not implemented: accuracy\")","sourceCodeStart":3806,"sourceCodeEnd":3842,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/lowering.py#L3806-L3842","documentation":"The Mosaic sign lowering handles floats, signed ints (via comparison trick), and unsigned ints. Any other dtype (e.g. complex, bool) raises ValueError('Unsupported dtype for sign: ...'). Note this is a ValueError, not NotImplementedError, and comes from the Python-level _lower_fun helper.","triggerScenarios":"Calling jnp.sign/lax.sign on complex or boolean arrays inside a Pallas kernel.","commonSituations":"Sign of complex numbers (ill-defined anyway); sign of booleans; unexpected dtype promotion.","solutions":["Define sign semantics yourself for complex (e.g. z/|z|) and implement manually","Cast to float32/int32 before sign","Avoid jnp.sign on bools — use the mask directly"],"exampleFix":"// before\ns = jnp.sign(z)  # complex\n// after\nmag = jnp.sqrt(z.real**2 + z.imag**2)\ns = jnp.where(mag > 0, z / mag, 0)","handlingStrategy":"type-guard","validationCode":"import jax.numpy as jnp\ndef sign_dtype_ok(dt):\n    return (jnp.issubdtype(dt, jnp.floating) or jnp.issubdtype(dt, jnp.signedinteger)\n            or jnp.issubdtype(dt, jnp.unsignedinteger))","typeGuard":"def is_sign_supported(dt): return not jnp.issubdtype(dt, jnp.complexfloating)","tryCatchPattern":null,"preventionTips":["Avoid jnp.sign on complex/bool in kernels","Define custom complex sign semantics if needed"],"tags":["jax","pallas","tpu","dtype","sign"],"backgroundTag":"unsupported-dtype-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}