jax-ml/jax · error · TypeError
Custom JVP rule {jvp_name} for function {primal_name} must p
Error message
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}. What it means
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.).
Source
Thrown at jax/_src/custom_derivatives.py:314
flat_fun, out_type1 = _flatten_fun_nokwargs(f_, in_tree)
flat_jvp, out_type2 = _flatten_jvp(jvp, primal_name, debug_jvp.func_name,
in_tree, out_type1)
out_flat = custom_jvp_call_p.bind(*args_flat, subfuns=(flat_fun, flat_jvp),
symbolic_zeros=self.symbolic_zeros)
_, (out_tree, _, _) = lu.merge_linear_aux(out_type1, out_type2)
return tree_unflatten(out_tree, out_flat)
@partial(lu.transformation_with_aux2, use_eq_store=True)
def _flatten_jvp(f, store, primal_name, jvp_name, in_tree, maybe_out_type, *args):
primals_in, tangents_in = split_list(args, [len(args) // 2])
py_primals = tree_unflatten(in_tree, primals_in)
py_tangents = tree_unflatten(in_tree, tangents_in)
pair_out = f(py_primals, py_tangents)
if not isinstance(pair_out, (list, tuple)) or len(pair_out) != 2:
msg = (f"Custom JVP rule {jvp_name} for function {primal_name} "
"must produce a pair (list or tuple of length two) representing "
f"primal and tangent outputs, but got {pair_out}.")
raise TypeError(msg)
py_primals_out, py_tangents_out = pair_out
primals_out, out_tree = tree_flatten(py_primals_out)
tangents_out, out_tree2 = tree_flatten(py_tangents_out)
primal_avals = [core.typeof(x) for x in primals_out]
if out_tree != out_tree2:
msg = (f"Custom JVP rule {jvp_name} for function {primal_name} must "
"produce primal and tangent outputs with equal container (pytree) "
f"structures, but got {out_tree} and {out_tree2} respectively.")
raise TypeError(msg)
# If the primal function already ran, check out_tree agreement.
try: out_type_ = maybe_out_type()
except lu.StoreException: out_type_ = None
if out_type_ is not None:
out_tree_, primal_avals_, () = out_type_
ty_tree = tree_unflatten(out_tree , [a.str_short() for a in primal_avals])
ty_tree_ = tree_unflatten(out_tree_, [a.str_short() for a in primal_avals_])
if out_tree_ != out_tree:
m = (f"Custom JVP rule {jvp_name} for function {primal_name} must "View on GitHub (pinned to 1e1c6a8fc0)
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
Example fix
# before
@f.defjvp
def f_jvp(p, t):
return f(*p) # missing tangent
# after
@f.defjvp
def f_jvp(p, t):
return f(*p), jnp.cos(p[0]) * t[0] Defensive patterns
Strategy: validation
Validate before calling
def is_jvp_pair(out) -> bool:
return isinstance(out, (list, tuple)) and len(out) == 2 Prevention
- Always end defjvp bodies with `return (primal, tangent)`
- Copy the canonical defjvp template from JAX docs
When it happens
Trigger: Writing @f.defjvp that returns only the primal output, or returns (primal, tangent, extra); returning a dict or non-sequence.
Common situations: Adapting a jvp function from jax.custom_vjp or custom_gradient styles (TensorFlow habit of returning (out, grad_fn)); partially edited rules.
Related errors
- Pure callbacks do not support JVP. Please use `jax.custom_jv
- Can't use ``defjvps`` with ``nondiff_argnums``.
- No JVP defined for custom_jvp function {primal_name} using d
- Custom JVP rule {jvp_name} for function {primal_name} must p
- Custom JVP rule {jvp_name} for function {primal_name} must p
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/25158e6012f5a2fd.
Report an issue: GitHub.