def f(x, y): return x ** 2 + x * y # Gradient with respect to x (first argument, default) df_dx = jax.grad(f, argnums=0) # Gradient with respect to y (second argument) df_dy = jax.grad(f, argnums=1) # Gradients with respect to both df_both = jax.grad(f, argnums=(0, 1)) x, y = 2.0, 3.0 print(f"df/dx: {df_dx(x, y)}") # 2*2 + 3 = 7 print(f"df/dy: {df_dy(x, y)}") # 2 print(f"Both: {df_both(x, y)}") # (7.0, 2.0) __ __