import jaximport jax.numpy as jnpdef f(x): return x**4 + 3*x**2 + 2*xddf = jax.grad(jax.grad(f))x = 2.0print(f"f''(2) = {ddf(x)}")# Output: f''(2) = 102.0