Wielokrotna aktualizacja wag
Teraz wykonasz wiele aktualizacji wag, co pozwoli ci znacznie poprawić model i zobaczyć, jak prognozy ulepszają się z każdą kolejną aktualizacją.
Aby kod był przejrzysty, przygotowano gotową funkcję get_slope(), która przyjmuje argumenty input_data, target i weights. Dostępna jest również funkcja get_mse() przyjmująca te same argumenty. Zmienne input_data, target i weights są wstępnie załadowane.
Ta sieć nie ma żadnych ukrytych warstw – dane przechodzą bezpośrednio z wejścia (złożonego z 3 węzłów) do węzła wyjściowego. Zwróć uwagę, że weights to pojedyncza tablica.
Załadowano również bibliotekę matplotlib.pyplot – historia błędów zostanie przedstawiona na wykresie po wykonaniu kroków gradientu prostego.
To ćwiczenie jest częścią kursu
Wprowadzenie do uczenia głębokiego w Pythonie
Instrukcje do ćwiczenia
- Użyj pętli
for, aby iteracyjnie aktualizować wagi:- Oblicz nachylenie, korzystając z funkcji
get_slope(). - Zaktualizuj wagi, stosując współczynnik uczenia
0.01. - Oblicz błąd średniokwadratowy (
mse) na podstawie zaktualizowanych wag, używając funkcjiget_mse(). - Dołącz
msedomse_hist.
- Oblicz nachylenie, korzystając z funkcji
- Kliknij „Prześlij odpowiedź", aby zwizualizować
mse_hist. Jaką tendencję zauważasz?
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
n_updates = 20
mse_hist = []
# Iterate over the number of updates
for i in range(n_updates):
# Calculate the slope: slope
slope = ____(____, ____, ____)
# Update the weights: weights
weights = ____ - ____ * ____
# Calculate mse with new weights: mse
mse = ____(____, ____, ____)
# Append the mse to mse_hist
____
# Plot the mse history
plt.plot(mse_hist)
plt.xlabel('Iterations')
plt.ylabel('Mean Squared Error')
plt.show()