Zacznij terazZacznij za darmo

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

Zobacz kurs

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 funkcji get_mse().
    • Dołącz mse do mse_hist.
  • 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()
Edytuj i uruchom kod