patrick-kidger / patrick-kidger/diffrax

`EQX_ON_ERROR=breakpoint` and `diffeqsolve` causes tracer leakage ?

Open
#512 2 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

question
Dominant language
Python
Stars
2.1k
Forks
189
Avg merge
3d 18h
Merged PRs (30d)
1

Description

Hi Patrick, i'm getting weird behavior with HEAD and its use with EQX_ON_ERROR=breakpoint.
Here is a MWE :


import jax
import jax.numpy as jnp
import diffrax

with jax.checking_leaks():

    ts = jnp.linspace(0.0, 1.0, 10)
    ys = diffrax.diffeqsolve(
        diffrax.ODETerm(lambda t, y, args: -y),
        diffrax.Bosh3(),
        t0=ts[0],
        t1=ts[-1],
        dt0=ts[1] - ts[0],
        y0=jnp.ones((1, )),
        saveat=diffrax.SaveAt(ts=ts),
    )

Launching the file once is OK.
Then if I set the bash variable export EQX_ON_ERROR=breakpoint in another terminal and re-run the MWE, the error stack outputs leaked tracers (and EQX_ON_ERROR=breakpoint doesn't open a jax.debug.breakpoint where the error arises) :

Traceback (most recent call last):
  File "/home/monsel/Desktop/sandbox_diffrax/mwe.py", line 10, in <module>
    ys = diffrax.diffeqsolve(
  File "/home/monsel/miniconda3/envs/sandbox_diffrax/lib/python3.10/site-packages/equinox/_jit.py", line 239, in __call__
    return self._call(False, args, kwargs)
  File "/home/monsel/miniconda3/envs/sandbox_diffrax/lib/python3.10/site-packages/equinox/_module.py", line 1093, in __call__
    return self.__func__(self.__self__, *args, **kwargs)
  File "/home/monsel/miniconda3/envs/sandbox_diffrax/lib/python3.10/site-packages/equinox/_jit.py", line 212, in _call
    out = self._cached(dynamic_donate, dynamic_nodonate, static)
  File "/home/monsel/miniconda3/envs/sandbox_diffrax/lib/python3.10/site-packages/equinox/_errors.py", line 187, in fixed_jit_impl
    return jit_fun(*args2, **kwargs2)
  File "/home/monsel/miniconda3/envs/sandbox_diffrax/lib/python3.10/contextlib.py", line 142, in __exit__
    next(self.gen)
Exception: Leaked trace MainTrace(1,DynamicJaxprTrace). Leaked tracer(s):

Traced<ShapedArray(float32[1])>with<DynamicJaxprTrace(level=1/0)>
This DynamicJaxprTracer was created on line /home/monsel/Desktop/sandbox_diffrax/mwe.py:10 (<module>)
<DynamicJaxprTracer 139338190526784> is referred to by <function 139338169797072> (_allocate_output) closed-over variable y0
<function 139338169797072> is referred to by <list 139338138048256>[11]
<list 139338138048256> is referred to by <tuple 139338137653056>[1]

Traced<ShapedArray(float32[])>with<DynamicJaxprTrace(level=1/0)>
This DynamicJaxprTracer was created on line /home/monsel/Desktop/sandbox_diffrax/mwe.py:10 (<module>)
<DynamicJaxprTracer 139338169959664> is referred to by <function 139338169797072> (_allocate_output) closed-over variable t0
<function 139338169797072> is referred to by <list 139338138048256>[11]
<list 139338138048256> is referred to by <tuple 139338137653056>[1]

Traced<ShapedArray(int32[], weak_type=True)>with<DynamicJaxprTrace(level=1/0)>
This DynamicJaxprTracer was created on line /home/monsel/Desktop/sandbox_diffrax/mwe.py:10 (<module>)
<DynamicJaxprTracer 139338169486448> is referred to by <function 139338169797072> (_allocate_output) closed-over variable direction
<function 139338169797072> is referred to by <list 139338138048256>[11]
<list 139338138048256> is referred to by <tuple 139338137653056>[1]

Traced<ShapedArray(float32[])>with<DynamicJaxprTrace(level=1/0)>
This DynamicJaxprTracer was created on line /home/monsel/Desktop/sandbox_diffrax/mwe.py:10 (<module>)
<DynamicJaxprTracer 139338190525104> is referred to by <function 139338169246064> (_check_subsaveat_ts) closed-over variable t1
<function 139338169246064> is referred to by <list 139338138048256>[5]
<list 139338138048256> is referred to by <tuple 139338137653056>[1]

equinox==0.11.7
jax==0.4.34
jaxlib==0.4.34
jaxtyping==0.2.34
lineax==0.0.6
ml_dtypes==0.5.0
numpy==2.1.2
opt_einsum==3.4.0
optimistix==0.0.8
scipy==1.14.1
typeguard==2.13.3
typing_extensions==4.12.2

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Reproduce the report using the MWE in mwe.py with the listed JAX, Equinox, and Diffrax versions, comparing normal execution with EQX_ON_ERROR=breakpoint. Start at diffrax.diffeqsolve and the traceback path through equinox/_errors.py, then trace the leaked values held by _allocate_output and _check_subsaveat_ts. Done means no tracer leakage is reported and the breakpoint behavior works at the originating error.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
backend
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
35/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.