@nnx.jit def train_step(model, optimizer, x, y): """Execute one training step.""" def loss_fn(model): logits = model(x) loss = optax.softmax_cross_entropy_with_integer_labels( logits=logits, labels=y ).mean() return loss, logits # Compute loss and gradients (loss, logits), grads = nnx.value_and_grad(loss_fn, has_aux=True)(model) # Update parameters optimizer.update(model, grads) # Compute accuracy for logging predictions = jnp.argmax(logits, axis=-1) accuracy = jnp.mean(predictions == y) return loss, accuracy __ __