import jax import jax.numpy as jnp def compute_grad_norm(grads): """Compute the global L2 norm of gradients.""" leaves = jax.tree_util.tree_leaves(grads) return jnp.sqrt(sum(jnp.sum(g ** 2) for g in leaves)) # In training: loss, grads = nnx.value_and_grad(loss_fn)(model) grad_norm = compute_grad_norm(grads) print(f"Gradient norm: {grad_norm:.4f}") __ __