Kom igångKom igång gratis

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

Visa kurs

Ö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 funktionen get_mse().
    • Lägg till mse i mse_hist.
  • 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()
Redigera och kör kod