
.. DO NOT EDIT.
.. THIS FILE WAS AUTOMATICALLY GENERATED BY SPHINX-GALLERY.
.. TO MAKE CHANGES, EDIT THE SOURCE PYTHON FILE:
.. "auto_examples/15_adaptive_ode.py"
.. LINE NUMBERS ARE GIVEN BELOW.

.. only:: html

    .. note::
        :class: sphx-glr-download-link-note

        :ref:`Go to the end <sphx_glr_download_auto_examples_15_adaptive_ode.py>`
        to download the full example code.

.. rst-class:: sphx-glr-example-title

.. _sphx_glr_auto_examples_15_adaptive_ode.py:


Adaptive-step ODE integration (diffrax)
========================================

Every other ODE example uses a fixed-step integrator (``"rk4"``, ``"rk6"``,
...): the same ``dt`` is both the sampling interval between QR
renormalizations *and* the integrator's own internal step size, so getting
a smaller local truncation error means shrinking ``dt`` globally, which also
means more (cheap) raw steps between renormalizations.

``lyapax.adaptive.diffrax_adaptive_step`` decouples those two meanings of
``dt``. It builds an integrator, backed by `diffrax
<https://docs.kidger.site/diffrax/>`_, whose internal step size is chosen
adaptively (via a PID step-size controller and an embedded Runge-Kutta error
estimate) to hit a requested local error tolerance -- while ``dt`` itself
keeps meaning only "the interval between renormalizations". It drops
straight into ``lyapax.ode_problem`` alongside the fixed-step builtins, with
no other code changes: ``lyapunov_spectrum`` cannot tell the difference
between an integrator that takes one internal step per call and one that
takes an internally-controlled, variable number of them.

**Reading the plot.** The left panel sweeps the controller's tolerance
(``rtol``, with ``atol`` scaled alongside it) and shows the Lyapunov
exponent estimate converging toward the published Lorenz value as the
tolerance tightens -- the adaptive analogue of a ``dt``-convergence sweep.
The right panel cross-checks the tightest-tolerance adaptive estimate
against a fine-``dt`` fixed-step ``rk4`` run of the same system: both should
agree with each other and with the published value, confirming the adaptive
path is solving the same problem, not just converging to *something*.

**A real caveat.** The adaptive integrator's internal accept/reject/step
-size loop is a ``jax.lax.while_loop`` with a data-dependent trip count.
Forward-mode differentiation (``jax.jvp``/``jax.jacfwd``) works through it
fine -- that's what the Lyapunov engine's own tangent propagation already
uses. Reverse-mode (``jax.grad``/``jax.jacrev``) does not: JAX cannot
replay a dynamic-trip-count ``while_loop`` backward. The last section below
demonstrates differentiating a Lyapunov exponent with respect to a system
parameter using ``jax.jacfwd`` (which works and matches finite differences)
and shows ``jax.grad`` raising instead of silently returning something
wrong.

.. GENERATED FROM PYTHON SOURCE LINES 43-57

.. code-block:: Python

    import os

    os.environ["JAX_PLATFORMS"] = "cpu"

    import jax
    import jax.numpy as jnp
    import matplotlib.pyplot as plt

    jax.config.update("jax_enable_x64", True)

    from lyapax import lyapunov_spectrum, ode_problem, systems
    from lyapax.adaptive import diffrax_adaptive_step
    from lyapax.integrators import rk4_step








.. GENERATED FROM PYTHON SOURCE LINES 58-74

.. code-block:: Python

    sigma, rho, beta = 10.0, 28.0, 8.0 / 3.0
    rhs = systems.lorenz(sigma, rho, beta)
    state0 = jnp.array([1.0, 1.0, 1.0])
    published_lambda1 = 0.9056

    rtols = [1e-3, 1e-4, 1e-6, 1e-8, 1e-10]
    errors = []
    for rtol in rtols:
        integrator = diffrax_adaptive_step(rtol=rtol, atol=rtol * 1e-2)
        problem = ode_problem(rhs, state0=state0, dt=0.1, integrator=integrator)
        result = lyapunov_spectrum(
            problem, n_steps=5_000, renorm_every=1, t_transient=100.0)
        errors.append(abs(float(result.exponents[0]) - published_lambda1))
        print(f"rtol={rtol:.0e}  lambda1={float(result.exponents[0]):.4f}  "
              f"error={errors[-1]:.2e}")





.. rst-class:: sphx-glr-script-out

 .. code-block:: none

    rtol=1e-03  lambda1=0.8935  error=1.21e-02
    rtol=1e-04  lambda1=0.9030  error=2.59e-03
    rtol=1e-06  lambda1=0.8946  error=1.10e-02
    rtol=1e-08  lambda1=0.9161  error=1.05e-02
    rtol=1e-10  lambda1=0.9136  error=7.97e-03




.. GENERATED FROM PYTHON SOURCE LINES 75-77

Cross-check: tightest-tolerance adaptive run vs. a fine-dt fixed-step rk4
run of the same system.

.. GENERATED FROM PYTHON SOURCE LINES 77-90

