from flax import nnx import optax @nnx.jit def train_step(model, optimizer, batch): 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) # Peek at real numbers, right inside the compiled step jax.debug.print("loss = {loss}", loss=loss) jax.debug.print("max logit = {m}", m=jnp.max(logits)) optimizer.update(model, grads) return loss __ __