def evaluate_model(state, test_loader): accuracies = [] for batch in test_loader: acc = compute_accuracy(state.params, state, batch) accuracies.append(acc) final_accuracy = jnp.mean(jnp.array(accuracies)) print(f"Test Accuracy: {final_accuracy * 100:.2f}%") return final_accuracy __ __