jax>=0.6
optax
scikit-learn
scipy
numpy
pandas
matplotlib
