def loss_single(params, x, y): """Loss for a single data point.""" pred = params[0] * x + params[1] # Linear: y = mx + b return (pred - y) ** 2 # Stack the transforms: # 1. grad: compute gradients with respect to params # 2. vmap: do this for a batch of (x, y) pairs # 3. jit: compile the whole thing batched_grad_fn = jax.jit( jax.vmap( jax.grad(loss_single), in_axes=(None, 0, 0) # Same params, batch of x, batch of y ) ) params = jnp.array([1.0, 0.0]) # Initial guess: y = 1*x + 0 x_batch = jnp.array([1.0, 2.0, 3.0]) y_batch = jnp.array([2.0, 4.0, 6.0]) # True relationship: y = 2x # Get gradients for each example in the batch grads_per_example = batched_grad_fn(params, x_batch, y_batch) print("Gradients per example:") print(grads_per_example) # Average them for a batch gradient batch_grad = jnp.mean(grads_per_example, axis=0) print(f"Batch gradient: {batch_grad}") __ __