.. code-block:: Python

    rk4_problem = ode_problem(rhs, state0=state0, dt=1e-2, integrator=rk4_step)
    rk4_result = lyapunov_spectrum(
        rk4_problem, n_steps=50_000, renorm_every=10, t_transient=100.0)

    tight_integrator = diffrax_adaptive_step(rtol=1e-10, atol=1e-12)
    adaptive_problem = ode_problem(rhs, state0=state0, dt=0.1, integrator=tight_integrator)
    adaptive_result = lyapunov_spectrum(
        adaptive_problem, n_steps=5_000, renorm_every=1, t_transient=100.0)

    print(f"\nrk4 (dt=1e-2):      lambda1={float(rk4_result.exponents[0]):.4f}")
    print(f"adaptive (rtol=1e-10): lambda1={float(adaptive_result.exponents[0]):.4f}")
    print(f"published:          lambda1={published_lambda1:.4f}")





.. rst-class:: sphx-glr-script-out

 .. code-block:: none


    rk4 (dt=1e-2):      lambda1=0.9017
    adaptive (rtol=1e-10): lambda1=0.9136
    published:          lambda1=0.9056




.. GENERATED FROM PYTHON SOURCE LINES 91-109

.. code-block:: Python

    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4))

    ax1.loglog(rtols, errors, "o-")
    ax1.set_xlabel("rtol")
    ax1.set_ylabel(r"$|\lambda_1 - \mathrm{published}|$")
    ax1.set_title("Adaptive-integrator convergence vs. tolerance")
    ax1.invert_xaxis()

    methods = ["rk4\n(fixed dt=1e-2)", "adaptive\n(rtol=1e-10)", "published"]
    values = [float(rk4_result.exponents[0]), float(adaptive_result.exponents[0]),
              published_lambda1]
    ax2.bar(methods, values, color=["C0", "C1", "C2"])
    ax2.axhline(published_lambda1, color="C2", linestyle="--", linewidth=1)
    ax2.set_ylabel(r"$\lambda_1$")
    ax2.set_title("Cross-check: adaptive vs. fixed-step rk4")
    fig.tight_layout()
    plt.show()




.. image-sg:: /auto_examples/images/sphx_glr_15_adaptive_ode_001.png
   :alt: Adaptive-integrator convergence vs. tolerance, Cross-check: adaptive vs. fixed-step rk4
   :srcset: /auto_examples/images/sphx_glr_15_adaptive_ode_001.png
   :class: sphx-glr-single-img





.. GENERATED FROM PYTHON SOURCE LINES 110-112

Differentiating a Lyapunov exponent through the adaptive integrator: works
with jax.jacfwd (forward-mode), not with jax.grad (reverse-mode).

.. GENERATED FROM PYTHON SOURCE LINES 112-141

.. code-block:: Python



    def lambda_max(rho_param):
        def linear_rhs(y):
            # A small linear system whose exact fundamental matrix (and hence
            # exact Lyapunov exponent, its log-growth rate) is known, so the
            # jacfwd result can be checked against a closed-form derivative
            # rather than only a finite-difference estimate.
            return jnp.array([[-rho_param, 1.0], [-1.0, -rho_param]]) @ y

        integrator = diffrax_adaptive_step(rtol=1e-8, atol=1e-10)
        problem = ode_problem(linear_rhs, state0=jnp.array([1.0, 0.3]), dt=0.2,
                               integrator=integrator)
        result = lyapunov_spectrum(problem, n_steps=50, renorm_every=1, k=1)
        return result.exponents[0]


    a0 = 0.5
    grad_fwd = jax.jacfwd(lambda_max)(a0)
    eps = 1e-5
    finite_diff = (lambda_max(a0 + eps) - lambda_max(a0 - eps)) / (2 * eps)
    print(f"\njax.jacfwd(lambda_max)({a0}) = {float(grad_fwd):.6f}")
    print(f"finite-difference reference   = {float(finite_diff):.6f}")

    try:
        jax.grad(lambda_max)(a0)
    except ValueError as exc:
        print(f"\njax.grad raises (expected): {type(exc).__name__}: "
              f"{str(exc)[:88]}...")




.. rst-class:: sphx-glr-script-out

 .. code-block:: none


    jax.jacfwd(lambda_max)(0.5) = -1.000000
    finite-difference reference   = -1.000000

    jax.grad raises (expected): ValueError: Reverse-mode differentiation does not work for lax.while_loop or lax.fori_loop with dyna...





.. rst-class:: sphx-glr-timing

   **Total running time of the script:** (0 minutes 22.149 seconds)


.. _sphx_glr_download_auto_examples_15_adaptive_ode.py:

.. only:: html

  .. container:: sphx-glr-footer sphx-glr-footer-example

    .. container:: sphx-glr-download sphx-glr-download-jupyter

      :download:`Download Jupyter notebook: 15_adaptive_ode.ipynb <15_adaptive_ode.ipynb>`

    .. container:: sphx-glr-download sphx-glr-download-python

      :download:`Download Python source code: 15_adaptive_ode.py <15_adaptive_ode.py>`

    .. container:: sphx-glr-download sphx-glr-download-zip

      :download:`Download zipped: 15_adaptive_ode.zip <15_adaptive_ode.zip>`


.. only:: html

 .. rst-class:: sphx-glr-signature

    `Gallery generated by Sphinx-Gallery <https://sphinx-gallery.github.io>`_
