Zacznij terazZacznij za darmo

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

Zobacz kurs

Instrukcje do ćwiczenia

  • Dodaj warstwę SimpleRNN do 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.____()
Edytuj i uruchom kod