Problème de gradient qui s'annule
L'autre problème possible avec les gradients survient lorsqu'ils s'annulent, c'est-à-dire qu'ils tendent vers zéro. C'est bien plus difficile à résoudre, surtout parce qu'il est moins facile à détecter. Si la fonction de perte ne s'améliore pas à chaque étape, est-ce parce que les gradients sont tombés à zéro et n'ont donc pas mis à jour les poids? Ou est-ce plutôt que le modèle n'arrive pas à apprendre?
Ce problème survient plus souvent dans les modèles RNN lorsque l'on a besoin d'une mémoire longue (phrases longues).
Dans cet exercice, vous allez observer ce phénomène sur les données IMDB, en sélectionnant des phrases plus longues. Les données sont chargées dans les variables X et y, et les classes Sequential, SimpleRNN, Dense ainsi que matplotlib.pyplot sous le nom plt sont disponibles. Le modèle a été préentraîné pendant 100 époques; ses poids et son historique sont enregistrés dans le fichier model_weights.h5 et la variable history.
Cette activité fait partie du cours
Réseaux de neurones récurrents (RNN) pour la modélisation du langage avec Keras
Instructions de l’exercice
- Ajoutez une couche
SimpleRNNau modèle. - Chargez les poids préentraînés dans le modèle avec la méthode
.load_weights(). - Ajoutez au graphique l'exactitude des données d'entraînement disponible sous l'attribut
'acc'. - Affichez le graphique avec la méthode
.show().
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
# Create the model
model = Sequential()
model.add(____(units=600, input_shape=(None, 1)))
model.add(Dense(1, activation='sigmoid'))
model.compile(loss='binary_crossentropy', optimizer='sgd', metrics=['accuracy'])
# Load pre-trained weights
model.____('model_weights.h5')
# Plot the accuracy x epoch graph
plt.plot(history.history[____])
plt.plot(history.history['val_acc'])
plt.legend(['train', 'val'], loc='upper left')
plt.____()