{"record":{"id":"10b74f7997168219","repo":"jax-ml/jax","slug":"transpose-of-einsum-with-multiple-linear-inputs-is","errorCode":null,"errorMessage":"Transpose of Einsum with multiple linear inputs is not supported.","messagePattern":"Transpose of Einsum with multiple linear inputs is not supported\\.","errorType":"exception","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/numpy/hijax.py","lineNumber":443,"sourceCode":"    return batched_prim(*args), 0\n\n  def jvp(self, primals: tuple[Array, ...], tangents: Any) -> tuple[Array, Array]:\n    primal_out = self(*primals)\n    tangent_outs = []\n    for i, t in enumerate(tangents):\n      if not isinstance(t, ad_util.Zero):\n        tangent_outs.append(self(*primals[:i], t, *primals[i+1:]))\n    if not tangent_outs:\n      return primal_out, ad_util.zeros_like_aval(self.out_aval)\n    return primal_out, functools.reduce(lax.add, tangent_outs)\n\n  def transpose(self, out_ct, *maybe_accums):\n    in_subs_list = self.subscripts.split('->')[0].split(',')\n    out_sub = self.subscripts.split('->')[1]\n\n    accums = [acc for acc in maybe_accums if isinstance(acc, ad.GradAccum)]\n    if len(accums) > 1:\n      raise NotImplementedError(\"Transpose of Einsum with multiple linear inputs is not supported.\")\n\n    for i, accum in enumerate(maybe_accums):\n      if isinstance(accum, ad.GradAccum):\n        if isinstance(out_ct, ad_util.Zero):\n          accum.accum(ad_util.zeros_like_aval(self.in_avals[i]))\n          continue\n\n        orig_sub = in_subs_list[i]\n        orig_aval = self.in_avals[i]\n\n        all_ct_input_chars = set(out_sub)\n        for k, sub in enumerate(in_subs_list):\n          if k != i:\n            all_ct_input_chars.update(sub)\n\n        missing_chars = sorted(\n            set(orig_sub) - all_ct_input_chars, key=orig_sub.index\n        )","sourceCodeStart":425,"sourceCodeEnd":461,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/numpy/hijax.py#L425-L461","documentation":"Raised by Einsum.transpose (the VJP rule) when gradient accumulation is requested for more than one input simultaneously. The linearized transpose implementation only supports a single GradAccum per einsum; multiple linear inputs under custom accumulation (e.g. jax.checkpoint/accumulation APIs) are unimplemented.","triggerScenarios":"Computing a VJP of an einsum with two or more linear operands while using gradient accumulation (ad.GradAccum), e.g. inside remat/checkpointed code with per-input accumulation; the trace reaches _nonzero_impl-adjacent transpose machinery and hits this NotImplementedError.","commonSituations":"Backprop through large einsum-based models under jax.checkpoint or custom gradient accumulation; upgrading code that previously differentiated without accumulation.","solutions":["Accumulate gradients for only one einsum input at a time (run separate VJPs per input)","Avoid gradient accumulation around this einsum: remove jax.checkpoint/custom accum on that call or split the einsum so each has one linear input","File/track an upstream issue in JAX for multi-input accumulation support"],"exampleFix":"# before\npull = jax.vjp(lambda a, b: jnp.einsum('ij,jk->ik', a, b), a, b)\n# with accumulation on both inputs -> NotImplementedError\n# after: accumulate one input at a time\n_, vjp_a = jax.vjp(lambda a: jnp.einsum('ij,jk->ik', a, b_stop), a)\n_, vjp_b = jax.vjp(lambda b: jnp.einsum('ij,jk->ik', a_stop, b), b)","handlingStrategy":"fallback","validationCode":null,"typeGuard":null,"tryCatchPattern":"try:\n    grads = jax.grad(loss)(params)\nexcept NotImplementedError as e:\n    if 'multiple linear inputs' in str(e):\n        # differentiate one einsum input at a time\n        grads = {k: jax.grad(lambda p: loss_with_fixed(others, k, p))(p) for k, p in params.items()}\n    else:\n        raise","preventionTips":["Avoid gradient accumulation (checkpoint/remat accum) around multi-input einsums","Split multi-input einsums into single-linear-input stages","Pin and test the JAX version when relying on HiJAX transpose rules"],"tags":["jax","einsum","autodiff","not-implemented"],"backgroundTag":"autodiff-unsupported-operation","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}