def loss_fn(model, x, y): logits = model(x) loss = compute_loss(logits, y) return loss # Get both the loss value and the gradients loss, grads = nnx.value_and_grad(loss_fn)(model, x, y) __ __