@nnx.jit def train_step(model, optimizer, batch): """Single training step.""" def loss_fn(model): logits = model(batch['image']) loss = optax.softmax_cross_entropy_with_integer_labels( logits=logits, labels=batch['label'] ).mean() return loss, logits (loss, logits), grads = nnx.value_and_grad(loss_fn, has_aux=True)(model) optimizer.update(model, grads) accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == batch['label']) return loss, accuracy @nnx.jit def eval_step(model, batch): """Single evaluation step.""" logits = model(batch['image']) loss = optax.softmax_cross_entropy_with_integer_labels( logits=logits, labels=batch['label'] ).mean() accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == batch['label']) return loss, accuracy __ __