patrick-kidger / patrick-kidger/diffrax

Zero gradient when using jnp.piecewise inside an ODE

Open
#363 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,
applying jax.grad to a function which uses diffrax to integrate a piecewise defined ODE, I observe that one partial derivative is unexpectedly zero. The ODE solver returns correct function values, just the gradient is wrong. I’m wondering whether this is a bug, or whether I’m doing something wrong.
Thanks in advance!
David

Example:

Consider the piecewise defined ODE

$$\frac{dy}{dt} = -k(t) \cdot y, \qquad y(0) = y_0, \qquad \mathrm{with} \ k(t) = \begin{cases} k_0, \ t \leq T \
0, \ t > T \end{cases},$$

to which the solution reads

$$y(t) = y_0 \cdot \begin{cases} e^{-k_0 t}, \ t \leq T \
e^{-k_0 T}, \ t > T \end{cases}$$

I'm interested in the partial derivatives w.r.t. $T$ and $k_0$. In the code example below, I compare the gradient obtained from integrating the ODE using diffrax to the analytical solution and to a finite difference calculation.

import jax
import jax.numpy as jnp
import diffrax
jax.config.update("jax_enable_x64", True)
jax.config.update("jax_debug_nans", True)


def calc_y_analytically(params, t, y0):
    T, k0 = params
    return y0 * jnp.piecewise(t, [t <= T, t > T], [lambda x: jnp.exp(-k0*x), lambda x: jnp.exp(-k0*T)])


def calc_y_ode(params, t, y0):

    def ode(t, y, args):
        T, k0 = args
        k_of_t = jnp.piecewise(t, [t <= T, t > T], [k0, 0.0])
        d_y = -1 * k_of_t * y
        return d_y

    term = diffrax.ODETerm(ode)
    solver = diffrax.Tsit5()
    sol = diffrax.diffeqsolve(term, solver, t0=0.0, t1=t, dt0=0.00001, y0=y0, args=params, max_steps=200000)
    return sol.ys[0]


if __name__ == '__main__':

    params = jnp.array([1.0, 5.5])  # (T, k0)
    t = 1.2
    y0 = 100.0

    # calculate y(t) and the gradient w.r.t. T and k0 analytically
    y_ana, grads_ana = jax.value_and_grad(calc_y_analytically)(params, t, y0)

    # propagate y(0)=y0 until t by solving the ODE and calculate the gradient w.r.t. T and k0
    y_diff, grads_diff = jax.value_and_grad(calc_y_ode)(params, t, y0)

    # perform finite differences method on 0th parameter for verification
    eps = 1e-4
    param0 = params[0]
    params_plus = params.at[0].set(param0 + eps)
    params_minus = params.at[0].set(param0 - eps)

    y_ana_plus = calc_y_analytically(params_plus, t, y0)
    y_ana_minus = calc_y_analytically(params_minus, t, y0)
    part_deriv_ana = (y_ana_plus - y_ana_minus) / (2 * eps)

    y_diff_plus = calc_y_ode(params_plus, t, y0)
    y_diff_minus = calc_y_ode(params_minus, t, y0)
    part_deriv_diff = (y_diff_plus - y_diff_minus) / (2 * eps)

    print('\ny(t):')
    print('Analytical: {y:.6f}'.format(y=y_ana))
    print('Diffrax:    {y:.6f}'.format(y=y_diff))

    print('\nGradient:')
    print('Analytical: ' + str(grads_ana))
    print('Diffrax:    ' + str(grads_diff))

    print('\nPartial derivative w.r.t. parameter 0 via finite difference:')
    print('Analytical: {p:.6f}'.format(p=part_deriv_ana))
    print('Diffrax:    {p:.6f}'.format(p=part_deriv_diff))

prints out the following:

y(t):
Analytical: 0.408677
Diffrax:    0.408675

Gradient:
Analytical: [-2.24772429 -0.40867714]
Diffrax:    [ 0.         -0.40867537]

Partial derivative w.r.t. parameter 0 via finite difference:
Analytical: -2.247724
Diffrax:    -2.247712

I'm using Python 3.11.7, jax 0.4.23, jaxlib 0.4.23.dev20231223, diffrax 0.5.0, MacOS 14.2.1, x86_64, running on CPU

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

Start with the provided calc_y_ode reproducer and its jax.value_and_grad call, then trace the diffrax.diffeqsolve path for the piecewise ODE. Compare the solver gradient with the analytical and finite-difference results; done when the derivative with respect to T is correctly propagated without regressing the k0 derivative.

Written by the indexing model from the issue text.

Assessment

Tech stack
python
Domain
machine-learning
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.