from flax import nnx import optax import chex @nnx.jit def bulletproof_train_step(model, optimizer, x, y): def loss_fn(m): logits = m(x) loss = optax.softmax_cross_entropy_with_integer_labels(logits, y).mean() return loss loss, grads = nnx.value_and_grad(loss_fn)(model) # The Firewall: check before we let bad gradients touch the model chex.assert_tree_all_finite(grads) optimizer.update(model, grads) return loss __ __