jax-ml/jax · error · NotImplementedError

Effects not supported in partial-eval of `checkpoint`/`remat

Error message

Effects not supported in partial-eval of `checkpoint`/`remat`: {disallowed_effects}

What it means

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.

Source

Thrown at jax/_src/ad_checkpoint.py:1098

class RematTraced(HiPrim):
  jaxpr: core.Jaxpr
  policy: Any
  prevent_cse: bool | tuple[bool, ...]

  def __init__(self, jaxpr, policy, prevent_cse=True):
    assert (isinstance(prevent_cse, bool) or
            len(prevent_cse) == len(jaxpr.in_avals))
    self.in_avals = tuple(jaxpr.in_avals)
    self.out_aval = jaxpr.out_avals
    self.params = dict(jaxpr=jaxpr, policy=policy, prevent_cse=prevent_cse)
    self.effects = frozenset(core.positional_effects(jaxpr))
    super().__init__()

  def _check_differentiable(self):
    disallowed = effects.remat_allowed_effects.filter_not_in(self.jaxpr.effects)
    if disallowed:
      raise NotImplementedError(
          'Effects not supported in partial-eval of `checkpoint`/`remat`: '
          f'{disallowed}')

  @source_info_util.extend_name_stack('checkpoint')
  def expand(self, *args):
    return core.eval_jaxpr_p.bind(*args, call_jaxpr=self.jaxpr)

  def vjp_fwd(self, nzs_in, *primals):
    # TODO eval_jaxpr_p trace time
    self._check_differentiable()
    traced = core.jaxpr_as_fun(self.jaxpr)
    primals_out, fwd2 = remat_transform(self.policy, traced, *primals,
                                        custom_vjp_rules=True)
    in_nzs = tuple(tree_leaves(nzs_in))
    out_nzs_cell = []
    def make_vjp(*xs):
      _, f_vjp = api.vjp(fwd2, *xs, in_nzs=in_nzs)
      out_nzs_cell.append(f_vjp.out_nzs)  # pyrefly: ignore[missing-attribute]

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Create all child PyTreeDefs from the same registry object
  2. Re-derive mixed treedefs via tree_structure on objects under one registry
  3. Check def.registry() equality before calling Tuple

Example fix

# before
combined = tuple(def_global, def_custom)
# after
def_custom2 = rebuild via same registry as def_global
combined = tuple(def_global, def_custom2)
Defensive patterns

Strategy: validation

Validate before calling

regs = {id(d.registry()) for d in defs}
assert len(regs) == 1, 'PyTreeDefs come from different registries'
combined = treedef_tuple(defs)

Type guard

def all_same_registry(defs) -> bool:
    return all(d.registry() is defs[0].registry() for d in defs)

Try / catch

try:
    PyTreeDef.Tuple(registry, defs)
except ValueError as e:
    if 'Tuple()' in str(e):
        defs = [_rebuild_under(d, registry) for d in defs]
    else:
        raise

Prevention

When it happens

Trigger: 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.

Common situations: Combining treedefs from jax.tree_util with ones from jax.extend.treeutil registry experiments; migrating code to custom registries piecemeal.

Related errors


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/ae3f5e76cfaa5db9. Report an issue: GitHub.