CommencezCommencez gratuitement

Réduire le surapprentissage avec le dropout

Un problème courant avec les réseaux de neurones est leur tendance à faire du surapprentissage des données d'entraînement. Cela signifie que la mesure de pointage, comme R\(^2\) ou l'exactitude, est élevée pour l'ensemble d'entraînement, mais faible pour les ensembles de test et de validation, et que le modèle s'ajuste au bruit présent dans les données d'entraînement.

Pour prévenir le surapprentissage, on peut utiliser le dropout. Cette technique élimine au hasard certains neurones pendant l'entraînement, ce qui empêche le réseau de s'ajuster au bruit des données d'entraînement. keras offre une couche Dropout que nous pouvons utiliser à cette fin. Il faut définir le taux de dropout, c'est-à-dire la fraction de connexions éliminées pendant l'entraînement. Ce taux se fixe avec un nombre décimal entre 0 et 1 dans la couche Dropout().

Pour ce modèle, nous allons revenir à la fonction de perte de l'erreur quadratique moyenne.

Cette activité fait partie du cours

Machine Learning pour la finance en Python

Voir le cours

Instructions de l’exercice

  • Ajoutez une couche de dropout (Dropout()) après la première couche Dense du modèle et utilisez 20 % (0,2) comme taux de dropout.
  • Utilisez l'optimiseur adam et la fonction de perte mse lors de la compilation du modèle avec .compile().
  • Entraînez le modèle sur scaled_train_features et train_targets pendant 25 époques.

Exercice interactif pratique

Essayez cet exercice en complétant ce code d’exemple.

from keras.layers import Dropout

# Create model with dropout
model_3 = Sequential()
model_3.add(Dense(100, input_dim=scaled_train_features.shape[1], activation='relu'))
model_3.add(____)
model_3.add(Dense(20, activation='relu'))
model_3.add(Dense(1, activation='linear'))

# Fit model with mean squared error loss function
model_3.compile(optimizer=____, loss=____)
history = model_3.fit(____, ____, epochs=____)
plt.plot(history.history['loss'])
plt.title('loss:' + str(round(history.history['loss'][-1], 6)))
plt.show()
Modifier et exécuter le code