Prise en main¶
En un quart d'heure : des données, un modèle, un entraînement, et le modèle entraîné qu'on retrouve plus tard sur le disque. Exécutez les cellules au fil de la lecture, dans Jupyter ou sur Colab.
!pip install -qq -U yaxlib==0.3.*
import jax
import jax.numpy as jnp
import jax.random as jr
import matplotlib.pyplot as plt
import yax
Des données à séparer¶
Les données sont des tableaux jax: X de forme
(nb, 2) et Y de forme (nb, 1). Deux classes dans le plan, un nuage
central et un anneau autour, à séparer.
def nuages(rkey, nb):
"""Deux classes dans le plan : un nuage central (0) et un anneau autour (1)."""
rkey_r, rkey_theta, rkey_c = jr.split(rkey, 3)
classe = jr.bernoulli(rkey_c, 0.5, (nb,))
rayon = jnp.where(classe, 2.0, 0.0) + 0.5 * jr.normal(rkey_r, (nb,))
theta = jr.uniform(rkey_theta, (nb,), maxval=2 * jnp.pi)
X = jnp.stack([rayon * jnp.cos(theta), rayon * jnp.sin(theta)], axis=1)
return X, classe.astype(jnp.float32)[:, None]
X_train, Y_train = nuages(jr.key(1), 1000)
X_val, Y_val = nuages(jr.key(2), 400)
plt.scatter(X_train[:, 0], X_train[:, 1], c=Y_train[:, 0], s=6, cmap="coolwarm")
plt.axis("equal"); plt.show()
Un modèle¶
Un modèle yax est une classe qui hérite de yax.Module. On y déclare ses
briques (ici un perceptron multicouche et une couche linéaire de sortie), on les
construit dans __init__ à partir d'une clé aléatoire. La méthode apply calcule la
sortie pour un exemple, un vecteur de taille 2. La sortie est un logit : après entrainement on espère qu'il sera positif pour la classe 1, négatif pour la classe 0.
class Classifieur(yax.Module):
mlp: yax.layers.MLP
tete: yax.layers.Linear
def __init__(self, largeur, rkey):
rkey_mlp, rkey_tete = jr.split(rkey) # une clé par brique qui tire des poids
self.mlp = yax.layers.MLP(layer_sizes=(2, largeur, largeur), activation="relu", rkey=rkey_mlp)
self.tete = yax.layers.Linear(dim_in=largeur, dim_out=1, rkey=rkey_tete)
def apply(self, x, rkey=None):
h = jax.nn.relu(self.mlp.apply(x))
return self.tete.apply(h) # un logit
model = Classifieur(largeur=32, rkey=jr.key(0))
print("un exemple :", model.apply(X_train[0]))
un exemple : [-0.01391045]
Le modèle s'affiche comme un arbre à déplier :
yax.ipprint(model)
Classifieur 1185 paramètres
mlp : MLP 1152 paramètres
layers : list[2]
0 : Linear 96 paramètres
weight : f32[2,32]
[[-0.3598 -0.1155 -0.0539 -0.238 0.1046 -0.116 0.1369 -0.0541 -0.0336 -0.1744 -0.2447 -0.1779 -0.0527 0.1882 0.2428 0.0453 0.1297 0.1925 0.2246 0.2974 0.3839 0.326 0.3693 -0.0565 -0.2503 0.3227 0.2745 0.0161 -0.0297 0.2348 -0.3431 -0.3409] [ 0.1912 -0.2085 -0.3039 -0.0864 0.0527 0.0101 -0.0286 0.2227 -0.2835 -0.1947 0.3191 0.126 0.3252 -0.3635 0.3801 -0.3063 0.2029 -0.3998 0.3362 -0.2598 0.2689 -0.1004 0.201 -0.2597 -0.1778 -0.142 -0.2199 0.3768 0.1617 0.2701 -0.3977 0.3151]]
bias : f32[32]
[0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0.]
1 : Linear 1056 paramètres
weight : f32[32,32]
[[-0.2962 0.2069 -0.0855 ... -0.1299 -0.2927 -0.1062] [-0.1544 -0.0333 0.1655 ... 0.2068 -0.2814 0.2789] [ 0.0989 0.1529 0.1433 ... 0.0693 0.2259 0.1739] ... [-0.2644 -0.2933 -0.0032 ... -0.0844 0.258 0.0558] [ 0.1734 0.1652 -0.1152 ... 0.2895 -0.0998 0.1274] [ 0.0689 -0.1334 0.0383 ... -0.1093 -0.1052 -0.1349]]
bias : f32[32]
[0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0. 0.]
tete : Linear 33 paramètres
weight : f32[32,1]
[[ 0.0672] [-0.2383] [-0.3377] [-0.1359] [-0.2923] [ 0.0962] [ 0.29 ] [ 0.1628] [-0.2314] [ 0.3178] [-0.0033] [ 0.1481] [-0.0998] [-0.1662] [ 0.0929] [ 0.2327] [-0.0692] [-0.2567] [ 0.286 ] [ 0.1744] [-0.281 ] [ 0.269 ] [ 0.179 ] [-0.4113] [-0.3665] [ 0.0177] [ 0.3771] [-0.2671] [-0.0526] [-0.0646] [ 0.1371] [-0.005 ]]
bias : f32[1]
[0.]
apply traite un exemple à la fois. Pour un lot d'exemples, jax.vmap
applique la même fonction à chaque ligne :
logits = jax.vmap(model.apply)(X_train[:5])
print(logits.shape)
(5, 1)
Entraîner¶
yax.training.train lance la boucle d'entraînement :
- mélange les données et crée les batchs
- régle le pas d'optimisation (learning rate)
- valide à chaque époque, et conserve le meilleur modèle.
Il lui faut, dans cet ordre :
- un dossier où enregistrer (
mother_folder) ; - une configuration : valeur et décroissante du learning rate, nombre d'époques ;
- l'objectif, c'est-à-dire la perte à minimiser — ici une classification
binaire sur un logit :
yax.obj.bce; - le modèle ;
- les données d'entraînement
(X, Y, batch_size)et de validation(X, Y);
Il rend un run, dont run.history.plot() trace les pertes.
config = yax.configs.TrainConfig(learning_rate=1e-2, nb_epochs=40)
run = yax.training.train("out/anneau", config, yax.obj.bce, model,
(X_train, Y_train, 50), # entraînement, par batchs de 50
(X_val, Y_val) # validation
)
print("enregistré dans :", run.folder)
run.history.plot(title=f"perte de validation du meilleur modèle : {run.loss:.3f}")
plt.show()
enregistré dans : out/anneau/2
⇑ En rouge, la perte d'entraînement, un point par pas et sa moyenne par époque ; en bleu, la perte de validation à la fin de chaque époque ; en vert, les époques où elle a battu son record. Le modèle gardé est celui du dernier record.
Utiliser le modèle entraîné¶
run.trained_model est le modèle ayant eu la plus faible perte de validation.
yax.batch_apply l'applique à tout un tableau d'exemples.
logits = yax.batch_apply(run.trained_model, X_val, None)
print("justesse sur la validation :", float(jnp.mean((logits > 0) == (Y_val > 0.5))))
# la probabilité de la classe 1 sur tout le plan
gx, gy = jnp.meshgrid(jnp.linspace(-3.5, 3.5, 120), jnp.linspace(-3.5, 3.5, 120))
grille = jnp.stack([gx.ravel(), gy.ravel()], axis=1)
proba = jax.nn.sigmoid(yax.batch_apply(run.trained_model, grille, None))[:, 0]
plt.contourf(gx, gy, proba.reshape(gx.shape), levels=20, cmap="coolwarm", alpha=0.5)
plt.scatter(X_val[:, 0], X_val[:, 1], c=Y_val[:, 0], s=6, cmap="coolwarm")
plt.axis("equal"); plt.show()
justesse sur la validation : 0.98499995470047
Retrouver le modèle plus tard¶
Tout est sur le disque : le modèle, la configuration, l'historique. Chaque
entraînement crée un sous-dossier numéroté de out/anneau — 0 pour le
premier, 1 pour le suivant, etc. ; le chemin a été affiché à la fin de
l'entraînement.
Imaginons qu'on revienne demain, dans une nouvelle session : la variable run
n'existe plus, il ne reste que le dossier. On le désigne par son chemin, et
yax.training.load_run relit le tout. Seule condition : la classe
Classifieur doit être définie dans cette nouvelle session (c'est le principe de pickle).
recharge = yax.training.load_run("out/anneau/0")
modele = recharge.trained_model
print("sa configuration :", recharge.config)
print("sa perte de validation :", recharge.loss)
sa configuration : TrainConfig(learning_rate=0.01, nb_epochs=40, lr_final_ratio=1.0, optimizer='adam', optimizer_options=None, patience=None) sa perte de validation : 0.03834846243262291
Le modèle rechargé est bien celui qu'on avait entraîné : il rend exactement les mêmes sorties.
logits_recharge = yax.batch_apply(modele, X_val, None)
print("rechargé identique :", bool(jnp.all(logits_recharge == logits)))
rechargé identique : True
Et ensuite ?¶
- Le tutoriel Aller plus loin ouvre le capot : ce qu'est vraiment un modèle yax, l'aléa et le mode d'évaluation, les différentes façons de fournir les données, les options de l'entraînement.
- La référence de l'API liste les couches (yax.layers) et les modèles complets (yax.models) fournis.
- Les démos montrent une famille de modèles par script.