{"record":{"id":"24497941684bb8cb","repo":"jax-ml/jax","slug":"unexpected-precision-specifier-value-precision","errorCode":null,"errorMessage":"Unexpected precision specifier value {precision}","messagePattern":"Unexpected precision specifier value (.+?)","errorType":"exception","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/experimental/rnn.py","lineNumber":267,"sourceCode":"  #   if precision is None and config.jax_default_matmul_precision is not None:\n  #     precision = Precision(config.jax_default_matmul_precision)\n  #   else:\n  #     precision = None\n  #\n  # but we prefer to still invoke it here for consistency\n  precision = lax.canonicalize_precision(precision)\n  if precision is None or not (isinstance(precision, tuple) and len(precision) == 2):\n    return True\n  # cuDNN allows only one precision specifier per RNN op\n  match precision:\n    case (lax.Precision.HIGHEST, _):\n      return False\n    case (lax.Precision.HIGH, _):\n      return True\n    case (lax.Precision.DEFAULT, _): # bfloat16\n      raise NotImplementedError(\"bfloat16 support not implemented for LSTM\")\n    case _:\n      raise ValueError(f\"Unexpected precision specifier value {precision}\")\n\n\n@partial(custom_vjp, nondiff_argnums=(5, 6, 7, 8, 9, 10))\ndef lstm(x: Array, h_0: Array, c_0: Array, weights: Array, seq_lengths: Array,\n         input_size: int, hidden_size: int, num_layers: int, dropout: float,\n         bidirectional: bool, precision: lax.PrecisionLike = None) -> tuple[Array, Array, Array]:\n  \"\"\"LSTM via CuDNN or HIPDNN (not-yet-supported).\n\n  Assume batch-first inputs.\n\n  Arguments:\n    x: (batch_size, max_seq_length, input_size)\n    h_0: (num_directions * num_layers, batch_size, hidden_size)\n    c_0: (num_directions * num_layers, batch_size, hidden_size)\n    weights: (num_params,) where num_params = get_num_params_in_lstm(...)\n    seq_lengths: (batch_size,)\n  Returns: (y, h_n, c_n, reserve_space).\n    y: (batch_size, max_seq_length, hidden_size * num_directions)","sourceCodeStart":249,"sourceCodeEnd":285,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/experimental/rnn.py#L249-L285","documentation":"The LSTM precision parser accepts only specific lax.Precision values (HIGHEST, HIGH, DEFAULT) in tuple or scalar form. Anything else — an invalid enum, a string like 'high', or a malformed tuple — raises ValueError.","triggerScenarios":"Calling jax.experimental.rnn.lstm with precision set to a raw string ('high'), a lax.Precision combined with lax.Precision via unsupported ops, or a nested/incorrect tuple structure.","commonSituations":"Copy-pasting precision strings from JAX docs for matmul-like APIs; passing precision=(lax.Precision.DEFAULT, None) or enum combos the matcher doesn't handle.","solutions":["Pass a valid lax.Precision enum or tuple of enums, e.g. lax.Precision.HIGHEST or (lax.Precision.HIGHEST, lax.Precision.HIGHEST)","Use None only if the version supports it; otherwise pick an explicit enum","Check the installed JAX version's accepted values in rnn.py"],"exampleFix":"# before\nlstm(..., precision='high')\n# after\nlstm(..., precision=lax.Precision.HIGH)","handlingStrategy":"type-guard","validationCode":"assert precision is None or precision in (lax.Precision.HIGHEST, lax.Precision.HIGH, lax.Precision.DEFAULT) or all(p in lax.Precision for p in precision)","typeGuard":"def valid_precision(p) -> bool:\n    ok = {lax.Precision.HIGHEST, lax.Precision.HIGH, lax.Precision.DEFAULT}\n    if isinstance(p, tuple):\n        return all(x in ok for x in p)\n    return p is None or p in ok","tryCatchPattern":null,"preventionTips":["Only pass lax.Precision enums, never strings"],"tags":["jax","lstm","rnn","invalid-argument","precision"],"backgroundTag":"invalid-enum-value","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}