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
- 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
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
- Create treedefs from one registry instance
- Assert registry equality in helpers that combine treedefs
- Document which registry your library's treedefs use
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
- the rematted computation's closure contains a mutable array
- Expected tuple, got %s.
- Tuple arity mismatch: %d != %d; tuple: %s.
- Could not find type: %s.
- numpy masked arrays are not supported as direct inputs to JA
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/ae3f5e76cfaa5db9.
Report an issue: GitHub.