import jaximport jax.numpy as jnp@jax.jitdef f(x): return x**4 + 3*x**2 + 2*xdf = jax.grad(f)x = 2.0print(f"f(2) = {f(x)}, f'(2) = {df(x)}")# Output: f(2) = 42.0, f'(2) = 58.0# Vectorized computationvdf = jax.vmap(df)x_vec = jnp.array([1.0, 2.0, 3.0])print(f"f'(x) for x=[1,2,3]: {vdf(x_vec)}")# Output: f'(x) for x=[1,2,3]: [10. 58. 154.]