NNX – Neural Networks for JAX
github.com
NNX – Neural Networks for JAX
1–2 of 2 posts
Re: NNX – Neural Networks for JAX
#2NNX is a Neural Networks library for JAX that provides a simple yet powerful module system that adheres to standard Python semantics. Its aim is to combine the robustness of Flax with a simplified, Pythonic API akin to that of PyTorch.