import chex @chex.assert_max_traces(n=2) @nnx.jit def strict_train_step(model, optimizer, batch): # ... training logic ... return loss __ __