(params, batch)
| 162 | |
| 163 | |
| 164 | def accuracy(params, batch): |
| 165 | inputs, targets = batch |
| 166 | target_class = jnp.argmax(targets, axis=1) |
| 167 | predicted_class = jnp.argmax(predict(params, inputs), axis=1) |
| 168 | return jnp.mean(predicted_class == target_class) |
| 169 | |
| 170 | |
| 171 | def eval_fn(params): |