{"record":{"id":"5be5d140593d80cd","repo":"jax-ml/jax","slug":"custom-vjp-inputs-marked-with-nondiff-argnums-must","errorCode":null,"errorMessage":"custom_vjp inputs marked with nondiff_argnums must be static, not Tracers","messagePattern":"custom_vjp inputs marked with nondiff_argnums must be static, not Tracers","errorType":"exception","errorClass":"UnexpectedTracerError","httpStatus":null,"severity":"error","filePath":"jax/_src/hijax.py","lineNumber":949,"sourceCode":"    self.bwd = bwd\n    self.symz = symbolic_zeros\n    self.opt_remat = optimize_remat\n\n  def defvjp_with_logs(self, fwd, bwd, *, symbolic_zeros=False,\n                       optimize_remat=False):\n    self.defvjp(fwd, bwd, symbolic_zeros=symbolic_zeros,\n                optimize_remat=optimize_remat)\n    self.with_logs = True\n\n  def __call__(self, *args, **kwargs):\n    if not self.fwd or not self.bwd:\n      msg = f\"No VJP defined for custom_vjp function {self.f.__name__} using defvjp.\"\n      raise AttributeError(msg)\n\n    args = resolve_kwargs(self.f, args, kwargs)\n    if any(isinstance(l, core.Tracer) for i in self.static_argnums\n           for l in tree_leaves(args[i])):\n      raise UnexpectedTracerError(\"custom_vjp inputs marked with nondiff_argnums \"\n                                  \"must be static, not Tracers\")\n    if all(is_hashable(args[i]) for i in self.static_argnums):\n      traced = api.jit(self.f, static_argnums=(*self.static_argnums,)).trace(*args)\n    else:\n      # jit requires hashable static_argnums values, but classic custom_vjp\n      # accepted unhashable nondiff_argnums values, so close over them instead\n      which_static = [i in self.static_argnums for i in range(len(args))]\n      dyn_args, static_args = partition_list(which_static, args)\n      f = dyn_args_fun(self.f, self.static_argnums,\n                       tuple(map(WrapHashably, static_args)), len(args))\n      traced = api.jit(f).trace(*dyn_args)\n    args = tuple(Static(x) if i in self.static_argnums else x for i, x in enumerate(args))\n    consts, traced = traced.with_consts_as_arg()\n    fwd_ = update_wrapper(lambda _, __, *args: self.fwd(*args), self.fwd)\n    static_argnums = frozenset(i + 2 for i in self.static_argnums)\n    in_avals = tree_map(typeof, (consts, (), *args))\n    prim = CustomVJPTraced(traced, fwd_, self.bwd, in_avals, self.symz,\n                           static_argnums, self.opt_remat, self.with_logs)","sourceCodeStart":931,"sourceCodeEnd":967,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/hijax.py#L931-L967","documentation":"Arguments of a @jax.custom_vjp function marked as nondiff_argnums must be static (concrete Python values), but one of them was a core.Tracer — i.e. an abstract/traced value arising from jit/grad/vmap/pmap transformation. JAX raises UnexpectedTracerError because static args are passed as Python constants to the rule functions and cannot contain tracers.","triggerScenarios":"Calling a custom_vjp function with nondiff_argnums while a nondiff argument is produced inside a jax.jit, jax.grad, lax.scan, or vmap transformation (e.g. passing a traced array as the 'static' integer argument).","commonSituations":"Passing a shape or index computed from traced arrays (e.g. x.shape[0] inside jit works, but array-derived values traced); reusing the same function with and without jit; forgetting that nondiff_argnums indexes positional args only.","solutions":["Move the traced value out of the nondiff argnum and into a differentiable argument","Compute the static value outside the transformation and pass it in as a plain Python int/str","If the value must remain dynamic, handle it inside the fwd/bwd rules as an array input instead of a static one"],"exampleFix":"// before\n@jax.custom_vjp, nondiff_argnums=(1,)\ndef f(x, n): ...\njit(lambda x, n: f(x, n))(x, jnp.asarray(3))  # traced static arg\n// after\njit(lambda x, n: f(x, int(n)))(x, jnp.asarray(3))\n# or make n a traced diff-arg and use it inside the rules","handlingStrategy":"type-guard","validationCode":"from jax.core import Tracer\ndef safe_static_args(fn, args, static_idx):\n    for i in static_idx:\n        if isinstance(args[i], Tracer):\n            raise TypeError(f'arg {i} is traced; convert to static or make it differentiable')","typeGuard":"from jax.core import Tracer\ndef is_static(v) -> bool:\n    return not isinstance(v, Tracer)","tryCatchPattern":null,"preventionTips":["Convert traced scalars with int()/float() outside jit before passing as nondiff args","Keep nondiff args to plain Python scalars/strings","Prefer passing arrays as diff args"],"tags":["jax","custom-vjp","tracer-error","nondiff-argnums"],"backgroundTag":"jax-unexpected-tracer-error","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}