yax — des réseaux de neurones au plus près de jax¶
yax sert à écrire des réseaux de neurones avec jax, sans concepts ajoutés (bio).
Un modèle est une classe Python dont les
tableaux sont les paramètres : jax.grad, jax.jit, jax.vmap et la fameuse lib d'optimisation optax
s'appliquent sur les modèles directement. Si on le souhaite, une fonction yax.training.train fait la
boucle d'entraînement et sauvegarde le meilleur modèle.
pip install yaxlib
En trente secondes¶
Une sinusoïde bruitée, un perceptron multicouche, un entraînement :
import jax.numpy as jnp
import jax.random as jr
import yax
X = jr.uniform(jr.key(0), (256, 1), minval=-3.0, maxval=3.0)
Y = jnp.sin(X) + 0.1 * jr.normal(jr.key(1), X.shape)
model = yax.layers.MLP(layer_sizes=(1, 32, 32, 1), activation="tanh", rkey=jr.key(2))
config = yax.configs.TrainConfig(learning_rate=1e-2, nb_epochs=100)
run = yax.training.train("out/sinus", config,
yax.obj.mse, # l'objectif à minimiser
model,
(X[:200], Y[:200], 32), # entraînement, par batchs de 32
(X[200:], Y[200:])) # validation
print(run.loss) # perte de validation du meilleur modèle
Y_pred = yax.batch_apply(run.trained_model, X)
Le dossier out/sinus/0 contient l'historique de l'entrainement ainsi que le modèle entrainé, sauvegardé à sa meilleure validation.
Dans une autre session, yax.training.load_run("out/sinus/0") permet de retrouver tout cela.
Un second entrainement (en variant la config par exemple) s'enregistrera dans "out/sinus/1" etc.
Ce que yax vous donne¶
- Un modèle est un pytree. Déclarez une classe qui hérite de
yax.Module: ses champs sont les paramètres entrainables ou statiques. Rien à envelopper, rien à filtrer :jax.grad(loss)(model)rend un modèle de même forme dont les feuilles sont les gradients. - Une seule signature, partout.
apply(x,rkey)traite un exemplex. Ensuitejax.vmappermet de traiter un batch;rkeyest la source d'aléatoire quand il y en a (dropout, tirages). Le mode d'évaluation se bascule avec la méthode:model.set_inference(True). - Un entraînement qui laisse des traces. Chaque appel à
yax.training.traincrée un dossier numéroté avec le meilleur modèle, l'état de l'optimiseur, la configuration et l'historique. Les runs d'un même dossier se comparent, etyax.training.find_best_rundésigne le meilleur au sein dumother_folder. - Des briques prêtes à l'emploi, toutes écrites comme vous les écririez : couches (MLP, convolutions n-dimensionnelles, GRU et LSTM, attention multi-têtes, blocs transformer, message passing), modèles complets (U-Net, opérateur de Fourier, VAE, flot normalisant, diffusion, mini-YOLO), fonctions de pertes, optimiseurs, prétraitements de données.
- Des mini-fonctionnalités sympathiques. Essayez par exemple
yax.ipprint(model)dans un notebook.
Pourquoi pas flax ou equinox ?¶
Par simplicité. Chaque framework se mesure aux concepts qu'il ajoute à jax : flax apporte ses collections de variables, ses scopes et son cycle init/apply ; equinox réduit cela à des modules-pytrees, mais y ajoute sa machinerie de filtrage. yax ajoute le strict minimum pour construire des modèles à base de layers imbriqués. Tout ce qui n'est pas dans yax se code en jax ordinaire, sans friction.
Pour continuer¶
- Prise en main : un quart d'heure, des données, un modèle, un entraînement, un rechargement.
- Aller plus loin : sous le capot du modèle, l'aléa et le mode inférence, les samplers, la configuration de l'entraînement, les pertes à soi.
- Référence de l'API et démos : une démonstration par famille de modèles, qui converge en quelques secondes sur CPU.
Développement¶
pip install yaxlib[dev] # ajoute pytest
pytest tests/