import jaximport jax.numpy as jnpdef f(x): return jnp.array([x**2, x**3])x = 2.0y, vjp_fn = jax.vjp(f, x)print(f"VJP: {vjp_fn(jnp.array([1.0, 1.0]))[0]}")primal, jvp_fn = jax.jvp(f, (x,), (1.0,))print(f"JVP: {jvp_fn}")# Output:# VJP: 16.0# JVP: [4. 12.]