Проблема затухающих градиентов
Ещё одна возможная проблема — затухание градиентов, то есть их стремление к нулю. Обнаружить её значительно сложнее, чем взрывной рост. Если функция потерь не улучшается на каждом шаге, причина может быть двоякой: либо градиенты обнулились и веса не обновились, либо модель попросту не способна обучиться.
Эта проблема чаще всего возникает в моделях RNN, когда требуется долгая память — например, при работе с длинными предложениями.
В этом упражнении вы исследуете данную проблему на данных IMDB, где намеренно отобраны более длинные предложения. Данные загружены в переменные X и y, классы Sequential, SimpleRNN, Dense и библиотека matplotlib.pyplot импортированы как plt. Модель предварительно обучена в течение 100 эпох; её веса сохранены в файле model_weights.h5, а история обучения — в переменной history.
Это упражнение является частью курса
Рекуррентные нейронные сети (RNN) для языкового моделирования с Keras
Инструкции к упражнению
- Добавьте слой
SimpleRNNк модели. - Загрузите предварительно обученные веса в модель с помощью метода
.load_weights(). - Добавьте на график значения точности на обучающих данных, доступные через атрибут
'acc'. - Отобразите график с помощью метода
.show().
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# 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.____()