import jax.nn as jnn def cross_entropy_loss(params, state, batch): logits = state.apply_fn({'params': params}, batch['images']) labels = jnn.one_hot(batch['labels'], num_classes=10) return -jnp.sum(labels * jnn.log_softmax(logits)) / batch['labels'].shape[0] __ __