def loss_fn(params, x, y): prediction = params * x return (prediction - y) ** 2 # Returns (loss_value, gradients) loss_and_grad_fn = jax.value_and_grad(loss_fn) params = 1.0 x, y = 2.0, 6.0 # We want params=3.0 so that 3*2=6 loss, grad = loss_and_grad_fn(params, x, y) print(f"Loss: {loss}") # 16.0 (because (1*2 - 6)² = 16) print(f"Grad: {grad}") # -16.0 __ __