{"record":{"id":"ee6b9b051840c2cb","repo":"jax-ml/jax","slug":"hlo-comparison-direction-for-extended-dtype-ava","errorCode":null,"errorMessage":"HLO comparison {direction} for extended dtype {avals_in[0].dtype}","messagePattern":"HLO comparison (.+?) for extended dtype (.+?)","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/lax.py","lineNumber":5236,"sourceCode":"  return mlir.delegate_lowering(\n      ctx, partial(_unary_reduce_lower, reduction_op, identity,\n                   axes=reduce_axes),\n      res, avals_in=[base_aval_out], avals_out=[aval_out])\n\n_opaque_eq_hlo = partial(\n    _opaque_comparison_hlo, 'EQ', hlo.AndOp, _get_bitwise_and_identity)\n_opaque_ne_hlo = partial(\n    _opaque_comparison_hlo, 'NE', hlo.OrOp, _get_bitwise_or_identity)\n\ndef _compare_lower_hlo_opaque(direction: str, ctx, avals_in, aval_out, x, y):\n  broadcast_avals_in = tuple(\n      core.ShapedArray(aval_out.shape, aval.dtype) for aval in avals_in)\n  if direction == 'EQ':\n    return _opaque_eq_hlo(ctx, broadcast_avals_in, aval_out, x, y)\n  elif direction == 'NE':\n    return _opaque_ne_hlo(ctx, broadcast_avals_in, aval_out, x, y)\n  else:\n    raise NotImplementedError(\n        f\"HLO comparison {direction} for extended dtype {avals_in[0].dtype}\")\n\n\ndef _compare_lower_hlo(direction: str, total_order: bool, ctx, x, y):\n  avals_in, (aval_out,) = ctx.avals_in, ctx.avals_out\n  x_dtype = avals_in[0].dtype\n  x, y = mlir.multi_broadcast_in_dim(ctx, (x, y), avals_in, aval_out.shape,\n                                     aval_out.sharding)\n  if dtypes.issubdtype(x_dtype, dtypes.extended):\n    assert not total_order\n    return _compare_lower_hlo_opaque(direction, ctx, avals_in, aval_out, x, y)\n  if dtypes.issubdtype(x_dtype, np.inexact):\n    compare_type = \"TOTALORDER\" if total_order else \"FLOAT\"\n  elif dtypes.issubdtype(x_dtype, np.signedinteger):\n    compare_type = \"SIGNED\"\n  else:\n    compare_type = \"UNSIGNED\"\n  return [mlir.compare_hlo(x, y, direction, compare_type)]","sourceCodeStart":5218,"sourceCodeEnd":5254,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/lax.py#L5218-L5254","documentation":"When lowering a comparison operation on an extended dtype (like jax's ml_dtypes float8 or custom extension types), only EQ and NE comparisons are implemented via special HLO helpers. Requesting any other direction (LT, LE, GT, GE) on an opaque/extended dtype has no lowering, so NotImplementedError is raised.","triggerScenarios":"Calling jnp.less/greater/less_equal/greater_equal (or lax.lt etc.) on arrays with an ExtendedDType whose comparison is opaque, e.g. comparing float8_* or custom extension dtypes; equality (==, !=) works but ordering does not.","commonSituations":"Using float8 dtypes with jnp.sort, jnp.maximum, clipping, or boolean masks that lower to ordered comparisons; assuming all numpy comparison semantics carry over to new extended dtypes.","solutions":["Use equality comparisons (jnp.equal / jnp.not_equal) which are supported","Convert to a standard dtype (e.g. .astype(jnp.float32)) before ordering comparisons","For float8, compare via bitcast to uint8 only if you understand the bit layout","Implement/extend the dtype rules if it is a custom extended dtype"],"exampleFix":"// before\nmask = x_f8 < y_f8\n\n// after\nmask = x_f8.astype(jnp.float32) < y_f8.astype(jnp.float32)","handlingStrategy":"type-guard","validationCode":"def is_orderable(dtype):\n    return not isinstance(dtype, jax.dtypes.ExtendedDType)","typeGuard":"def orderable(x):\n    return not isinstance(x.dtype, jax.dtypes.ExtendedDType)","tryCatchPattern":"try:\n    mask = x < y\nexcept NotImplementedError:\n    mask = x.astype(jnp.float32) < y.astype(jnp.float32)","preventionTips":["Reserve ordering comparisons for standard dtypes","Convert extended dtypes to float32 before sort/min/max/mask logic","Document which ops are EQ/NE-only for extended dtypes"],"tags":["jax","extended-dtype","comparison","lax","hlo"],"backgroundTag":"unsupported-dtype-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}