def messy_function(x): return jnp.sin(x) * jnp.exp(-x ** 2) + jnp.tanh(x) gradient_fn = jax.grad(messy_function) # The gradient at x=1.0 print(gradient_fn(1.0)) # -0.5047... __ __