НачатьНачать бесплатно

Проблема затухающих градиентов

Ещё одна возможная проблема — затухание градиентов, то есть их стремление к нулю. Обнаружить её значительно сложнее, чем взрывной рост. Если функция потерь не улучшается на каждом шаге, причина может быть двоякой: либо градиенты обнулились и веса не обновились, либо модель попросту не способна обучиться.

Эта проблема чаще всего возникает в моделях 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.____()
Редактировать и запускать код