
Tech • IA • Robotique
L’intégration récente de Keras avec le système modulaire Flax et NNX ouvre de nouvelles perspectives, combinant la simplicité de Keras avec la puissance et la flexibilité de JAX pour la gestion avancée des variables et les entraînements personnalisés.
La nouvelle liaison fait de keras.Variable une instance de nnx.Variable, ce qui permet une coexistence harmonieuse des états entre Keras et NNX. Cette compatibilité facilite le mélange des composants issus des deux écosystèmes sans rupture dans la gestion des variables.
Pour bénéficier de cette intégration, il faut configurer deux variables d’environnement avant d’importer Keras: définir le backend sur JAX et activer explicitement le mode NNX. Cette étape assure que Keras utilise les capacités avancées de NNX et Flax.
Le système permet de vérifier que les variables Keras sont bien reconnues et suivies par NNX: elles apparaissent dans la liste des variables gérées, disposent d’un état de trace pour la compilation juste-à-temps et peuvent être accédées directement via NNX, renforçant le contrôle lors des phases d’entraînement.
NNX offre une approche modulaire avec des classes Python standard, introduisant une modularité native dans Keras grâce à l’intégration. Par exemple, un module NNX peut contenir des couches linéaires, des variables personnalisées Keras et manipuler leur interaction dans la méthode d’appel du modèle.
Le premier mode correspond à la méthode classique de Keras: model.compile et model.fit fonctionnent normalement, avec NNX et JAX en coulisses pour optimiser les performances, sans nécessiter de refonte du code.
Le second mode permet d’exécuter des boucles d’entraînement personnalisées en traitant un modèle Keras comme un nnx.Module. Cela donne accès au vaste écosystème JAX, y compris Optax pour les optimisateurs, et aux fonctions nnx.jit et nnx.grad pour accélérer et différencier efficacement les modèles.
L’utilisation du décorateur nnx.jit améliore les performances en compilant à la volée les fonctions d’entraînement, ce qui donne un gain notable en vitesse même sur des modèles Keras intégrés à NNX.
Cette symbiose permet de débuter facilement avec les API familières de Keras tout en accédant progressivement à des fonctionnalités avancées de recherche et optimisation propres à JAX, créant ainsi un cadre unifié pour tous les niveaux d’expertise.
Les modèles Keras ainsi intégrés peuvent exploiter pleinement les outils JAX comme Optax, un large éventail de fonctions statistiques et différentielles, étendant considérablement les possibilités au-delà du cadre Keras traditionnel.
Un guide complet et des exemples de code sont proposés sur le site keras.io, offrant un point de départ clair pour maîtriser cette nouvelle approche et expérimenter concrètement cette intégration.
L’intégration de Keras avec Flax et NNX révolutionne la gestion d’état et les possibilités d’entraînement, offrant à la fois simplicité d’utilisation et puissance de personnalisation dans l’écosystème JAX. Ce pont ouvre la voie à des développements plus flexibles et performants en machine learning.
Explique-moi