{"record":{"id":"8e3e4ba5386a1f7c","repo":"jax-ml/jax","slug":"per-operand-conv-precision-unsupported","errorCode":null,"errorMessage":"Per-operand conv precision unsupported","messagePattern":"Per-operand conv precision unsupported","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic/lowering.py","lineNumber":2986,"sourceCode":"  rhs_spec = dimension_numbers.rhs_spec\n  out_spec = dimension_numbers.out_spec\n\n  def format_dims(dims):\n    return \"[\" + \", \".join(str(d) for d in dims) + \"]\"\n\n  tpu_conv_dims_str = (\n      f\"#tpu.conv_dimension_numbers<{lhs_spec[0]}, {lhs_spec[1]}, \"\n      f\"{format_dims(lhs_spec[2:])}, {rhs_spec[1]}, {rhs_spec[0]}, \"\n      f\"{format_dims(rhs_spec[2:])}, {out_spec[0]}, {out_spec[1]}, \"\n      f\"{format_dims(out_spec[2:])}>\"\n  )\n  return ir.Attribute.parse(tpu_conv_dims_str)\n\n\ndef _parse_precision_attr(precision):\n  if precision is not None:\n    if isinstance(precision, tuple) and precision[0] != precision[1]:\n      raise NotImplementedError(\"Per-operand conv precision unsupported\")\n    precision = precision[0] if isinstance(precision, tuple) else precision\n  if precision is None or precision == lax.Precision.DEFAULT:\n    return None\n  elif precision == lax.Precision.HIGHEST:\n    return ir.Attribute.parse(\"#tpu.contract_precision<fp32>\")\n  else:\n    raise NotImplementedError(f\"Unsupported conv precision: {precision}\")\n\n\n@register_lowering_rule(lax.conv_general_dilated_p)\ndef _conv_general_dilated_lowering_rule(\n    ctx: LoweringRuleContext,\n    lhs,\n    rhs,\n    *,\n    window_strides,\n    padding,\n    lhs_dilation,","sourceCodeStart":2968,"sourceCodeEnd":3004,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic/lowering.py#L2968-L3004","documentation":"The convolution precision parser in Mosaic only supports a single contract precision, so a precision tuple with different values per operand (e.g. (DEFAULT, HIGHEST)) for conv_general_dilated is rejected.","triggerScenarios":"lax.conv_general_dilated with precision=(lax.Precision.DEFAULT, lax.Precision.HIGH) inside a Pallas Mosaic TPU kernel.","commonSituations":"Conv layers with per-operand precision hints ported from XLA/TPU pipelining code into Pallas.","solutions":["Use a uniform precision (scalar or matching tuple)","Pass precision=None"],"exampleFix":"// before\nout = lax.conv_general_dilated(x, w, ..., precision=(Precision.DEFAULT, Precision.HIGH))\n// after\nout = lax.conv_general_dilated(x, w, ..., precision=lax.Precision.HIGH)","handlingStrategy":"validation","validationCode":"if isinstance(precision, tuple):\n    assert precision[0] == precision[1]\nprecision = precision[0] if isinstance(precision, tuple) else precision","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Use uniform conv precision on TPU","Don't share precision configs between dot and conv blindly"],"tags":["jax","pallas","tpu","convolution","precision"],"backgroundTag":"per-operand-precision-unsupported","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}