Flera viktuppdateringar i rad
Nu ska du göra flera uppdateringar för att avsevärt förbättra modellens vikter och se hur prediktionerna förbättras för varje steg.
För att hålla koden ren finns en fördefinierad funktion get_slope() som tar input_data, target och weights som argument. Det finns också en funktion get_mse() som tar samma argument. input_data, target och weights är fördefinierade.
Detta nätverk har inga dolda lager – det går direkt från indata (med 3 noder) till en utdatanod. Observera att weights är en enda array.
Även matplotlib.pyplot är fördefinierat, och felhistoriken plottas när du har genomfört dina steg med gradientmetoden.
Den här övningen är en del av kursen
Introduktion till djupinlärning i Python
Övningsinstruktioner
- Använd en
for-slinga för att iterativt uppdatera vikterna:- Beräkna lutningen med funktionen
get_slope(). - Uppdatera vikterna med en inlärningshastighet på
0.01. - Beräkna medelkvadratfelet (
mse) med de uppdaterade vikterna med hjälp av funktionenget_mse(). - Lägg till
mseimse_hist.
- Beräkna lutningen med funktionen
- Klicka på Skicka in svar för att visualisera
mse_hist. Vilken trend kan du se?
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
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()