I want to compute per-example gradients to perform per-example clipping before the final gradient descent step.
I wanted to ensure that the per-example gradients are correct. Therefore I implemented this minimal example below to test standard gradient computation and per-example gradient computation with JAX.
The problem I have is, that the average of the per-example gradients differs from the gradients of the standard computation.
Does someone see where I went wrong?
import jax
import jax.numpy as jnp
from jax import random
def loss(params, x, t):
w, b = params
y = jnp.dot(w, x.T) + b
return ((t.T - y)**2).sum()
def main():
n_samples = 3
dims_in = 7
dims_out = 5
key = random.PRNGKey(0)
# Random data
x = random.normal(key, (n_samples, dims_in), dtype=jnp.float32)
t = random.normal(key, (n_samples, dims_out), dtype=jnp.float32)
# Random weights
w = random.normal(key, (dims_out, dims_in), dtype=jnp.float32)
b = random.normal(key, (dims_out, 1), dtype=jnp.float32)
params = (w, b)
# Standard gradient
reduced_grads = jax.grad(loss)
dw0, db0 = reduced_grads(params, x, t)
print(f"{dw0.shape = }")
print(f"{db0.shape = }")
# Per-example gradients
perex_grads = jax.vmap(jax.grad(loss), in_axes=((None, None), 0, 0))
dw1, db1 = perex_grads(params, x, t)
print(f"{dw1.shape = }")
print(f"{db1.shape = }")
# Gradients are different!
print(jnp.allclose(dw0, jnp.mean(dw1, axis=0))) # should be True
print(jnp.allclose(db0, jnp.mean(db1, axis=0))) # should be True
if __name__ == "__main__":
main()