Problem zanikających gradientów
Drugim możliwym problemem z gradientami jest ich zanikanie, czyli zbieganie do zera. To znacznie trudniejszy problem do rozwiązania, ponieważ nie jest łatwy do wykrycia. Jeśli funkcja straty nie poprawia się z każdym krokiem, czy to dlatego, że gradienty zanikły i nie zaktualizowały wag? A może model po prostu nie jest w stanie się uczyć?
Problem ten częściej występuje w modelach RNN, gdy wymagana jest długa pamięć (np. przy przetwarzaniu długich zdań).
W tym ćwiczeniu zaobserwujesz ten problem na danych IMDB, dla których wybrano dłuższe zdania. Dane są załadowane do zmiennych X i y, a także zaimportowane klasy Sequential, SimpleRNN, Dense oraz matplotlib.pyplot jako plt. Model został wstępnie wytrenowany przez 100 epok – jego wagi i historia są zapisane w pliku model_weights.h5 oraz zmiennej history.
To ćwiczenie jest częścią kursu
Rekurencyjne sieci neuronowe (RNN) do modelowania języka w Keras
Instrukcje do ćwiczenia
- Dodaj warstwę
SimpleRNNdo modelu. - Wczytaj wstępnie wytrenowane wagi do modelu, używając metody
.load_weights(). - Dodaj do wykresu dokładność na danych treningowych, dostępną pod atrybutem
'acc'. - Wyświetl wykres, używając metody
.show().
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# 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.____()