def complex_operation(x): """A function with multiple steps.""" a = jnp.sin(x) b = jnp.exp(-x ** 2) c = a * b + jnp.log(1 + jnp.abs(x)) return c # Without vmap, you'd need to think about broadcasting at each step # With vmap, you just wrap it batched_complex = jax.vmap(complex_operation) x_batch = jnp.linspace(-3, 3, 1000) results = batched_complex(x_batch) __ __