def compute_accuracy(params, state, batch): logits = state.apply_fn({'params': params}, batch['images']) predicted_labels = jnp.argmax(logits, axis=1) return jnp.mean(predicted_labels == batch['labels']) __ __