Hacker News new | ask | show | jobs
NNX – Neural Networks for JAX (github.com)
2 points by cgarciae 1105 days ago
1 comments

NNX 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.