{"record":{"id":"ae3f5e76cfaa5db9","repo":"jax-ml/jax","slug":"effects-not-supported-in-partial-eval-of-checkpoi","errorCode":null,"errorMessage":"Effects not supported in partial-eval of `checkpoint`/`remat`: {disallowed_effects}","messagePattern":"Effects not supported in partial-eval of `checkpoint`/`remat`: (.+?)","errorType":"error_code","errorClass":"NotImplementedError","httpStatus":null,"severity":"error","filePath":"jax/_src/ad_checkpoint.py","lineNumber":1098,"sourceCode":"\nclass RematTraced(HiPrim):\n  jaxpr: core.Jaxpr\n  policy: Any\n  prevent_cse: bool | tuple[bool, ...]\n\n  def __init__(self, jaxpr, policy, prevent_cse=True):\n    assert (isinstance(prevent_cse, bool) or\n            len(prevent_cse) == len(jaxpr.in_avals))\n    self.in_avals = tuple(jaxpr.in_avals)\n    self.out_aval = jaxpr.out_avals\n    self.params = dict(jaxpr=jaxpr, policy=policy, prevent_cse=prevent_cse)\n    self.effects = frozenset(core.positional_effects(jaxpr))\n    super().__init__()\n\n  def _check_differentiable(self):\n    disallowed = effects.remat_allowed_effects.filter_not_in(self.jaxpr.effects)\n    if disallowed:\n      raise NotImplementedError(\n          'Effects not supported in partial-eval of `checkpoint`/`remat`: '\n          f'{disallowed}')\n\n  @source_info_util.extend_name_stack('checkpoint')\n  def expand(self, *args):\n    return core.eval_jaxpr_p.bind(*args, call_jaxpr=self.jaxpr)\n\n  def vjp_fwd(self, nzs_in, *primals):\n    # TODO eval_jaxpr_p trace time\n    self._check_differentiable()\n    traced = core.jaxpr_as_fun(self.jaxpr)\n    primals_out, fwd2 = remat_transform(self.policy, traced, *primals,\n                                        custom_vjp_rules=True)\n    in_nzs = tuple(tree_leaves(nzs_in))\n    out_nzs_cell = []\n    def make_vjp(*xs):\n      _, f_vjp = api.vjp(fwd2, *xs, in_nzs=in_nzs)\n      out_nzs_cell.append(f_vjp.out_nzs)  # pyrefly: ignore[missing-attribute]","sourceCodeStart":1080,"sourceCodeEnd":1116,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/ad_checkpoint.py#L1080-L1116","documentation":"PyTreeDef::Tuple rejects building a tuple treedef from child PyTreeDefs bound to different registries than the output treedef's registry. All inputs must share the exact same PyTreeRegistry instance.","triggerScenarios":"Calling PyTreeDef.Tuple(registry, [def1, def2]) where some def was created under another registry (global vs custom), or from Python tuple(def1, def2) mixing registries.","commonSituations":"Combining treedefs from jax.tree_util with ones from jax.extend.treeutil registry experiments; migrating code to custom registries piecemeal.","solutions":["Create all child PyTreeDefs from the same registry object","Re-derive mixed treedefs via tree_structure on objects under one registry","Check def.registry() equality before calling Tuple"],"exampleFix":"# before\ncombined = tuple(def_global, def_custom)\n# after\ndef_custom2 = rebuild via same registry as def_global\ncombined = tuple(def_global, def_custom2)","handlingStrategy":"validation","validationCode":"regs = {id(d.registry()) for d in defs}\nassert len(regs) == 1, 'PyTreeDefs come from different registries'\ncombined = treedef_tuple(defs)","typeGuard":"def all_same_registry(defs) -> bool:\n    return all(d.registry() is defs[0].registry() for d in defs)","tryCatchPattern":"try:\n    PyTreeDef.Tuple(registry, defs)\nexcept ValueError as e:\n    if 'Tuple()' in str(e):\n        defs = [_rebuild_under(d, registry) for d in defs]\n    else:\n        raise","preventionTips":["Create treedefs from one registry instance","Assert registry equality in helpers that combine treedefs","Document which registry your library's treedefs use"],"tags":["pytree","jax","registry","tuple"],"backgroundTag":"pytree-registry-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}