model.eval() # Ensure evaluation mode (deterministic) @nnx.jit def predict(model, images): logits = model(images) return jnp.argmax(logits, axis=-1) # Get predictions test_batch = next(test_ds.as_numpy_iterator()) predictions = predict(model, jnp.array(test_batch['image'])) print(f"Predictions: {predictions[:10]}") print(f"Actual: {test_batch['label'][:10]}") __ __