Aller plus loin¶
Vous avez suivi la prise en main. Ce tutoriel suit un
seul problème, une régression en dimension 1, à travers plusieurs
entraînements : on change la fonction d'activation, on combat le
sur-apprentissage, on distribue les données autrement, puis on compare tous
les runs sur un même graphique. En chemin, on ouvre le capot : ce qu'est un
modèle yax, comment l'aléa est géré, et tout ce que yax.training.train
accepte.
Télécharger ce notebook. Pour l'ouvrir sur Colab: Fichier → Importer le notebook.
!pip install -qq -U yaxlib==0.3.*
import shutil
import jax
import jax.numpy as jnp
import jax.random as jr
import optax
import matplotlib.pyplot as plt
import yax
Le problème : retrouver une courbe sous le bruit¶
Une grandeur suit une loi $y = f(x)$ que l'on ne connaît pas. On n'en a que 20 mesures, entachées d'un bruit d'écart-type $0.3$. Deux cents autres mesures forment le jeu de validation, qui jugera chaque modèle.
def f(x):
"""La loi à retrouver (inconnue en pratique)."""
return 0.8 * jnp.sin(4 * x) + 0.3 * x
def mesure(rkey, nb):
"""nb mesures bruitées, en des points tirés au hasard dans [-1, 1]."""
rkey_x, rkey_bruit = jr.split(rkey)
X = jr.uniform(rkey_x, (nb, 1), minval=-1, maxval=1)
return X, f(X) + 0.3 * jr.normal(rkey_bruit, (nb, 1))
X_train, Y_train = mesure(jr.key(1), 20)
X_val, Y_val = mesure(jr.key(2), 200)
plancher = float(jnp.mean((f(X_val) - Y_val) ** 2))
print(f"erreur quadratique de la vraie loi f sur la validation : {plancher:.3f}")
erreur quadratique de la vraie loi f sur la validation : 0.109
⇑ Même la vraie loi fait une erreur de 0.109, proche de $0.3^2 = 0.09$ : c'est le plancher que le bruit impose. Aucun modèle ne descendra nettement dessous.
x_grille = jnp.linspace(-1, 1, 200)[:, None]
def trace_donnees(ax):
ax.plot(X_val, Y_val, ".", color="lightgray", label="validation")
ax.plot(x_grille, f(x_grille), "k--", lw=1, label="f (inconnue)")
ax.plot(X_train, Y_train, "o", color="tab:red", ms=4, label="entraînement")
_, ax = plt.subplots(figsize=(7, 3.5))
trace_donnees(ax)
ax.legend()
plt.show()
Un modèle est un pytree¶
Un yax.Module est un pytree jax dont les feuilles sont exactement les
paramètres. Ses champs sont de deux familles :
- les champs dynamiques (annotation seule) : des tableaux jax, des sous-modules, ou des listes, tuples et dictionnaires de ceux-ci — et rien d'autre. Ce sont les paramètres, ceux que le gradient et l'optimiseur voient ;
- les champs statiques (
= yax.StaticField()) : tout le reste — tailles, options, noms — rangé dans la structure du pytree, invisible pour le gradient. Un flottant ou une chaîne glissés dans un champ dynamique sont refusés à la construction, avec un message explicite.
Notre régresseur enchaîne un MLP, une activation, un dropout et une tête
linéaire. L'activation est donnée par son nom, pris dans le dictionnaire
yax.activations.ACTIVATIONS ; ce nom est un champ statique, enregistré avec
le modèle.
class Regresseur(yax.Module):
corps: yax.layers.MLP # trois champs dynamiques : des sous-modules
dropout: yax.layers.Dropout
tete: yax.layers.Linear
activation: str = yax.StaticField() # un champ statique : un nom
def __init__(self, largeur, activation, rate, rkey):
rkey_corps, rkey_tete = jr.split(rkey)
self.corps = yax.layers.MLP(layer_sizes=(1, largeur, largeur), activation=activation, rkey=rkey_corps)
self.dropout = yax.layers.Dropout(rate)
self.tete = yax.layers.Linear(dim_in=largeur, dim_out=1, rkey=rkey_tete)
self.activation = activation
def apply(self, x, rkey=None):
h = self.corps.apply(x) # le MLP n'active pas sa sortie :
h = yax.activations.ACTIVATIONS[self.activation](h) # on le fait ici
h = self.dropout.apply(h, rkey)
return self.tete.apply(h)
model = Regresseur(largeur=128, activation="tanh", rate=0.1, rkey=jr.key(0))
yax.ipprint(model)
Regresseur 16897 paramètres
corps : MLP 16768 paramètres
layers : list[2]
0 : Linear 256 paramètres
weight : f32[1,128]
[[-0.1847 -0.0593 -0.0277 -0.1222 0.0537 -0.0595 0.0703 -0.0278 -0.0172 -0.0895 -0.1256 -0.0913 -0.027 0.0966 0.1247 0.0233 0.0666 0.0988 0.1153 0.1527 0.1971 0.1674 0.1896 -0.029 -0.1285 0.1657 0.1409 0.0082 -0.0153 0.1206 -0.1762 -0.175 0.0981 -0.107 -0.156 -0.0444 0.0271 0.0052 -0.0147 0.1144 -0.1456 -0.1 0.1638 0.0647 0.167 -0.1866 0.1951 -0.1573 0.1042 -0.2052 0.1726 -0.1334 0.1381 -0.0515 0.1032 -0.1333 -0.0913 -0.0729 -0.1129 0.1935 0.083 0.1387 -0.2042 0.1618 0.0837 -0.003 -0.03 -0.0371 0.0174 -0.0317 -0.1031 -0.0413 0.0358 0.0357 0.1515 0.0153 0.0941 -0.0719 -0.1387 -0.1768 -0.0604 0.1408 -0.0406 -0.0401 -0.0449 0.146 -0.0365 -0.1596 0.108 -0.1838 -0.1821 0.0665 0.1469 -0.0241 0.1036 -0.0836 0.1262 0.178 -0.1711 0.1667 -0.0997 0.013 0.0681 0.1506 -0.1242 -0.1459 -0.0297 -0.0367 -0.1431 -0.1942 0.1773 -0.1174 -0.0573 -0.1695 -0.09 -0.1882 -0.0464 0.1937 -0.0137 0.1915 -0.206 0.1211 -0.1458 -0.2147 -0.0418 -0.1354 0.124 0.1593]]
bias : f32[128]
[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. 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. 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. 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 16512 paramètres
weight : f32[128,128]
[[-0.1481 0.1034 -0.0428 ... 0.1049 0.1021 0.028 ] [ 0.0461 0.0036 0.0081 ... -0.0189 -0.0314 0.1091] [-0.064 0.1061 -0.0791 ... 0.095 0.0127 0.1138] ... [ 0.1144 0.1395 0.1072 ... -0.0636 -0.0423 0.1445] [-0.1378 0.0309 0.1028 ... 0.121 -0.0612 -0.0214] [-0.0654 0.061 -0.1013 ... 0.1319 0.1301 0.1332]]
bias : f32[128]
[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. 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. 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. 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.]
dropout : Dropout 0 paramètre
tete : Linear 129 paramètres
weight : f32[128,1]
[[ 0.034 ] [-0.1205] [-0.1708] [-0.0687] [-0.1478] [ 0.0486] [ 0.1467] [ 0.0823] [-0.1171] [ 0.1607] [-0.0016] [ 0.0749] [-0.0505] [-0.084 ] [ 0.047 ] [ 0.1177] [-0.035 ] [-0.1298] [ 0.1446] [ 0.0882] [-0.1421] [ 0.136 ] [ 0.0905] [-0.208 ] [-0.1854] [ 0.0089] [ 0.1908] [-0.1351] [-0.0266] [-0.0327] [ 0.0694] [-0.0025] [-0.2083] [ 0.1278] [ 0.1696] [ 0.1387] [ 0.1978] [-0.0663] [ 0.0599] [-0.0008] [ 0.1946] [ 0.2094] [ 0.0226] [ 0.0673] [-0.1341] [-0.0407] [-0.0722] [ 0.0708] [-0.1131] [ 0.2059] [ 0.0618] [ 0.0893] [-0.0046] [ 0.1609] [-0.0552] [ 0.0577] [ 0.1605] [ 0.1215] [ 0.1624] [ 0.2013] [-0.0547] [ 0.1302] [-0.1493] [-0.2119] [ 0.1893] [-0.1942] [ 0.1861] [ 0.2123] [-0.0231] [ 0.0875] [-0.1584] [ 0.1448] [-0.2152] [-0.047 ] [ 0.0818] [ 0.0068] [ 0.1914] [-0.1118] [ 0.1212] [-0.0572] [-0.0072] [-0.0827] [-0.1799] [ 0.1691] [-0.1037] [-0.024 ] [ 0.1054] [ 0.2063] [-0.038 ] [-0.0523] [ 0.0678] [ 0.0347] [-0.0508] [ 0.1432] [ 0.1105] [ 0.0167] [-0.1634] [ 0.1852] [ 0.0443] [ 0.0785] [ 0.1259] [ 0.1258] [-0.2146] [-0.1389] [ 0.0268] [ 0.071 ] [-0.0699] [ 0.1569] [-0.0123] [ 0.1087] [ 0.135 ] [-0.1522] [-0.0901] [ 0.1164] [ 0.1082] [-0.0814] [ 0.1949] [ 0.1807] [-0.0172] [-0.0756] [ 0.1298] [ 0.1016] [-0.1008] [ 0.136 ] [-0.0755] [ 0.1381] [ 0.1197] [ 0.2104]]
bias : f32[1]
[0.]
⇑ Les tableaux (f32[1,128]…) sont les feuilles ; activation, rate,
layer_sizes ou inference sont des statiques, rangés dans la structure. Le
nombre de paramètres se compte donc avec l'outil de base des pytrees :
print("paramètres :", sum(p.size for p in jax.tree.leaves(model)))
paramètres : 16897
L'aléa et le mode inférence¶
Deux choses distinctes, et découplées :
- l'aléa passe par l'argument
rkeydeapply. Le dropout en a besoin pour tirer son masque ; en mode entraînement, l'oublier est une erreur explicite ; - le mode est le champ statique
inference, porté par tout module et basculé parmodel.set_inference(True), qui rend une copie. En inférence, le dropout devient l'identité et la clé n'a plus d'effet.
yax.training.train gère les deux : il entraîne en inference=False avec une
clé par pas, valide en inference=True, et rend le meilleur modèle en
inference=True, prêt à évaluer.
x = jnp.array([0.5])
model_eval = model.set_inference(True)
try:
model.apply(x) # entraînement, sans clé
except ValueError as e:
print("sans clé, mode entraînement :", str(e)[:60], "...")
print("avec clé, deux tirages :", model.apply(x, jr.key(1)), model.apply(x, jr.key(2)))
print("inférence, avec ou sans clé :", model_eval.apply(x), model_eval.apply(x, jr.key(1)))
sans clé, mode entraînement : Dropout en mode entrainement (inference=False) : une cle est ... avec clé, deux tirages : [-0.01956649] [-0.05608208] inférence, avec ou sans clé : [-0.03535716] [-0.03535716]
Un premier entraînement¶
La configuration¶
yax.configs.TrainConfig rassemble les réglages de l'optimisation, et rien
d'autre : la taille des lots appartient aux données. La configuration est
enregistrée avec chaque run ; ses champs sont donc des données simples.
| Attribut | Rôle | Par défaut |
|---|---|---|
learning_rate |
pas d'apprentissage (learning rate) initial | obligatoire |
nb_epochs |
nombre d'époques (epochs) : de passages sur les données | obligatoire |
lr_final_ratio |
rapport du pas final au pas initial ; entre les deux, le pas décroît selon un cosinus | 1.0 : pas constant |
optimizer |
le nom d'un optimiseur de yax.optimizers |
"adam" |
optimizer_options |
ses réglages, en dictionnaire : {"weight_decay": 0.1} pour "adamw", {"momentum": 0.9} pour "sgd" |
None |
patience |
arrête l'entraînement après ce nombre d'époques sans record de validation (early stopping) | None : jusqu'au bout |
Les noms d'optimiseurs acceptés :
print(list(yax.optimizers.OPTIMIZERS))
['adam', 'adamw', 'adamax', 'adabelief', 'adagrad', 'nadam', 'radam', 'rmsprop', 'lion', 'sgd', 'lbfgs']
Un dossier pour toute l'expérience¶
Tous les runs de ce tutoriel iront dans le même mother_folder : ils
partagent le jeu de validation et la perte, leurs pertes sont donc
comparables. On le vide d'abord, pour que la comparaison finale ne porte que
sur les runs de cette session.
Chaque run reçoit un titre, l'argument title de yax.training.train. Il
est enregistré avec le run et se retrouve dans run.title : il servira à
légender les courbes. La petite fonction entraine évite de répéter les
arguments communs à tous les runs.
mother_folder = "out/regression"
shutil.rmtree(mother_folder, ignore_errors=True) # on repart d'un dossier vide
def entraine(title, model, config, training):
run = yax.training.train(mother_folder, config, yax.obj.mse, model, training, (X_val, Y_val),
rkey=jr.key(3), title=title)
print(f"{run.title} : perte {run.loss:.4f} après {len(run.history.val_losses)} époques, dans {run.folder}")
return run
Premier essai : activation tanh, pas de dropout, toutes les données en un
seul lot. L'objectif à minimiser, le 3e argument de yax.training.train,
est l'erreur quadratique moyenne yax.obj.mse.
config = yax.configs.TrainConfig(learning_rate=1e-2, nb_epochs=1000)
run_tanh = entraine("tanh", Regresseur(largeur=128, activation="tanh", rate=0.0, rkey=jr.key(0)),
config, (X_train, Y_train))
tanh : perte 0.1343 après 1000 époques, dans out/regression/0
run_tanh.history.plot(title=run_tanh.title)
plt.show()
⇑ La perte d'entraînement ne cesse de baisser. La validation, elle, atteint son minimum vers le pas 200, puis remonte jusqu'à 0.35 : le modèle apprend par cœur le bruit des 20 mesures. C'est le sur-apprentissage (overfitting). Le modèle rendu est le meilleur (gros point vert), pas le dernier.
def trace_predictions(ax, run):
trace_donnees(ax)
ax.plot(x_grille, yax.batch_apply(run.trained_model, x_grille), color="tab:blue", lw=2, label="modèle")
ax.set_title(f"{run.title} — perte {run.loss:.3f}", fontsize=10)
ax.set_ylim(-2, 2)
_, ax = plt.subplots(figsize=(7, 3.5))
trace_predictions(ax, run_tanh)
ax.legend(fontsize=8)
plt.show()
⇑ Le meilleur modèle suit bien la loi au centre ; il s'en écarte au bord gauche, où les mesures sont rares.
Changer la fonction d'activation¶
L'activation se choisit par son nom. Essayons relu, qui est affine par
morceaux, et gelu, une version adoucie de relu. La liste complète est
dans yax.activations ; on peut aussi passer
directement une fonction jax.
run_relu = entraine("relu", Regresseur(largeur=128, activation="relu", rate=0.0, rkey=jr.key(0)),
config, (X_train, Y_train))
run_gelu = entraine("gelu", Regresseur(largeur=128, activation="gelu", rate=0.0, rkey=jr.key(0)),
config, (X_train, Y_train))
relu : perte 0.1268 après 1000 époques, dans out/regression/1 gelu : perte 0.1186 après 1000 époques, dans out/regression/2
_, axs = plt.subplots(1, 3, figsize=(12, 3), sharey=True)
for ax, run in zip(axs, [run_tanh, run_relu, run_gelu]):
trace_predictions(ax, run)
plt.show()
⇑ Avec relu, la courbe est faite de segments : on voit ses coudes. gelu,
plus douce, donne ici la meilleure perte.
Limiter le sur-apprentissage¶
yax.training.train garde le meilleur modèle : il arrête de fait
l'apprentissage au bon moment. Mais un modèle qui sur-apprend vite est
fragile, et son meilleur moment dépend du hasard. Deux remèdes classiques
freinent le sur-apprentissage lui-même :
- le dropout : pendant l'entraînement, une fraction
ratedes neurones est annulée à chaque pas, au hasard. Le réseau ne peut plus compter sur un neurone précis pour retenir un point ; - la pénalisation des poids (weight decay) : l'optimiseur
"adamw"tire les poids vers zéro à chaque pas, ce qui favorise les fonctions douces. Son intensité est une option de l'optimiseur.
run_dropout = entraine("tanh + dropout", Regresseur(largeur=128, activation="tanh", rate=0.1, rkey=jr.key(0)),
config, (X_train, Y_train))
config_adamw = yax.configs.TrainConfig(learning_rate=1e-2, nb_epochs=1000,
optimizer="adamw", optimizer_options={"weight_decay": 1.0})
run_adamw = entraine("tanh + weight decay", Regresseur(largeur=128, activation="tanh", rate=0.0, rkey=jr.key(0)),
config_adamw, (X_train, Y_train))
tanh + dropout : perte 0.1319 après 1000 époques, dans out/regression/3 tanh + weight decay : perte 0.1271 après 1000 époques, dans out/regression/4
_, ax = plt.subplots(figsize=(7, 3.5))
for run in [run_tanh, run_dropout, run_adamw]:
ax.plot(run.history.val_losses, label=run.title)
ax.axhline(plancher, color="gray", ls=":", label="plancher du bruit")
ax.set(yscale="log", xlabel="époque", ylabel="perte de validation")
ax.legend()
plt.show()
⇑ Sans remède, la validation remonte après l'époque 200. Le dropout freine cette remontée (0.18 en fin d'entraînement au lieu de 0.35), au prix d'une courbe agitée : chaque pas tire un masque différent. La pénalisation des poids la supprime : la courbe reste à plat, vers 0.13.
Distribuer les données autrement¶
training et validation de yax.training.train acceptent la même chose :
(X, Y): toutes les données en un seul lot, un pas par époque ;(X, Y, batch_size): mélange à chaque époque, puis lots (batchs) de cette taille ;- un sampler : tout objet qui a un attribut
nb_batcheset qui, appelé avec une clé, produit les lots(x, y)d'une époque.
Le triplet est un raccourci pour yax.training.DatasetSampler :
sampler = yax.training.DatasetSampler(X_train, Y_train, batch_size=5)
x, y = next(iter(sampler(jr.key(0))))
print(sampler, "| un lot :", x.shape, y.shape)
DatasetSampler(nb_data=20, batch_size=5, nb_batches=4) | un lot : (5, 1) (5, 1)
Avec des lots de 5, une époque fait 4 pas : 250 époques font autant de pas que les 1000 époques précédentes. Chaque pas voit un lot différent, le gradient est plus bruité. On en profite pour faire décroître le pas jusqu'au dixième de sa valeur initiale.
config_lots = yax.configs.TrainConfig(learning_rate=1e-2, nb_epochs=250, lr_final_ratio=0.1)
run_lots = entraine("tanh, lots de 5", Regresseur(largeur=128, activation="tanh", rate=0.0, rkey=jr.key(0)),
config_lots, (X_train, Y_train, 5))
tanh, lots de 5 : perte 0.1265 après 250 époques, dans out/regression/5
Et si l'on pouvait mesurer à volonté ? yax.training.FunctionSampler(f, nb_batches) fabrique chaque lot par calcul, f(rkey) -> (x, y), sans jeu de
données : c'est le cas d'un simulateur, ou d'un problème physique où l'on
tire des points au hasard (y peut alors valoir None). Ici, chaque lot est
un nouveau tirage de 20 mesures : le modèle ne voit jamais deux fois le même
point, il ne peut pas apprendre le bruit par cœur.
Les époques n'ont plus de fin naturelle : on en autorise 500, avec une
patience de 50.
mesures_a_volonte = yax.training.FunctionSampler(lambda rkey: mesure(rkey, 20), nb_batches=5)
config_patience = yax.configs.TrainConfig(learning_rate=1e-2, nb_epochs=500, patience=50)
run_volonte = entraine("tanh, mesures à volonté", Regresseur(largeur=128, activation="tanh", rate=0.0, rkey=jr.key(0)),
config_patience, mesures_a_volonte)
tanh, mesures à volonté : perte 0.1140 après 305 époques, dans out/regression/6
run_volonte.history.plot(title=run_volonte.title)
plt.show()
⇑ La perte d'entraînement est très agitée, puisque chaque lot est nouveau. Mais la validation ne remonte pas : sans données répétées, pas de sur-apprentissage. La patience a arrêté l'entraînement après 305 époques, 50 époques après le dernier record.
Comparer les runs¶
Chaque appel à yax.training.train a créé un sous-dossier numéroté du
mother_folder. Relisons-les tous depuis le disque, comme on le ferait le
lendemain dans un autre notebook : yax.training.load_runs rend la liste
des runs, dans l'ordre de leur création. Chacun est un yax.training.Run,
avec son titre, le modèle, la configuration et l'historique.
yax.training.find_best_run désigne le meilleur.
runs = yax.training.load_runs(mother_folder)
meilleur = yax.training.find_best_run(mother_folder)
for run in runs:
marque = " ← le meilleur" if run.folder == meilleur else ""
print(f"{run.folder:18} {run.title:24} {run.config.optimizer:6} perte {run.loss:.4f}{marque}")
out/regression/0 tanh adam perte 0.1343 out/regression/1 relu adam perte 0.1268 out/regression/2 gelu adam perte 0.1186 out/regression/3 tanh + dropout adam perte 0.1319 out/regression/4 tanh + weight decay adamw perte 0.1271 out/regression/5 tanh, lots de 5 adam perte 0.1265 out/regression/6 tanh, mesures à volonté adam perte 0.1140 ← le meilleur
fig, axs = plt.subplots(2, 4, figsize=(14, 6), sharex=True, sharey=True)
for ax, run in zip(axs.flat, runs):
trace_predictions(ax, run)
axs.flat[-1].axis("off")
plt.show()
⇑ Le meilleur run est celui qui disposait de mesures à volonté (0.114, tout près du plancher 0.109) : davantage de données reste le meilleur remède au sur-apprentissage. Avec 20 mesures seulement, les runs se tiennent de près, entre 0.119 et 0.134.
Écrire son propre objectif¶
Un objectif a la signature objective_fn(model, x, y, rkey) et reçoit un lot ;
yax.batch_apply(model, x, rkey) applique le modèle à chaque exemple en
distribuant la clé.
Supposons que trois des vingt mesures soient aberrantes : un capteur a
déraillé. L'erreur quadratique punit très fort les grands écarts, et le
modèle se tord pour s'approcher de ces trois points. L'erreur absolue
moyenne les pèse beaucoup moins. La voici écrite à la main — yax la fournit
aussi, sous le nom yax.obj.mae :
def perte_absolue(model, x, y, rkey):
predictions = yax.batch_apply(model, x, rkey)
return jnp.mean(jnp.abs(predictions - y))
Y_aberrant = Y_train.at[:3].add(2.0) # trois mesures faussées
La perte de validation change avec l'objectif : les deux runs vont chacun dans
leur propre mother_folder.
run_quadratique = yax.training.train("out/aberrant_quadratique", config, yax.obj.mse,
Regresseur(largeur=128, activation="tanh", rate=0.0, rkey=jr.key(0)),
(X_train, Y_aberrant), (X_val, Y_val), rkey=jr.key(3), title="erreur quadratique")
run_absolue = yax.training.train("out/aberrant_absolue", config, perte_absolue,
Regresseur(largeur=128, activation="tanh", rate=0.0, rkey=jr.key(0)),
(X_train, Y_aberrant), (X_val, Y_val), rkey=jr.key(3), title="erreur absolue")
_, axs = plt.subplots(1, 2, figsize=(10, 3), sharey=True)
for ax, run in zip(axs, [run_quadratique, run_absolue]):
trace_predictions(ax, run)
ax.plot(X_train[:3], Y_aberrant[:3], "x", color="black", ms=8)
ax.set_ylim(-2, 3)
plt.show()
⇑ Avec l'erreur quadratique, la courbe se tord vers les deux mesures aberrantes de droite : elle monte à 1.4, quand la loi plafonne à 0.9. L'erreur absolue les ignore presque et suit la loi. Les deux pertes affichées ne se comparent pas : l'une est quadratique, l'autre absolue.
Sous le capot¶
Le gradient, sans machinerie¶
Puisque les feuilles sont exactement les paramètres, jax.grad d'une perte
par rapport au modèle rend… un modèle de même structure, dont les feuilles
sont les gradients. Rien à filtrer, rien à envelopper — et optax l'accepte
tel quel. C'est ce que fait yax.training.train à chaque pas ; vous pouvez
donc écrire votre propre boucle quand la sienne ne convient pas.
grads = jax.grad(yax.obj.mse)(model, X_train, Y_train, jr.key(0))
print("gradient de la tête :", grads.tete.weight.shape, grads.tete.bias.shape)
optimizer = optax.adam(1e-3)
opt_state = optimizer.init(model) # optax accepte le modèle tel quel
updates, opt_state = optimizer.update(grads, opt_state)
model_apres = optax.apply_updates(model, updates) # un nouveau modèle
gradient de la tête : (128, 1) (1,)
Modifier un modèle : yax.tree_at¶
Un module est immuable : on ne lui affecte rien après construction. Pour
changer une feuille, yax.tree_at rend une copie où elle est remplacée ;
l'original ne bouge pas. Ajoutons 1 au biais de la tête du premier modèle :
toute sa courbe monte d'autant.
modele = run_tanh.trained_model
modele_decale = yax.tree_at(lambda m: m.tete.bias, modele, modele.tete.bias + 1.0)
print("écart des prédictions :", yax.batch_apply(modele_decale, x_grille[:3]) - yax.batch_apply(modele, x_grille[:3]))
print("biais de l'original :", modele.tete.bias)
écart des prédictions : [[1.] [1.] [1.]] biais de l'original : [0.07900926]
Et ensuite ?¶
- yax.layers et yax.models pour les briques et les modèles complets ; yax.preprocessing pour l'enrichissement d'images, les masques d'attention et le fenêtrage des séries.
- Les démos : un script par famille de modèles, à lire comme des exemples complets.