import jax import jax.numpy as jnp # vmap: Automatic Vectorization def single_fn(x): return x ** 2 batched_fn = jax.vmap(single_fn) results = batched_fn(jnp.array([1, 2, 3])) # With multiple arguments def dot(w, x): return jnp.dot(w, x) # Shared weights, batched inputs batch_dot = jax.vmap(dot, in_axes=(None, 0)) # grad: Automatic Differentiation def loss(params): return params ** 2 grad_fn = jax.grad(loss) gradient = grad_fn(3.0) # 6.0 # Get both value and gradient loss_val, grad_val = jax.value_and_grad(loss)(3.0) # Gradient with respect to specific argument def f(x, y): return x * y df_dy = jax.grad(f, argnums=1) # Combining Transforms fast_batched_grad = jax.jit(jax.vmap(jax.grad(loss_fn), in_axes=(None, 0))) __ __