{"record":{"id":"25158e6012f5a2fd","repo":"jax-ml/jax","slug":"custom-jvp-rule-jvp-name-for-function-primal-na","errorCode":null,"errorMessage":"Custom JVP rule {jvp_name} for function {primal_name} must produce a pair (list or tuple of length two) representing primal and tangent outputs, but got {pair_out}.","messagePattern":"Custom JVP rule (.+?) for function (.+?) must produce a pair \\(list or tuple of length two\\) representing primal and tangent outputs, but got (.+?)\\.","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/custom_derivatives.py","lineNumber":314,"sourceCode":"    flat_fun, out_type1 = _flatten_fun_nokwargs(f_, in_tree)\n    flat_jvp, out_type2 = _flatten_jvp(jvp, primal_name, debug_jvp.func_name,\n                                       in_tree, out_type1)\n    out_flat = custom_jvp_call_p.bind(*args_flat, subfuns=(flat_fun, flat_jvp),\n                                      symbolic_zeros=self.symbolic_zeros)\n    _, (out_tree, _, _) = lu.merge_linear_aux(out_type1, out_type2)\n    return tree_unflatten(out_tree, out_flat)\n\n@partial(lu.transformation_with_aux2, use_eq_store=True)\ndef _flatten_jvp(f, store, primal_name, jvp_name, in_tree, maybe_out_type, *args):\n  primals_in, tangents_in = split_list(args, [len(args) // 2])\n  py_primals = tree_unflatten(in_tree, primals_in)\n  py_tangents = tree_unflatten(in_tree, tangents_in)\n  pair_out = f(py_primals, py_tangents)\n  if not isinstance(pair_out, (list, tuple)) or len(pair_out) != 2:\n    msg = (f\"Custom JVP rule {jvp_name} for function {primal_name} \"\n           \"must produce a pair (list or tuple of length two) representing \"\n           f\"primal and tangent outputs, but got {pair_out}.\")\n    raise TypeError(msg)\n  py_primals_out, py_tangents_out = pair_out\n  primals_out, out_tree = tree_flatten(py_primals_out)\n  tangents_out, out_tree2 = tree_flatten(py_tangents_out)\n  primal_avals = [core.typeof(x) for x in primals_out]\n  if out_tree != out_tree2:\n    msg = (f\"Custom JVP rule {jvp_name} for function {primal_name} must \"\n           \"produce primal and tangent outputs with equal container (pytree) \"\n           f\"structures, but got {out_tree} and {out_tree2} respectively.\")\n    raise TypeError(msg)\n  # If the primal function already ran, check out_tree agreement.\n  try: out_type_ = maybe_out_type()\n  except lu.StoreException: out_type_ = None\n  if out_type_ is not None:\n    out_tree_, primal_avals_, () = out_type_\n    ty_tree  = tree_unflatten(out_tree , [a.str_short() for a in primal_avals])\n    ty_tree_ = tree_unflatten(out_tree_, [a.str_short() for a in primal_avals_])\n    if out_tree_ != out_tree:\n      m = (f\"Custom JVP rule {jvp_name} for function {primal_name} must \"","sourceCodeStart":296,"sourceCodeEnd":332,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/custom_derivatives.py#L296-L332","documentation":"A custom_jvp rule must return a pair — a list/tuple of length two: (primal_output, tangent_output). This TypeError fires when the rule returns something else (a bare array, a 3-tuple, None, etc.).","triggerScenarios":"Writing @f.defjvp that returns only the primal output, or returns (primal, tangent, extra); returning a dict or non-sequence.","commonSituations":"Adapting a jvp function from jax.custom_vjp or custom_gradient styles (TensorFlow habit of returning (out, grad_fn)); partially edited rules.","solutions":["Return exactly (primal_out, tangent_out) as a 2-element tuple/list","If you have multiple outputs, both elements are pytrees of the same shape: ((outs...), (tangents...))","Add a smoke test calling jax.grad on the function"],"exampleFix":"# before\n@f.defjvp\ndef f_jvp(p, t):\n    return f(*p)  # missing tangent\n# after\n@f.defjvp\ndef f_jvp(p, t):\n    return f(*p), jnp.cos(p[0]) * t[0]","handlingStrategy":"validation","validationCode":"def is_jvp_pair(out) -> bool:\n    return isinstance(out, (list, tuple)) and len(out) == 2","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Always end defjvp bodies with `return (primal, tangent)`","Copy the canonical defjvp template from JAX docs"],"tags":["jax","custom-jvp","autodiff","return-shape"],"backgroundTag":"invalid-callback-return-value","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}