pip
numpy

[jax]
flax

[pytorch]
torch>=2.4.0

[tensorflow]
tensorflow>=2.18.1

[torax]
torax
epednn_mit[jax]
