Aller au contenu principal

Module 9 — Entraînement distribué sur plusieurs accélérateurs

Quand un modèle tient en mémoire mais que l'entraînement prend trois jours, la distribution devient intéressante. TensorFlow la rend étonnamment simple à écrire, ce qui masque deux réglages qu'il faut absolument ajuster à la main.

Le principe : répliquer le modèle, partager le lot

La stratégie la plus courante réplique le modèle entier sur chaque accélérateur. Chaque réplique reçoit une fraction du lot, calcule ses gradients, puis tous les gradients sont moyennés et appliqués de façon identique partout. Les répliques restent ainsi rigoureusement synchronisées.

C'est le parallélisme de données. Il suppose que le modèle tient dans la mémoire d'un seul accélérateur. Quand ce n'est plus le cas — les grands modèles de langage du cours 16 — il faut du parallélisme de modèle, qui découpe le réseau lui-même et sort du cadre de ce module.

Trois lignes de code

import tensorflow as tf
from tensorflow import keras

strategie = tf.distribute.MirroredStrategy()
print(f"Repliques : {strategie.num_replicas_in_sync}")

with strategie.scope():
modele = construire_modele()
modele.compile(
optimizer=keras.optimizers.Adam(1e-3),
loss="sparse_categorical_crossentropy",
metrics=["accuracy"],
)

modele.fit(jeu, validation_data=jeu_val, epochs=20)

Tout ce qui crée des variables doit se trouver dans le scope : la construction du modèle et la compilation. Le fit, lui, reste à l'extérieur. Une variable créée hors du champ n'est pas répliquée, et l'erreur qui en résulte est difficile à relier à sa cause.

StratégiePortéeUsage
MirroredStrategyplusieurs accélérateurs, une machinele cas courant
MultiWorkerMirroredStrategyplusieurs machinesexige une configuration réseau
TPUStrategyunités de traitement tensorielaprès connexion au cluster
OneDeviceStrategyun seul dispositifdéboguer le code distribué

La taille de lot est globale, et c'est contre-intuitif

Premier réglage à ne pas manquer. Le batch_size que vous fournissez est la taille globale : elle est divisée entre les répliques.

LOT_PAR_REPLIQUE = 64
LOT_GLOBAL = LOT_PAR_REPLIQUE * strategie.num_replicas_in_sync

jeu = jeu.batch(LOT_GLOBAL).prefetch(tf.data.AUTOTUNE)

Conserver batch_size=64 sur quatre accélérateurs donne seize exemples par réplique. Le calcul devient inefficace — les accélérateurs sont sous-occupés — et les statistiques de normalisation par lots se dégradent, puisqu'elles sont calculées par réplique et non globalement. Sous trente-deux exemples par réplique, la normalisation par lots devient franchement bruitée.

Le lot global doit être divisible par le nombre de répliques

Un reste provoque des répliques de tailles inégales, et selon les versions soit une erreur, soit un déséquilibre silencieux. Calculez toujours le lot global à partir du lot par réplique, jamais l'inverse.

Le taux d'apprentissage doit suivre

Second réglage, celui qu'on oublie le plus souvent. Multiplier la taille de lot par quatre divise par quatre le nombre de mises à jour par époque. À taux constant, le modèle apprend donc quatre fois moins par époque, et l'entraînement paraît régresser alors qu'il est simplement ralenti.

Deux règles de mise à l'échelle circulent :

  • linéaire : multiplier le taux par le nombre de répliques. Convient jusqu'à des lots de quelques milliers d'exemples.
  • racine carrée : multiplier par la racine du facteur. Plus prudent sur les très grands lots.
TAUX_BASE = 1e-3
taux = TAUX_BASE * strategie.num_replicas_in_sync

with strategie.scope():
modele.compile(optimizer=keras.optimizers.Adam(taux), loss="mse")

Un taux multiplié par quatre appliqué dès le premier lot déstabilise l'entraînement. La montée en régime évoquée au module 6 devient ici quasi indispensable : partir du taux de base et atteindre le taux mis à l'échelle en quelques centaines de pas.

La précision mixte, souvent plus rentable que la distribution

Avant d'ajouter des accélérateurs, il existe un levier moins coûteux. La précision mixte stocke les poids en float32 mais effectue les calculs en float16, ce qui exploite les unités matérielles dédiées et divise à peu près par deux la mémoire occupée.

keras.mixed_precision.set_global_policy("mixed_float16")

Une précaution accompagne ce réglage : la couche de sortie doit rester en float32. En float16, une softmax sature et l'entropie croisée perd toute précision numérique.

sortie = layers.Dense(nb_classes, activation="softmax", dtype="float32")(x)

Keras gère de lui-même la mise à l'échelle de la perte, qui évite que de petits gradients ne deviennent nuls dans la plage réduite du float16. Si vous avez redéfini train_step comme au module 4, cette mise à l'échelle devient votre responsabilité : optimizer.get_scaled_loss avant la dérivation, get_unscaled_gradients après.

Dans quel ordre optimiser

Le premier goulot est presque toujours le pipeline de données du module 5, et le corriger ne coûte rien. Vient ensuite la précision mixte, qui apporte souvent un facteur proche de deux pour une ligne de code. La distribution arrive en troisième position : elle multiplie le coût matériel et introduit des réglages supplémentaires. Vérifiez avec le profileur du module 7 que le calcul est bien le goulot avant d'y venir.

Ce qui change à l'échelle de plusieurs machines

MultiWorkerMirroredStrategy étend le principe entre machines, mais ajoute des contraintes qui n'existaient pas. La variable d'environnement TF_CONFIG doit décrire la topologie du cluster sur chaque nœud. La bande passante réseau devient déterminante, puisque les gradients traversent le réseau à chaque pas. Et surtout, l'écriture des points de contrôle doit être coordonnée : chaque travailleur écrit dans un répertoire temporaire distinct, et seul le travailleur principal conserve le fichier final. keras.callbacks.BackupAndRestore gère cette coordination et permet de reprendre après la panne d'un nœud, ce qui devient statistiquement inévitable au-delà de quelques machines.

En résumé

  • Le parallélisme de données réplique le modèle entier et moyenne les gradients ; il exige que le modèle tienne dans un seul accélérateur.
  • Tout ce qui crée des variables va dans le scope de la stratégie ; fit reste à l'extérieur.
  • Le batch_size est global : le calculer à partir du lot par réplique, sous peine de sous-occuper le matériel et de dégrader la normalisation par lots.
  • Le taux d'apprentissage doit être mis à l'échelle avec le nombre de répliques, et accompagné d'une montée en régime ; avant de distribuer, tester le pipeline de données puis la précision mixte, qui coûtent bien moins cher.

Module suivant : exporter le modèle au format SavedModel et le servir en production.