@jax.jit def train_step(state, batch): def loss_fn(params): logits = state.apply_fn({'params': params}, batch['images']) labels = jax.nn.one_hot(batch['labels'], num_classes=10) loss = -jnp.sum(labels * jax.nn.log_softmax(logits)) / batch['labels'].shape[0] return loss # Compute gradients grads = jax.grad(loss_fn)(state.params) # Update model state state = state.apply_gradients(grads=grads) return state __ __